mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-07-28 08:56:13 +00:00
276 lines
8.1 KiB
Python
276 lines
8.1 KiB
Python
"""Contract tests for the LangGraph-compatible run request boundary."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from typing import Any
|
|
|
|
import pytest
|
|
from fastapi import FastAPI
|
|
from fastapi.testclient import TestClient
|
|
|
|
from app.gateway.routers import runs
|
|
|
|
SUPPORTED_STREAM_MODES = {
|
|
"values",
|
|
"messages-tuple",
|
|
"updates",
|
|
"debug",
|
|
"tasks",
|
|
"checkpoints",
|
|
"custom",
|
|
}
|
|
|
|
|
|
@pytest.fixture
|
|
def client() -> TestClient:
|
|
app = FastAPI()
|
|
app.include_router(runs.router)
|
|
return TestClient(app, raise_server_exceptions=False)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("field", "value"),
|
|
[
|
|
("webhook", "https://example.com/callback"),
|
|
("stream_resumable", True),
|
|
("on_completion", "complete"),
|
|
("on_completion", "continue"),
|
|
("on_completion", "keep"),
|
|
("on_completion", "delete"),
|
|
("after_seconds", 1.5),
|
|
("if_not_exists", "reject"),
|
|
("feedback_keys", []),
|
|
("multitask_strategy", "enqueue"),
|
|
],
|
|
)
|
|
def test_run_request_rejects_each_unsupported_option_with_exact_422(
|
|
client: TestClient,
|
|
field: str,
|
|
value: Any,
|
|
) -> None:
|
|
response = client.post("/api/runs/stream", json={field: value})
|
|
|
|
assert response.status_code == 422
|
|
assert response.json() == {
|
|
"detail": [
|
|
{
|
|
"type": "unsupported_run_option",
|
|
"loc": ["body", field],
|
|
"msg": f"Run option '{field}' is not supported by DeerFlow",
|
|
"input": value,
|
|
"ctx": {"option": field},
|
|
}
|
|
]
|
|
}
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"stream_mode",
|
|
[
|
|
"messages",
|
|
"events",
|
|
"tools",
|
|
["values", "events"],
|
|
["events", "tools"],
|
|
],
|
|
)
|
|
def test_run_request_rejects_unsupported_stream_modes_with_exact_422(
|
|
client: TestClient,
|
|
stream_mode: str | list[str],
|
|
) -> None:
|
|
response = client.post("/api/runs/stream", json={"stream_mode": stream_mode})
|
|
requested = [stream_mode] if isinstance(stream_mode, str) else stream_mode
|
|
unsupported = list(dict.fromkeys(mode for mode in requested if mode not in SUPPORTED_STREAM_MODES))
|
|
mode_list = ", ".join(unsupported)
|
|
|
|
assert response.status_code == 422
|
|
assert response.json() == {
|
|
"detail": [
|
|
{
|
|
"type": "unsupported_stream_mode",
|
|
"loc": ["body", "stream_mode"],
|
|
"msg": f"Unsupported stream mode(s): {mode_list}",
|
|
"input": stream_mode,
|
|
"ctx": {"modes": mode_list},
|
|
}
|
|
]
|
|
}
|
|
|
|
|
|
def test_run_request_keeps_supported_modes_and_compatibility_defaults() -> None:
|
|
from app.gateway.routers.thread_runs import RunCreateRequest
|
|
|
|
body = RunCreateRequest(
|
|
stream_mode=list(SUPPORTED_STREAM_MODES),
|
|
webhook=None,
|
|
stream_resumable=None,
|
|
on_completion=None,
|
|
after_seconds=None,
|
|
if_not_exists="create",
|
|
feedback_keys=None,
|
|
)
|
|
|
|
assert set(body.stream_mode or []) == SUPPORTED_STREAM_MODES
|
|
assert body.on_completion is None
|
|
assert body.if_not_exists == "create"
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"artifact_delivery",
|
|
[
|
|
{"required": True, "tool": "present_files"},
|
|
{},
|
|
None,
|
|
"present_files",
|
|
[],
|
|
],
|
|
)
|
|
def test_run_request_rejects_removed_artifact_delivery_override(
|
|
client: TestClient,
|
|
artifact_delivery: Any,
|
|
) -> None:
|
|
response = client.post(
|
|
"/api/runs/stream",
|
|
json={"artifact_delivery": artifact_delivery},
|
|
)
|
|
|
|
assert response.status_code == 422
|
|
|
|
|
|
def _sdk_default_payload(method: str) -> dict[str, Any]:
|
|
"""Capture the body a stock ``langgraph_sdk`` client sends for a default run."""
|
|
from langgraph_sdk.client import get_client
|
|
|
|
client = get_client(url="http://gateway.invalid")
|
|
captured: dict[str, Any] = {}
|
|
|
|
def capture(*_args: Any, json: dict[str, Any] | None = None, **_kwargs: Any) -> None:
|
|
captured.update(json or {})
|
|
|
|
async def capture_post(*args: Any, **kwargs: Any) -> dict[str, Any]:
|
|
capture(*args, **kwargs)
|
|
return {}
|
|
|
|
if method == "stream":
|
|
client.runs.http.stream = capture # type: ignore[method-assign]
|
|
# ``stream()`` is a sync factory: the payload is built and handed to the
|
|
# transport before the returned async iterator is consumed.
|
|
client.runs.stream("thread-id", "deerflow", input={"messages": []})
|
|
else:
|
|
client.runs.http.post = capture_post # type: ignore[method-assign]
|
|
asyncio.run(client.runs.create("thread-id", "deerflow", input={"messages": []}))
|
|
return captured
|
|
|
|
|
|
@pytest.mark.parametrize("method", ["stream", "create"])
|
|
def test_gateway_accepts_langgraph_sdk_default_payload(client: TestClient, method: str) -> None:
|
|
"""The SDK's own defaults must never be rejected as unsupported options (#4466).
|
|
|
|
``stream_resumable`` defaults to ``False`` rather than ``None``, so the SDK's
|
|
drop-``None`` payload filter always forwards it; every IM channel run goes
|
|
through this path.
|
|
"""
|
|
payload = _sdk_default_payload(method)
|
|
|
|
assert payload["stream_resumable"] is False, "SDK contract changed; update this test"
|
|
response = client.post("/api/runs/stream", json=payload)
|
|
|
|
assert response.status_code != 422, response.text
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"payload",
|
|
[
|
|
{"multitask_strategy": []},
|
|
{"stream_mode": [{}]},
|
|
],
|
|
)
|
|
def test_malformed_option_types_remain_validation_errors(client: TestClient, payload: dict[str, Any]) -> None:
|
|
response = client.post("/api/runs/stream", json=payload)
|
|
|
|
assert response.status_code == 422
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"field",
|
|
[
|
|
"multitask_strategy",
|
|
"if_not_exists",
|
|
],
|
|
)
|
|
def test_string_run_options_keep_native_type_errors(client: TestClient, field: str) -> None:
|
|
response = client.post("/api/runs/stream", json={field: []})
|
|
|
|
assert response.status_code == 422
|
|
assert response.json()["detail"][0]["type"] == "literal_error"
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("field", "value"),
|
|
[
|
|
("checkpoint_during", False),
|
|
("durability", "sync"),
|
|
],
|
|
)
|
|
def test_run_request_rejects_unknown_sdk_options_with_exact_422(
|
|
client: TestClient,
|
|
field: str,
|
|
value: Any,
|
|
) -> None:
|
|
response = client.post("/api/runs/stream", json={field: value})
|
|
|
|
assert response.status_code == 422
|
|
assert response.json() == {
|
|
"detail": [
|
|
{
|
|
"type": "extra_forbidden",
|
|
"loc": ["body", field],
|
|
"msg": "Extra inputs are not permitted",
|
|
"input": value,
|
|
}
|
|
]
|
|
}
|
|
|
|
|
|
def test_openapi_stream_mode_enum_matches_runtime_support() -> None:
|
|
app = FastAPI()
|
|
app.include_router(runs.router)
|
|
openapi = app.openapi()
|
|
schemas = openapi["components"]["schemas"]
|
|
stream_mode_schema = schemas["RunCreateRequest"]["properties"]["stream_mode"]
|
|
|
|
enums: list[str] = []
|
|
|
|
def collect_enums(value: Any) -> None:
|
|
if isinstance(value, dict):
|
|
enum = value.get("enum")
|
|
if isinstance(enum, list):
|
|
enums.extend(item for item in enum if isinstance(item, str))
|
|
for child in value.values():
|
|
collect_enums(child)
|
|
elif isinstance(value, list):
|
|
for child in value:
|
|
collect_enums(child)
|
|
|
|
collect_enums(stream_mode_schema)
|
|
collect_enums(schemas["RunStreamMode"])
|
|
|
|
assert set(enums) == SUPPORTED_STREAM_MODES
|
|
assert "events" not in enums
|
|
|
|
|
|
def test_openapi_run_option_schema_exposes_only_supported_values() -> None:
|
|
app = FastAPI()
|
|
app.include_router(runs.router)
|
|
schema = app.openapi()["components"]["schemas"]["RunCreateRequest"]
|
|
properties = schema["properties"]
|
|
|
|
assert schema["additionalProperties"] is False
|
|
for field in ("webhook", "after_seconds", "feedback_keys"):
|
|
assert properties[field]["type"] == "null"
|
|
assert properties["stream_resumable"]["anyOf"] == [{"const": False, "type": "boolean"}, {"type": "null"}]
|
|
assert properties["on_completion"]["type"] == "null"
|
|
assert properties["if_not_exists"]["const"] == "create"
|
|
assert properties["multitask_strategy"]["enum"] == ["reject", "rollback", "interrupt"]
|