refactor(plugin): update HTTP request/response handling and improve type hints

- Replaced `parse_raw_request` and `convert_response_to_raw_data` with `deserialize_request` and `serialize_response` for better clarity and functionality.
- Enhanced the `PluginExecutor` methods to utilize the new request and response utilities.
- Updated type hints in the `Subscription` class to use `Mapping` instead of `dict` for improved type safety.
- Refactored HTTP parsing logic to support a wider range of content types and improve overall request handling.
This commit is contained in:
Harry
2025-09-01 14:19:48 +08:00
parent e9df3cbd4a
commit e9975e101d
3 changed files with 127 additions and 94 deletions
+6 -6
View File
@@ -41,7 +41,7 @@ from dify_plugin.core.entities.plugin.request import (
)
from dify_plugin.core.plugin_registration import PluginRegistration
from dify_plugin.core.runtime import Session
from dify_plugin.core.utils.http_parser import convert_response_to_raw_data, parse_raw_request
from dify_plugin.core.utils.http_parser import deserialize_request, serialize_response
from dify_plugin.entities.agent import AgentRuntime
from dify_plugin.entities.tool import ToolRuntime
from dify_plugin.entities.trigger import Subscription, TriggerRuntime
@@ -295,7 +295,7 @@ class PluginExecutor:
def invoke_endpoint(self, session: Session, data: EndpointInvokeRequest):
bytes_data = binascii.unhexlify(data.raw_http_request)
request = parse_raw_request(bytes_data)
request = deserialize_request(bytes_data)
try:
# dispatch request
@@ -360,7 +360,7 @@ class PluginExecutor:
def get_oauth_credentials(self, session: Session, data: OAuthGetCredentialsRequest):
provider_instance = self._get_oauth_provider_instance(data.provider)
bytes_data = binascii.unhexlify(data.raw_http_request)
request = parse_raw_request(bytes_data)
request = deserialize_request(bytes_data)
credentials = provider_instance.oauth_get_credentials(data.redirect_uri, data.system_credentials, request)
@@ -424,7 +424,7 @@ class PluginExecutor:
session_id=session.session_id,
)
trigger = trigger_cls(runtime=trigger_runtime, session=session)
event = trigger.trigger(parse_raw_request(binascii.unhexlify(request.raw_http_request)), request.parameters)
event = trigger.trigger(deserialize_request(binascii.unhexlify(request.raw_http_request)), request.parameters)
return TriggerInvokeResponse(
event=event.model_dump(),
)
@@ -453,10 +453,10 @@ class PluginExecutor:
bytes_data = binascii.unhexlify(request.raw_http_request)
provider_instance = trigger_provider_cls()
subscription = Subscription(**request.subscription)
dispatch_result = provider_instance.dispatch_event(subscription, parse_raw_request(bytes_data))
dispatch_result = provider_instance.dispatch_event(subscription, deserialize_request(bytes_data))
return TriggerDispatchResponse(
triggers=dispatch_result.triggers,
raw_http_response=binascii.hexlify(convert_response_to_raw_data(dispatch_result.response)).decode(),
raw_http_response=binascii.hexlify(serialize_response(dispatch_result.response)).decode(),
)
def subscribe_trigger(self, session: Session, request: TriggerSubscribeRequest) -> TriggerSubscriptionResponse:
+119 -86
View File
@@ -1,53 +1,13 @@
from dpkt.http import Request as DpktRequest
from dpkt.http import Response as DpktResponse
from werkzeug.test import EnvironBuilder
from werkzeug.wrappers import Request, Response
from io import BytesIO
from typing import Any
from flask import Request, Response
from werkzeug.datastructures import Headers
def parse_raw_request(raw_data: bytes):
"""
Parse raw HTTP data into a Request object.
Supports various content types including:
- application/json
- multipart/form-data (file uploads)
- application/x-www-form-urlencoded
- text/plain and others
Args:
raw_data: The raw HTTP data as bytes.
Returns:
A Werkzeug Request object.
"""
req = DpktRequest(raw_data)
# Extract content type to handle different request types
content_type = req.headers.get('content-type', '')
# Build the request with proper handling for different content types
builder = EnvironBuilder(
method=req.method,
path=req.uri,
headers=dict(req.headers),
data=req.body,
)
return Request(builder.get_environ())
def convert_request_to_raw_data(request: Request) -> bytes:
def serialize_request(request: Request) -> bytes:
"""
Convert a Request object to raw HTTP data.
Properly handles various content types including:
- application/json
- application/x-www-form-urlencoded
- text/plain and others
NOTE: For multipart/form-data (file uploads), this function will include
the raw body if it's still available in the request stream. However,
if the request has already been parsed (e.g., form fields accessed),
the multipart body may not be fully reconstructible.
Args:
request: The Request object to convert.
@@ -57,81 +17,104 @@ def convert_request_to_raw_data(request: Request) -> bytes:
"""
# Start with the request line
method = request.method
path = request.full_path if request.query_string else request.path
protocol = "HTTP/1.1"
path = request.full_path
# Remove trailing ? if there's no query string
path = path.removesuffix("?")
protocol = request.headers.get("HTTP_VERSION", "HTTP/1.1")
raw_data = f"{method} {path} {protocol}\r\n".encode()
# Get the body first to ensure Content-Length is correct
body = request.get_data(as_text=False)
# Build headers dict, excluding some werkzeug-specific headers
headers_to_skip = {'HTTP_VERSION', 'REQUEST_METHOD', 'PATH_INFO', 'QUERY_STRING',
'SCRIPT_NAME', 'SERVER_NAME', 'SERVER_PORT', 'SERVER_PROTOCOL',
'wsgi.version', 'wsgi.url_scheme', 'wsgi.input', 'wsgi.errors',
'wsgi.multithread', 'wsgi.multiprocess', 'wsgi.run_once'}
# Add important headers
if request.content_type:
raw_data += f"Content-Type: {request.content_type}\r\n".encode()
if body:
raw_data += f"Content-Length: {len(body)}\r\n".encode()
# Add other headers
# Add headers
for header_name, header_value in request.headers.items():
# Skip already added headers and internal headers
if header_name.lower() not in ('content-type', 'content-length') and header_name not in headers_to_skip:
raw_data += f"{header_name}: {header_value}\r\n".encode()
raw_data += f"{header_name}: {header_value}\r\n".encode()
# Add empty line to separate headers from body
raw_data += b"\r\n"
# Add body if exists
body = request.get_data(as_text=False)
if body:
raw_data += body
return raw_data
def parse_raw_response(raw_data: bytes) -> Response:
def deserialize_request(raw_data: bytes) -> Request:
"""
Parse raw HTTP response data into a Response object.
Convert raw HTTP data to a Request object.
Args:
raw_data: The raw HTTP response data as bytes.
raw_data: The raw HTTP data as bytes.
Returns:
A Werkzeug Response object.
A Flask Request object.
"""
resp = DpktResponse(raw_data)
lines = raw_data.split(b"\r\n")
# Extract status code from the status line
status_code = resp.status
# Parse request line
request_line = lines[0].decode("utf-8")
parts = request_line.split(" ", 2) # Split into max 3 parts
if len(parts) < 2:
raise ValueError(f"Invalid request line: {request_line}")
# Create Response object with body and status
response = Response(
response=resp.body,
status=status_code,
headers=dict(resp.headers),
)
method = parts[0]
path = parts[1]
protocol = parts[2] if len(parts) > 2 else "HTTP/1.1"
return response
# Parse headers
headers = Headers()
body_start = 0
for i in range(1, len(lines)):
line = lines[i]
if line == b"":
body_start = i + 1
break
if b":" in line:
header_line = line.decode("utf-8")
name, value = header_line.split(":", 1)
headers[name.strip()] = value.strip()
# Extract body
body = b""
if body_start > 0 and body_start < len(lines):
body = b"\r\n".join(lines[body_start:])
# Create environ for Request
environ = {
"REQUEST_METHOD": method,
"PATH_INFO": path.split("?")[0] if "?" in path else path,
"QUERY_STRING": path.split("?")[1] if "?" in path else "",
"SERVER_NAME": headers.get("Host", "localhost").split(":")[0],
"SERVER_PORT": headers.get("Host", "localhost:80").split(":")[1] if ":" in headers.get("Host", "") else "80",
"SERVER_PROTOCOL": protocol,
"wsgi.input": BytesIO(body),
"wsgi.url_scheme": "https" if headers.get("X-Forwarded-Proto") == "https" else "http",
"CONTENT_LENGTH": str(len(body)) if body else "0",
"CONTENT_TYPE": headers.get("Content-Type", ""),
}
# Add headers to environ
for header_name, header_value in headers.items():
env_name = f"HTTP_{header_name.upper().replace('-', '_')}"
if header_name.upper() not in ["CONTENT-TYPE", "CONTENT-LENGTH"]:
environ[env_name] = header_value
return Request(environ)
def convert_response_to_raw_data(response: Response) -> bytes:
def serialize_response(response: Response) -> bytes:
"""
Convert a Response object to raw HTTP response data.
Convert a Response object to raw HTTP data.
Args:
response: The Response object to convert.
Returns:
The raw HTTP response data as bytes.
The raw HTTP data as bytes.
"""
# Start with the status line
protocol = "HTTP/1.1"
status_code = response.status_code
status_text = response.status or "OK"
status_text = response.status.split(" ", 1)[1] if " " in response.status else "OK"
raw_data = f"{protocol} {status_code} {status_text}\r\n".encode()
# Add headers
@@ -147,3 +130,53 @@ def convert_response_to_raw_data(response: Response) -> bytes:
raw_data += body
return raw_data
def deserialize_response(raw_data: bytes) -> Response:
"""
Convert raw HTTP data to a Response object.
Args:
raw_data: The raw HTTP data as bytes.
Returns:
A Flask Response object.
"""
lines = raw_data.split(b"\r\n")
# Parse status line
status_line = lines[0].decode("utf-8")
parts = status_line.split(" ", 2)
if len(parts) < 2:
raise ValueError(f"Invalid status line: {status_line}")
protocol = parts[0]
status_code = int(parts[1])
status_text = parts[2] if len(parts) > 2 else "OK"
# Parse headers
headers: dict[str, Any] = {}
body_start = 0
for i in range(1, len(lines)):
line = lines[i]
if line == b"":
body_start = i + 1
break
if b":" in line:
header_line = line.decode("utf-8")
name, value = header_line.split(":", 1)
headers[name.strip()] = value.strip()
# Extract body
body = b""
if body_start > 0 and body_start < len(lines):
body = b"\r\n".join(lines[body_start:])
# Create Response object
response = Response(
response=body,
status=status_code,
headers=headers,
)
return response
+2 -2
View File
@@ -308,12 +308,12 @@ class Subscription(BaseModel):
endpoint: str = Field(..., description="The webhook endpoint URL allocated by Dify for receiving events")
parameters: dict[str, Any] | None = Field(
parameters: Mapping[str, Any] | None = Field(
default=None,
description="""The parameters of the subscription, this is the creation parameters.
Only available when creating a new subscription by credentials(auto subscription), not manual subscription""",
)
properties: dict[str, Any] = Field(
properties: Mapping[str, Any] = Field(
..., description="Subscription data containing all properties and provider-specific information"
)