mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-15 15:40:20 +08:00
Merge pull request #418 from Avilianb/feat/codex-responses-websocket
feat: support Codex Responses WebSocket transport
This commit is contained in:
+7
-1
@@ -90,6 +90,11 @@ RUN printf '%s\n' \
|
||||
' "" $remote_addr;' \
|
||||
'}' \
|
||||
'' \
|
||||
'map $http_upgrade $connection_upgrade {' \
|
||||
' default upgrade;' \
|
||||
' "" "";' \
|
||||
'}' \
|
||||
'' \
|
||||
'server {' \
|
||||
' listen 80;' \
|
||||
' server_name _;' \
|
||||
@@ -160,7 +165,8 @@ RUN printf '%s\n' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Connection "";' \
|
||||
' proxy_set_header Upgrade $http_upgrade;' \
|
||||
' proxy_set_header Connection $connection_upgrade;' \
|
||||
' proxy_set_header Accept $http_accept;' \
|
||||
' proxy_set_header Content-Type $content_type;' \
|
||||
' proxy_set_header Authorization $http_authorization;' \
|
||||
|
||||
@@ -112,6 +112,11 @@ RUN printf '%s\n' \
|
||||
' "" $remote_addr;' \
|
||||
'}' \
|
||||
'' \
|
||||
'map $http_upgrade $connection_upgrade {' \
|
||||
' default upgrade;' \
|
||||
' "" "";' \
|
||||
'}' \
|
||||
'' \
|
||||
'server {' \
|
||||
' listen 80;' \
|
||||
' server_name _;' \
|
||||
@@ -182,7 +187,8 @@ RUN printf '%s\n' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Connection "";' \
|
||||
' proxy_set_header Upgrade $http_upgrade;' \
|
||||
' proxy_set_header Connection $connection_upgrade;' \
|
||||
' proxy_set_header Accept $http_accept;' \
|
||||
' proxy_set_header Content-Type $content_type;' \
|
||||
' proxy_set_header Authorization $http_authorization;' \
|
||||
|
||||
@@ -63,6 +63,7 @@ disable_response_storage = true
|
||||
name = "OpenAI"
|
||||
base_url = "${baseUrl.value}/v1"
|
||||
wire_api = "responses"
|
||||
supports_websockets = true
|
||||
requires_openai_auth = true`)
|
||||
|
||||
const codexAuthConfig = computed(() => `{
|
||||
|
||||
@@ -0,0 +1,529 @@
|
||||
"""Codex Responses WebSocket bridge.
|
||||
|
||||
The current Codex client can use a WebSocket transport for the Responses API:
|
||||
it sends JSON text request frames and expects each Responses streaming event as
|
||||
one JSON text message. Aether already speaks Responses over HTTP/SSE, so this
|
||||
module keeps the business path intact and only adapts the wire protocol.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import codecs
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.requests import Request
|
||||
from starlette.websockets import WebSocketState
|
||||
|
||||
from src.api.base.adapter import ApiMode
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.handlers.openai_cli import OpenAICliAdapter
|
||||
from src.core.logger import logger
|
||||
from src.database.database import create_session
|
||||
|
||||
CODEX_TURN_STATE_HEADER = "x-codex-turn-state"
|
||||
|
||||
_HOP_BY_HOP_WS_HEADERS = {
|
||||
"connection",
|
||||
"upgrade",
|
||||
"sec-websocket-accept",
|
||||
"sec-websocket-extensions",
|
||||
"sec-websocket-key",
|
||||
"sec-websocket-protocol",
|
||||
"sec-websocket-version",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CodexWsRequest:
|
||||
kind: str
|
||||
body: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class CodexWsProtocolError(ValueError):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
status_code: int = 400,
|
||||
error_type: str = "invalid_request_error",
|
||||
) -> None:
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
self.error_type = error_type
|
||||
|
||||
|
||||
def parse_codex_ws_request(text: str) -> CodexWsRequest:
|
||||
"""Parse one Codex WebSocket request text frame."""
|
||||
try:
|
||||
payload = json.loads(text)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise CodexWsProtocolError("Request frame must be valid JSON") from exc
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise CodexWsProtocolError("Request frame must be a JSON object")
|
||||
|
||||
kind = payload.get("type")
|
||||
if not isinstance(kind, str) or not kind:
|
||||
raise CodexWsProtocolError("Request frame must include a string type")
|
||||
|
||||
if kind == "response.processed":
|
||||
response_id = payload.get("response_id")
|
||||
if response_id is not None and not isinstance(response_id, str):
|
||||
raise CodexWsProtocolError("response.processed response_id must be a string")
|
||||
return CodexWsRequest(kind=kind)
|
||||
|
||||
if kind == "response.create":
|
||||
return CodexWsRequest(kind=kind, body=response_create_http_body(payload))
|
||||
|
||||
raise CodexWsProtocolError(f"Unsupported Codex websocket request type: {kind}")
|
||||
|
||||
|
||||
def response_create_http_body(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Convert a Codex response.create frame to a Responses HTTP body."""
|
||||
if payload.get("type") != "response.create":
|
||||
raise CodexWsProtocolError("Expected response.create request frame")
|
||||
|
||||
body = dict(payload)
|
||||
body.pop("type", None)
|
||||
if not body:
|
||||
raise CodexWsProtocolError("response.create request body cannot be empty")
|
||||
return body
|
||||
|
||||
|
||||
def extract_sse_payloads(block: str) -> list[str]:
|
||||
"""Extract JSON payload strings from one SSE event block."""
|
||||
data_lines: list[str] = []
|
||||
stripped_block = block.strip()
|
||||
for raw_line in block.splitlines():
|
||||
line = raw_line.strip()
|
||||
if not line or line.startswith(":"):
|
||||
continue
|
||||
if line.startswith("data:"):
|
||||
data_lines.append(line[5:].lstrip())
|
||||
elif line.startswith("event:") and " data:" in line:
|
||||
data_lines.append(line.split(" data:", 1)[1].strip())
|
||||
|
||||
if data_lines:
|
||||
payload = "\n".join(data_lines).strip()
|
||||
if payload and payload != "[DONE]":
|
||||
return [payload]
|
||||
return []
|
||||
|
||||
if stripped_block.startswith("{"):
|
||||
return [stripped_block]
|
||||
return []
|
||||
|
||||
|
||||
async def iter_sse_payloads(chunks: AsyncIterator[bytes | str]) -> AsyncIterator[str]:
|
||||
"""Yield JSON payload strings from an SSE byte stream."""
|
||||
decoder = codecs.getincrementaldecoder("utf-8")()
|
||||
buffer = ""
|
||||
|
||||
async for chunk in chunks:
|
||||
if isinstance(chunk, bytes):
|
||||
buffer += decoder.decode(chunk)
|
||||
else:
|
||||
buffer += chunk
|
||||
|
||||
buffer = buffer.replace("\r\n", "\n").replace("\r", "\n")
|
||||
while "\n\n" in buffer:
|
||||
block, buffer = buffer.split("\n\n", 1)
|
||||
for payload in extract_sse_payloads(block):
|
||||
yield payload
|
||||
|
||||
tail = decoder.decode(b"", final=True)
|
||||
if tail:
|
||||
buffer += tail
|
||||
buffer = buffer.replace("\r\n", "\n").replace("\r", "\n")
|
||||
if buffer.strip():
|
||||
for payload in extract_sse_payloads(buffer):
|
||||
yield payload
|
||||
|
||||
|
||||
def build_codex_ws_error_event(
|
||||
*,
|
||||
status_code: int,
|
||||
message: str,
|
||||
error_type: str | None = None,
|
||||
code: str | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build the wrapped error shape parsed by Codex's websocket client."""
|
||||
status_code = int(status_code or 500)
|
||||
if error_type is None:
|
||||
error_type = "server_error" if status_code >= 500 else "invalid_request_error"
|
||||
|
||||
error: dict[str, Any] = {
|
||||
"type": error_type,
|
||||
"message": message,
|
||||
}
|
||||
if code:
|
||||
error["code"] = code
|
||||
|
||||
event: dict[str, Any] = {
|
||||
"type": "error",
|
||||
"status": status_code,
|
||||
"error": error,
|
||||
}
|
||||
if headers:
|
||||
event["headers"] = dict(headers)
|
||||
return event
|
||||
|
||||
|
||||
async def handle_codex_responses_websocket(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
pipeline: ApiRequestPipeline,
|
||||
db_factory: Callable[[], Session] | None = None,
|
||||
) -> None:
|
||||
"""Serve Codex Responses WebSocket traffic on top of the HTTP pipeline."""
|
||||
session_factory = db_factory or create_session
|
||||
turn_state = websocket.headers.get(CODEX_TURN_STATE_HEADER) or str(uuid.uuid4())
|
||||
await websocket.accept(headers=[(CODEX_TURN_STATE_HEADER.encode(), turn_state.encode())])
|
||||
|
||||
logger.debug("[CodexWS] accepted websocket responses connection")
|
||||
|
||||
while True:
|
||||
try:
|
||||
message = await websocket.receive()
|
||||
except WebSocketDisconnect:
|
||||
logger.debug("[CodexWS] websocket disconnected")
|
||||
return
|
||||
|
||||
message_type = message.get("type")
|
||||
if message_type == "websocket.disconnect":
|
||||
logger.debug("[CodexWS] websocket disconnected")
|
||||
return
|
||||
|
||||
if "bytes" in message and message.get("bytes") is not None:
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=400,
|
||||
message="Binary websocket frames are not supported",
|
||||
error_type="invalid_request_error",
|
||||
)
|
||||
await _close_websocket(websocket, code=1003)
|
||||
return
|
||||
|
||||
text = message.get("text")
|
||||
if text is None:
|
||||
continue
|
||||
|
||||
try:
|
||||
request = parse_codex_ws_request(text)
|
||||
except CodexWsProtocolError as exc:
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=exc.status_code,
|
||||
message=str(exc),
|
||||
error_type=exc.error_type,
|
||||
)
|
||||
continue
|
||||
|
||||
if request.kind == "response.processed":
|
||||
logger.debug("[CodexWS] response.processed acknowledged")
|
||||
continue
|
||||
|
||||
if request.body is None:
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=400,
|
||||
message="response.create frame is missing a request body",
|
||||
)
|
||||
continue
|
||||
|
||||
await _run_response_create(
|
||||
websocket,
|
||||
pipeline=pipeline,
|
||||
db_factory=session_factory,
|
||||
body=request.body,
|
||||
)
|
||||
|
||||
|
||||
async def _run_response_create(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
pipeline: ApiRequestPipeline,
|
||||
db_factory: Callable[[], Session],
|
||||
body: dict[str, Any],
|
||||
) -> None:
|
||||
db = db_factory()
|
||||
response: Any = None
|
||||
try:
|
||||
http_request = _build_synthetic_request(websocket, body)
|
||||
adapter = OpenAICliAdapter()
|
||||
response = await pipeline.run(
|
||||
adapter=adapter,
|
||||
http_request=http_request,
|
||||
db=db,
|
||||
mode=ApiMode.PROXY,
|
||||
api_format_hint=adapter.allowed_api_formats[0],
|
||||
)
|
||||
db.commit()
|
||||
except HTTPException as exc:
|
||||
db.rollback()
|
||||
await _send_http_exception(websocket, exc)
|
||||
return
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception("[CodexWS] response.create failed")
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=500,
|
||||
message=str(exc) or "Internal server error",
|
||||
error_type="server_error",
|
||||
)
|
||||
return
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
await _send_pipeline_response(websocket, response)
|
||||
|
||||
|
||||
def _build_synthetic_request(websocket: WebSocket, body: dict[str, Any]) -> Request:
|
||||
body_bytes = json.dumps(body, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
headers = _build_http_headers_from_websocket(websocket)
|
||||
path = websocket.scope.get("path") or "/v1/responses"
|
||||
scheme = "https" if websocket.url.scheme == "wss" else "http"
|
||||
sent_body = False
|
||||
|
||||
async def receive() -> dict[str, Any]:
|
||||
nonlocal sent_body
|
||||
if not sent_body:
|
||||
sent_body = True
|
||||
return {
|
||||
"type": "http.request",
|
||||
"body": body_bytes,
|
||||
"more_body": False,
|
||||
}
|
||||
return {
|
||||
"type": "http.request",
|
||||
"body": b"",
|
||||
"more_body": False,
|
||||
}
|
||||
|
||||
scope: dict[str, Any] = {
|
||||
"type": "http",
|
||||
"asgi": websocket.scope.get("asgi", {"version": "3.0"}),
|
||||
"http_version": websocket.scope.get("http_version", "1.1"),
|
||||
"method": "POST",
|
||||
"scheme": scheme,
|
||||
"path": path,
|
||||
"raw_path": str(path).encode("ascii", errors="ignore"),
|
||||
"query_string": websocket.scope.get("query_string", b""),
|
||||
"root_path": websocket.scope.get("root_path", ""),
|
||||
"headers": headers,
|
||||
"client": websocket.scope.get("client"),
|
||||
"server": websocket.scope.get("server"),
|
||||
"state": {},
|
||||
}
|
||||
return Request(scope, receive)
|
||||
|
||||
|
||||
def _build_http_headers_from_websocket(websocket: WebSocket) -> list[tuple[bytes, bytes]]:
|
||||
headers: list[tuple[bytes, bytes]] = []
|
||||
seen: set[str] = set()
|
||||
for key, value in websocket.headers.items():
|
||||
lower_key = key.lower()
|
||||
if lower_key in _HOP_BY_HOP_WS_HEADERS:
|
||||
continue
|
||||
seen.add(lower_key)
|
||||
headers.append((lower_key.encode("latin-1"), value.encode("latin-1")))
|
||||
|
||||
if "content-type" not in seen:
|
||||
headers.append((b"content-type", b"application/json"))
|
||||
if "accept" not in seen:
|
||||
headers.append((b"accept", b"text/event-stream"))
|
||||
return headers
|
||||
|
||||
|
||||
async def _send_pipeline_response(websocket: WebSocket, response: Any) -> None:
|
||||
if isinstance(response, StreamingResponse):
|
||||
await _send_streaming_response(websocket, response)
|
||||
return
|
||||
|
||||
status_code = int(getattr(response, "status_code", 200) or 200)
|
||||
headers = _string_headers(getattr(response, "headers", None))
|
||||
content = _response_content(response)
|
||||
|
||||
if status_code >= 400:
|
||||
await _send_error_from_content(
|
||||
websocket,
|
||||
status_code=status_code,
|
||||
content=content,
|
||||
headers=headers,
|
||||
)
|
||||
return
|
||||
|
||||
if isinstance(content, dict) and content.get("type") == "response.completed":
|
||||
await websocket.send_text(json.dumps(content, ensure_ascii=False, separators=(",", ":")))
|
||||
return
|
||||
|
||||
if isinstance(content, dict):
|
||||
await websocket.send_text(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": content,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
)
|
||||
return
|
||||
|
||||
await websocket.send_text(json.dumps(content, ensure_ascii=False, separators=(",", ":")))
|
||||
|
||||
|
||||
async def _send_streaming_response(websocket: WebSocket, response: StreamingResponse) -> None:
|
||||
terminal_seen = False
|
||||
try:
|
||||
async for payload in iter_sse_payloads(response.body_iterator):
|
||||
if _is_terminal_payload(payload):
|
||||
terminal_seen = True
|
||||
await websocket.send_text(payload)
|
||||
except WebSocketDisconnect:
|
||||
logger.debug("[CodexWS] client disconnected while streaming response")
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.exception("[CodexWS] failed while streaming response")
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=500,
|
||||
message=str(exc) or "Streaming response failed",
|
||||
error_type="server_error",
|
||||
)
|
||||
return
|
||||
finally:
|
||||
background = getattr(response, "background", None)
|
||||
if background is not None:
|
||||
await background()
|
||||
|
||||
if not terminal_seen:
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=502,
|
||||
message="Stream closed before response.completed",
|
||||
error_type="server_error",
|
||||
)
|
||||
|
||||
|
||||
def _is_terminal_payload(payload: str) -> bool:
|
||||
try:
|
||||
event = json.loads(payload)
|
||||
except json.JSONDecodeError:
|
||||
return False
|
||||
if not isinstance(event, dict):
|
||||
return False
|
||||
return event.get("type") in {"response.completed", "error"}
|
||||
|
||||
|
||||
def _response_content(response: Any) -> Any:
|
||||
if isinstance(response, Response):
|
||||
body = getattr(response, "body", b"")
|
||||
if not body:
|
||||
return None
|
||||
try:
|
||||
return json.loads(body.decode("utf-8"))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError):
|
||||
return body.decode("utf-8", errors="replace")
|
||||
return response
|
||||
|
||||
|
||||
async def _send_http_exception(websocket: WebSocket, exc: HTTPException) -> None:
|
||||
detail = exc.detail
|
||||
headers = {str(k): str(v) for k, v in (exc.headers or {}).items()}
|
||||
if isinstance(detail, dict):
|
||||
await _send_error_from_content(
|
||||
websocket,
|
||||
status_code=exc.status_code,
|
||||
content=detail,
|
||||
headers=headers,
|
||||
)
|
||||
return
|
||||
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=exc.status_code,
|
||||
message=str(detail),
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
async def _send_error_from_content(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
status_code: int,
|
||||
content: Any,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> None:
|
||||
if isinstance(content, dict):
|
||||
error_obj = content.get("error")
|
||||
if isinstance(error_obj, dict):
|
||||
message = str(error_obj.get("message") or content)
|
||||
error_type = str(error_obj.get("type") or "invalid_request_error")
|
||||
code = error_obj.get("code")
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=status_code,
|
||||
message=message,
|
||||
error_type=error_type,
|
||||
code=str(code) if code is not None else None,
|
||||
headers=headers,
|
||||
)
|
||||
return
|
||||
if "detail" in content:
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=status_code,
|
||||
message=str(content["detail"]),
|
||||
headers=headers,
|
||||
)
|
||||
return
|
||||
|
||||
await _send_ws_error(
|
||||
websocket,
|
||||
status_code=status_code,
|
||||
message=str(content),
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
async def _send_ws_error(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
status_code: int,
|
||||
message: str,
|
||||
error_type: str | None = None,
|
||||
code: str | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> None:
|
||||
if websocket.application_state != WebSocketState.CONNECTED:
|
||||
return
|
||||
await websocket.send_json(
|
||||
build_codex_ws_error_event(
|
||||
status_code=status_code,
|
||||
message=message,
|
||||
error_type=error_type,
|
||||
code=code,
|
||||
headers=headers,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _close_websocket(websocket: WebSocket, *, code: int) -> None:
|
||||
if websocket.application_state == WebSocketState.CONNECTED:
|
||||
await websocket.close(code=code)
|
||||
|
||||
|
||||
def _string_headers(headers: Any) -> dict[str, str]:
|
||||
if not headers:
|
||||
return {}
|
||||
return {str(key): str(value) for key, value in headers.items()}
|
||||
@@ -10,7 +10,7 @@ OpenAI API 端点
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from fastapi import APIRouter, Depends, Request, WebSocket
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
@@ -102,3 +102,16 @@ async def create_responses(
|
||||
mode=adapter.mode,
|
||||
api_format_hint=adapter.allowed_api_formats[0],
|
||||
)
|
||||
|
||||
|
||||
@router.websocket("/v1/responses")
|
||||
async def create_responses_websocket(websocket: WebSocket) -> None:
|
||||
"""
|
||||
OpenAI Responses WebSocket API (Codex CLI)
|
||||
|
||||
新版 Codex CLI 可通过 WebSocket 连接 /v1/responses。业务处理仍复用
|
||||
HTTP Responses pipeline,仅在这里完成 WS JSON frame 与 SSE 事件之间的转换。
|
||||
"""
|
||||
from src.api.handlers.openai_cli.ws_bridge import handle_codex_responses_websocket
|
||||
|
||||
await handle_codex_responses_websocket(websocket, pipeline=pipeline)
|
||||
|
||||
@@ -0,0 +1,267 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from src.api.handlers.openai_cli.ws_bridge import (
|
||||
CodexWsProtocolError,
|
||||
extract_sse_payloads,
|
||||
iter_sse_payloads,
|
||||
parse_codex_ws_request,
|
||||
)
|
||||
from src.api.public import openai as openai_routes
|
||||
|
||||
|
||||
class FakeDb:
|
||||
def __init__(self) -> None:
|
||||
self.commits = 0
|
||||
self.rollbacks = 0
|
||||
self.closed = False
|
||||
|
||||
def commit(self) -> None:
|
||||
self.commits += 1
|
||||
|
||||
def rollback(self) -> None:
|
||||
self.rollbacks += 1
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
def test_parse_response_create_removes_ws_type() -> None:
|
||||
request = parse_codex_ws_request(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "response.create",
|
||||
"model": "gpt-5",
|
||||
"stream": True,
|
||||
"input": [{"role": "user", "content": "hello"}],
|
||||
"generate": False,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert request.kind == "response.create"
|
||||
assert request.body == {
|
||||
"model": "gpt-5",
|
||||
"stream": True,
|
||||
"input": [{"role": "user", "content": "hello"}],
|
||||
"generate": False,
|
||||
}
|
||||
|
||||
|
||||
def test_parse_response_processed_is_noop_request() -> None:
|
||||
request = parse_codex_ws_request(
|
||||
json.dumps({"type": "response.processed", "response_id": "resp_123"})
|
||||
)
|
||||
|
||||
assert request.kind == "response.processed"
|
||||
assert request.body is None
|
||||
|
||||
|
||||
def test_parse_invalid_frame_raises_protocol_error() -> None:
|
||||
with pytest.raises(CodexWsProtocolError):
|
||||
parse_codex_ws_request("{not-json")
|
||||
|
||||
with pytest.raises(CodexWsProtocolError):
|
||||
parse_codex_ws_request(json.dumps({"type": "response.unknown"}))
|
||||
|
||||
|
||||
def test_extract_sse_payloads_skips_done() -> None:
|
||||
assert extract_sse_payloads('event: response.created\ndata: {"type":"response.created"}') == [
|
||||
'{"type":"response.created"}'
|
||||
]
|
||||
assert extract_sse_payloads("data: [DONE]") == []
|
||||
assert extract_sse_payloads('{"type":"response.completed"}') == [
|
||||
'{"type":"response.completed"}'
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_iter_sse_payloads_handles_split_chunks() -> None:
|
||||
async def chunks() -> AsyncIterator[bytes]:
|
||||
yield b'data: {"type":"response.created"}\n'
|
||||
yield b'\ndata: {"type":"response.completed"}\n\n'
|
||||
yield b"data: [DONE]\n\n"
|
||||
|
||||
payloads = [payload async for payload in iter_sse_payloads(chunks())]
|
||||
|
||||
assert payloads == [
|
||||
'{"type":"response.created"}',
|
||||
'{"type":"response.completed"}',
|
||||
]
|
||||
|
||||
|
||||
def _build_app(monkeypatch: pytest.MonkeyPatch, pipeline: Any, db: FakeDb | None = None) -> FastAPI:
|
||||
app = FastAPI()
|
||||
app.include_router(openai_routes.router)
|
||||
monkeypatch.setattr(openai_routes, "pipeline", pipeline)
|
||||
if db is not None:
|
||||
import src.api.handlers.openai_cli.ws_bridge as ws_bridge
|
||||
|
||||
monkeypatch.setattr(ws_bridge, "create_session", lambda: db)
|
||||
return app
|
||||
|
||||
|
||||
def test_websocket_response_create_streams_sse_as_text_frames(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
seen_bodies: list[dict[str, Any]] = []
|
||||
seen_headers: list[dict[str, str]] = []
|
||||
db = FakeDb()
|
||||
|
||||
class FakePipeline:
|
||||
async def run(self, *, http_request: Any, **_kwargs: Any) -> StreamingResponse:
|
||||
seen_bodies.append(json.loads((await http_request.body()).decode("utf-8")))
|
||||
seen_headers.append(dict(http_request.headers))
|
||||
|
||||
async def body() -> AsyncIterator[bytes]:
|
||||
yield b'data: {"type":"response.created","response":{"id":"resp_1"}}\n\n'
|
||||
yield b'data: {"type":"response.output_text.delta","delta":"hi"}\n\n'
|
||||
yield (
|
||||
b'data: {"type":"response.completed",'
|
||||
b'"response":{"id":"resp_1","status":"completed"}}\n\n'
|
||||
)
|
||||
yield b"data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(body(), media_type="text/event-stream")
|
||||
|
||||
client = TestClient(_build_app(monkeypatch, FakePipeline(), db))
|
||||
|
||||
with client.websocket_connect(
|
||||
"/v1/responses",
|
||||
headers={
|
||||
"Authorization": "Bearer test-key",
|
||||
"OpenAI-Beta": "responses_websockets=2026-02-06",
|
||||
"x-codex-window-id": "window-1",
|
||||
},
|
||||
) as websocket:
|
||||
websocket.send_json(
|
||||
{
|
||||
"type": "response.create",
|
||||
"model": "gpt-5",
|
||||
"stream": True,
|
||||
"input": [{"role": "user", "content": "hello"}],
|
||||
}
|
||||
)
|
||||
|
||||
assert websocket.receive_json() == {
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp_1"},
|
||||
}
|
||||
assert websocket.receive_json() == {
|
||||
"type": "response.output_text.delta",
|
||||
"delta": "hi",
|
||||
}
|
||||
assert websocket.receive_json() == {
|
||||
"type": "response.completed",
|
||||
"response": {"id": "resp_1", "status": "completed"},
|
||||
}
|
||||
|
||||
assert seen_bodies == [
|
||||
{
|
||||
"model": "gpt-5",
|
||||
"stream": True,
|
||||
"input": [{"role": "user", "content": "hello"}],
|
||||
}
|
||||
]
|
||||
assert seen_headers[0]["authorization"] == "Bearer test-key"
|
||||
assert seen_headers[0]["openai-beta"] == "responses_websockets=2026-02-06"
|
||||
assert seen_headers[0]["x-codex-window-id"] == "window-1"
|
||||
assert "sec-websocket-key" not in seen_headers[0]
|
||||
assert db.commits == 1
|
||||
assert db.rollbacks == 0
|
||||
assert db.closed is True
|
||||
|
||||
|
||||
def test_websocket_response_processed_does_not_call_pipeline(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = FakeDb()
|
||||
calls = SimpleNamespace(count=0)
|
||||
|
||||
class FakePipeline:
|
||||
async def run(self, **_kwargs: Any) -> StreamingResponse:
|
||||
calls.count += 1
|
||||
|
||||
async def body() -> AsyncIterator[bytes]:
|
||||
yield b'data: {"type":"response.completed","response":{"id":"resp_2"}}\n\n'
|
||||
|
||||
return StreamingResponse(body(), media_type="text/event-stream")
|
||||
|
||||
client = TestClient(_build_app(monkeypatch, FakePipeline(), db))
|
||||
|
||||
with client.websocket_connect("/v1/responses") as websocket:
|
||||
websocket.send_json({"type": "response.processed", "response_id": "resp_1"})
|
||||
websocket.send_json({"type": "response.create", "model": "gpt-5", "stream": True})
|
||||
assert websocket.receive_json() == {
|
||||
"type": "response.completed",
|
||||
"response": {"id": "resp_2"},
|
||||
}
|
||||
|
||||
assert calls.count == 1
|
||||
|
||||
|
||||
def test_websocket_invalid_json_returns_wrapped_error(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
class FakePipeline:
|
||||
async def run(self, **_kwargs: Any) -> None:
|
||||
raise AssertionError("pipeline should not be called")
|
||||
|
||||
client = TestClient(_build_app(monkeypatch, FakePipeline(), FakeDb()))
|
||||
|
||||
with client.websocket_connect("/v1/responses") as websocket:
|
||||
websocket.send_text("{bad-json")
|
||||
assert websocket.receive_json() == {
|
||||
"type": "error",
|
||||
"status": 400,
|
||||
"error": {
|
||||
"type": "invalid_request_error",
|
||||
"message": "Request frame must be valid JSON",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_websocket_http_exception_returns_wrapped_error(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = FakeDb()
|
||||
|
||||
class FakePipeline:
|
||||
async def run(self, **_kwargs: Any) -> None:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": {
|
||||
"type": "rate_limit_error",
|
||||
"message": "slow down",
|
||||
"code": "rate_limited",
|
||||
}
|
||||
},
|
||||
headers={"Retry-After": "1"},
|
||||
)
|
||||
|
||||
client = TestClient(_build_app(monkeypatch, FakePipeline(), db))
|
||||
|
||||
with client.websocket_connect("/v1/responses") as websocket:
|
||||
websocket.send_json({"type": "response.create", "model": "gpt-5", "stream": True})
|
||||
assert websocket.receive_json() == {
|
||||
"type": "error",
|
||||
"status": 429,
|
||||
"error": {
|
||||
"type": "rate_limit_error",
|
||||
"message": "slow down",
|
||||
"code": "rate_limited",
|
||||
},
|
||||
"headers": {"Retry-After": "1"},
|
||||
}
|
||||
|
||||
assert db.commits == 0
|
||||
assert db.rollbacks == 1
|
||||
assert db.closed is True
|
||||
@@ -0,0 +1,22 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
|
||||
|
||||
def test_frontend_codex_config_advertises_websocket_support() -> None:
|
||||
content = (ROOT / "frontend/src/views/public/home-config.ts").read_text(encoding="utf-8")
|
||||
|
||||
assert 'wire_api = "responses"' in content
|
||||
assert "supports_websockets = true" in content
|
||||
|
||||
|
||||
def test_nginx_api_proxy_forwards_websocket_upgrade() -> None:
|
||||
for dockerfile in ["Dockerfile.app", "Dockerfile.app.local"]:
|
||||
content = (ROOT / dockerfile).read_text(encoding="utf-8")
|
||||
|
||||
assert "map $http_upgrade $connection_upgrade" in content
|
||||
assert "proxy_set_header Upgrade $http_upgrade;" in content
|
||||
assert "proxy_set_header Connection $connection_upgrade;" in content
|
||||
Reference in New Issue
Block a user