337 lines
11 KiB
Python
337 lines
11 KiB
Python
"""Synchronous HTTP client for the DeerFlow Agent Profile Open API."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import secrets
|
|
from dataclasses import dataclass
|
|
from typing import Any, Dict, Iterator, List, Literal, Optional
|
|
from urllib.parse import quote
|
|
|
|
import httpx
|
|
|
|
|
|
AuthMode = Literal["bearer", "x-api-key"]
|
|
|
|
|
|
class OpenAgentAPIError(RuntimeError):
|
|
"""Raised when the API returns a non-success HTTP response."""
|
|
|
|
def __init__(
|
|
self,
|
|
status_code: int,
|
|
detail: Any,
|
|
response_text: str,
|
|
*,
|
|
request_id: Optional[str] = None,
|
|
) -> None:
|
|
self.status_code = status_code
|
|
self.detail = detail
|
|
self.response_text = response_text
|
|
self.request_id = request_id
|
|
suffix = f" (request_id={request_id})" if request_id else ""
|
|
super().__init__(f"Open Agent API request failed with HTTP {status_code}: {detail}{suffix}")
|
|
|
|
@property
|
|
def retryable(self) -> bool:
|
|
return self.status_code in {408, 429} or self.status_code >= 500
|
|
|
|
@property
|
|
def active_run_conflict(self) -> bool:
|
|
return self.status_code == 409
|
|
|
|
|
|
class OpenAgentTransportError(RuntimeError):
|
|
"""Raised when the HTTP exchange fails before a valid API response is received."""
|
|
|
|
|
|
class OpenAgentProtocolError(RuntimeError):
|
|
"""Raised when a successful API response does not match the documented shape."""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class OpenAgentEvent:
|
|
event: Optional[str]
|
|
data: Any
|
|
event_id: Optional[str] = None
|
|
retry_ms: Optional[int] = None
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
value: Dict[str, Any] = {"event": self.event, "data": self.data}
|
|
if self.event_id is not None:
|
|
value["id"] = self.event_id
|
|
if self.retry_ms is not None:
|
|
value["retry"] = self.retry_ms
|
|
return value
|
|
|
|
|
|
class OpenAgentAPIClient:
|
|
"""Small synchronous client matching the supplied Agent Open API contract."""
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
base_url: str,
|
|
api_key: str,
|
|
auth_mode: AuthMode = "bearer",
|
|
csrf_token: Optional[str] = None,
|
|
timeout: float = 60.0,
|
|
connect_timeout: float = 10.0,
|
|
http_client: Optional[httpx.Client] = None,
|
|
) -> None:
|
|
if auth_mode not in ("bearer", "x-api-key"):
|
|
raise ValueError("auth_mode must be 'bearer' or 'x-api-key'")
|
|
if not api_key or not api_key.strip():
|
|
raise ValueError("api_key must not be empty")
|
|
if timeout <= 0 or connect_timeout <= 0:
|
|
raise ValueError("timeouts must be greater than zero")
|
|
|
|
self.base_url = base_url.rstrip("/")
|
|
self.api_key = api_key.strip()
|
|
self.auth_mode = auth_mode
|
|
self.csrf_token = csrf_token or secrets.token_urlsafe(32)
|
|
self.timeout = httpx.Timeout(timeout, connect=connect_timeout)
|
|
self._client = http_client or httpx.Client(timeout=self.timeout)
|
|
self._owns_client = http_client is None
|
|
|
|
def close(self) -> None:
|
|
if self._owns_client:
|
|
self._client.close()
|
|
|
|
def __enter__(self) -> "OpenAgentAPIClient":
|
|
return self
|
|
|
|
def __exit__(self, exc_type: object, exc: object, traceback: object) -> None:
|
|
self.close()
|
|
|
|
def create_session(
|
|
self,
|
|
*,
|
|
external_subject_id: Optional[str] = None,
|
|
idempotency_key: Optional[str] = None,
|
|
metadata: Optional[Dict[str, Any]] = None,
|
|
) -> Dict[str, Any]:
|
|
payload = self._compact_payload(
|
|
{
|
|
"external_subject_id": external_subject_id,
|
|
"idempotency_key": idempotency_key,
|
|
"metadata": metadata,
|
|
}
|
|
)
|
|
return self._request_json("POST", "/api/open/agent-sessions", json_body=payload)
|
|
|
|
def send_message(
|
|
self,
|
|
session_id: str,
|
|
message: str,
|
|
*,
|
|
idempotency_key: Optional[str] = None,
|
|
metadata: Optional[Dict[str, Any]] = None,
|
|
) -> Dict[str, Any]:
|
|
payload = self._message_payload(message, idempotency_key=idempotency_key, metadata=metadata)
|
|
return self._request_json(
|
|
"POST",
|
|
f"/api/open/agent-sessions/{self._path_segment(session_id)}/messages",
|
|
json_body=payload,
|
|
)
|
|
|
|
def stream_message(
|
|
self,
|
|
session_id: str,
|
|
message: str,
|
|
*,
|
|
idempotency_key: Optional[str] = None,
|
|
metadata: Optional[Dict[str, Any]] = None,
|
|
include_trace: bool = False,
|
|
) -> Iterator[OpenAgentEvent]:
|
|
payload = self._message_payload(message, idempotency_key=idempotency_key, metadata=metadata)
|
|
path = f"/api/open/agent-sessions/{self._path_segment(session_id)}/messages/stream"
|
|
try:
|
|
with self._client.stream(
|
|
"POST",
|
|
self._url(path),
|
|
params={"include_trace": "true"} if include_trace else None,
|
|
headers=self._headers(include_csrf=True, accept="text/event-stream"),
|
|
json=payload,
|
|
timeout=self.timeout,
|
|
) as response:
|
|
self._raise_for_error(response)
|
|
yield from self._iter_sse_events(response)
|
|
except OpenAgentAPIError:
|
|
raise
|
|
except httpx.RequestError as exc:
|
|
raise OpenAgentTransportError(f"Open Agent API stream failed: {exc}") from exc
|
|
|
|
def get_run(self, session_id: str, run_id: str) -> Dict[str, Any]:
|
|
path = (
|
|
f"/api/open/agent-sessions/{self._path_segment(session_id)}"
|
|
f"/runs/{self._path_segment(run_id)}"
|
|
)
|
|
return self._request_json("GET", path)
|
|
|
|
def cancel_run(self, session_id: str, run_id: str) -> Dict[str, Any]:
|
|
path = (
|
|
f"/api/open/agent-sessions/{self._path_segment(session_id)}"
|
|
f"/runs/{self._path_segment(run_id)}/cancel"
|
|
)
|
|
return self._request_json("POST", path)
|
|
|
|
def _request_json(
|
|
self,
|
|
method: str,
|
|
path: str,
|
|
*,
|
|
json_body: Optional[Dict[str, Any]] = None,
|
|
) -> Dict[str, Any]:
|
|
try:
|
|
response = self._client.request(
|
|
method,
|
|
self._url(path),
|
|
headers=self._headers(
|
|
include_csrf=method.upper() in {"POST", "PUT", "PATCH", "DELETE"},
|
|
accept="application/json",
|
|
),
|
|
json=json_body,
|
|
timeout=self.timeout,
|
|
)
|
|
except httpx.RequestError as exc:
|
|
raise OpenAgentTransportError(f"Open Agent API request failed: {exc}") from exc
|
|
|
|
self._raise_for_error(response)
|
|
if not response.content:
|
|
return {}
|
|
try:
|
|
body = response.json()
|
|
except ValueError as exc:
|
|
raise OpenAgentProtocolError("Open Agent API returned invalid JSON") from exc
|
|
if not isinstance(body, dict):
|
|
raise OpenAgentProtocolError("Open Agent API JSON response must be an object")
|
|
return body
|
|
|
|
def _headers(self, *, include_csrf: bool, accept: str) -> Dict[str, str]:
|
|
headers = {
|
|
"Accept": accept,
|
|
"User-Agent": "arr-open-agent-integration/0.1",
|
|
}
|
|
if self.auth_mode == "x-api-key":
|
|
headers["X-DeerFlow-Open-API-Key"] = self.api_key
|
|
else:
|
|
headers["Authorization"] = f"Bearer {self.api_key}"
|
|
if include_csrf:
|
|
headers["X-CSRF-Token"] = self.csrf_token
|
|
headers["Cookie"] = f"csrf_token={self.csrf_token}"
|
|
return headers
|
|
|
|
def _url(self, path: str) -> str:
|
|
return f"{self.base_url}{path}"
|
|
|
|
@staticmethod
|
|
def _path_segment(value: str) -> str:
|
|
return quote(str(value), safe="")
|
|
|
|
@staticmethod
|
|
def _compact_payload(payload: Dict[str, Any]) -> Dict[str, Any]:
|
|
return {key: value for key, value in payload.items() if value is not None}
|
|
|
|
def _message_payload(
|
|
self,
|
|
message: str,
|
|
*,
|
|
idempotency_key: Optional[str],
|
|
metadata: Optional[Dict[str, Any]],
|
|
) -> Dict[str, Any]:
|
|
return self._compact_payload(
|
|
{
|
|
"message": message,
|
|
"idempotency_key": idempotency_key,
|
|
"metadata": metadata,
|
|
}
|
|
)
|
|
|
|
@staticmethod
|
|
def _iter_sse_events(response: httpx.Response) -> Iterator[OpenAgentEvent]:
|
|
event_name: Optional[str] = None
|
|
event_id: Optional[str] = None
|
|
retry_ms: Optional[int] = None
|
|
data_lines: List[str] = []
|
|
|
|
def finish_event() -> Optional[OpenAgentEvent]:
|
|
nonlocal event_name, event_id, retry_ms, data_lines
|
|
if event_name is None and event_id is None and retry_ms is None and not data_lines:
|
|
return None
|
|
data_text = "\n".join(data_lines)
|
|
if not data_text:
|
|
data: Any = None
|
|
else:
|
|
try:
|
|
data = json.loads(data_text)
|
|
except ValueError:
|
|
data = data_text
|
|
event = OpenAgentEvent(
|
|
event=event_name,
|
|
data=data,
|
|
event_id=event_id,
|
|
retry_ms=retry_ms,
|
|
)
|
|
event_name = None
|
|
event_id = None
|
|
retry_ms = None
|
|
data_lines = []
|
|
return event
|
|
|
|
for line in response.iter_lines():
|
|
if line == "":
|
|
event = finish_event()
|
|
if event is not None:
|
|
yield event
|
|
continue
|
|
if line.startswith(":"):
|
|
continue
|
|
|
|
field, separator, value = line.partition(":")
|
|
if not separator:
|
|
value = ""
|
|
elif value.startswith(" "):
|
|
value = value[1:]
|
|
|
|
if field == "event":
|
|
event_name = value
|
|
elif field == "data":
|
|
data_lines.append(value)
|
|
elif field == "id":
|
|
event_id = value
|
|
elif field == "retry":
|
|
try:
|
|
parsed_retry = int(value)
|
|
except ValueError:
|
|
continue
|
|
if parsed_retry >= 0:
|
|
retry_ms = parsed_retry
|
|
|
|
event = finish_event()
|
|
if event is not None:
|
|
yield event
|
|
|
|
@staticmethod
|
|
def _raise_for_error(response: httpx.Response) -> None:
|
|
if response.status_code < 400:
|
|
return
|
|
try:
|
|
response_text = response.text
|
|
except httpx.ResponseNotRead:
|
|
response.read()
|
|
response_text = response.text
|
|
try:
|
|
body = response.json()
|
|
except ValueError:
|
|
detail: Any = response_text or response.reason_phrase
|
|
else:
|
|
detail = body.get("detail", body) if isinstance(body, dict) else body
|
|
request_id = response.headers.get("x-request-id") or response.headers.get("request-id")
|
|
raise OpenAgentAPIError(
|
|
response.status_code,
|
|
detail,
|
|
response_text,
|
|
request_id=request_id,
|
|
)
|