feat: 导入 MemRelay 初始源码
This commit is contained in:
@@ -0,0 +1,756 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from memrelay.model_gateway import (
|
||||
GenerationResult,
|
||||
ModelGatewayError,
|
||||
OpenAICompatibleClient,
|
||||
contains_secret,
|
||||
parse_structured_json,
|
||||
redact_secrets,
|
||||
)
|
||||
from memrelay.models import CurationModelConfig
|
||||
from memrelay.security import SecretBox
|
||||
|
||||
|
||||
def _client(protocol: str = "auto", stream: bool = False) -> OpenAICompatibleClient:
|
||||
return OpenAICompatibleClient(
|
||||
CurationModelConfig(
|
||||
revision=1,
|
||||
name="gateway-test",
|
||||
base_url="https://model.example.test/v1",
|
||||
text_model="test-model",
|
||||
protocol=protocol,
|
||||
stream=stream,
|
||||
active=True,
|
||||
),
|
||||
SecretBox(b"0" * 32),
|
||||
)
|
||||
|
||||
|
||||
def _result(text: str, protocol: str) -> GenerationResult:
|
||||
return GenerationResult(
|
||||
text=text,
|
||||
protocol=protocol,
|
||||
response_model="test-model",
|
||||
request_id="request-1",
|
||||
input_tokens=3,
|
||||
output_tokens=2,
|
||||
latency_ms=10,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"secret",
|
||||
[
|
||||
"password: correct-horse-battery-staple",
|
||||
"Authorization: Bearer eyJhbGciOiJIUzI1NiJ9.payload.signature",
|
||||
"api_key=sk-1234567890abcdefghijklmnop",
|
||||
"github_pat_1234567890abcdefghijklmnop",
|
||||
"glpat-1234567890abcdefghijklmnop",
|
||||
"xoxb-1234567890abcdefghijklmnop",
|
||||
"AKIA1234567890ABCDEF",
|
||||
"otpauth://totp/MemRelay:user?secret=JBSWY3DPEHPK3PXP",
|
||||
"postgresql://admin:database-password@db.example.test/memrelay",
|
||||
"-----BEGIN PRIVATE KEY-----\nprivate-material\n-----END PRIVATE KEY-----",
|
||||
],
|
||||
)
|
||||
def test_secret_patterns_are_detected_and_redacted_without_echoing(secret: str) -> None:
|
||||
assert contains_secret(secret)
|
||||
redacted, count = redact_secrets(f"before {secret} after")
|
||||
assert count >= 1
|
||||
assert secret not in redacted
|
||||
assert "[REDACTED]" in redacted
|
||||
|
||||
|
||||
def test_probe_marks_transient_protocol_outage_as_waiting_dependency() -> None:
|
||||
client = _client("responses")
|
||||
client.list_models = AsyncMock(side_effect=ModelGatewayError("DOWN", "down", transient=True))
|
||||
client.generate = AsyncMock(
|
||||
side_effect=ModelGatewayError(
|
||||
"MODEL_UNAVAILABLE",
|
||||
"temporarily unavailable",
|
||||
transient=True,
|
||||
retry_after=45,
|
||||
)
|
||||
)
|
||||
|
||||
result = asyncio.run(client.probe(test_both=False))
|
||||
|
||||
assert result["available"] is False
|
||||
assert result["dependency_status"] == "waiting_dependency"
|
||||
assert result["error_code"] == "MODEL_UNAVAILABLE"
|
||||
assert result["retry_after"] == 45
|
||||
|
||||
|
||||
def test_insufficient_balance_is_reported_as_a_recoverable_quota_outage() -> None:
|
||||
client = _client()
|
||||
response = httpx.Response(
|
||||
403,
|
||||
json={"error": {"message": "Insufficient account balance"}},
|
||||
request=httpx.Request("POST", "https://model.example.test/v1/responses"),
|
||||
)
|
||||
|
||||
error = client._error(response)
|
||||
|
||||
assert error.code == "MODEL_QUOTA_EXHAUSTED"
|
||||
assert error.transient is True
|
||||
assert error.http_status == 403
|
||||
assert "额度不足" in error.message
|
||||
|
||||
|
||||
def test_auto_probe_falls_back_to_chat_completions_after_responses_is_unsupported() -> None:
|
||||
client = _client("auto")
|
||||
client.list_models = AsyncMock(return_value=[{"id": "test-model"}])
|
||||
|
||||
async def generate(
|
||||
prompt: str,
|
||||
*,
|
||||
protocol: str | None = None,
|
||||
stream: bool | None = None,
|
||||
max_output_tokens: int | None = None,
|
||||
**_kwargs: object,
|
||||
) -> GenerationResult:
|
||||
del max_output_tokens
|
||||
if protocol == "responses":
|
||||
raise ModelGatewayError(
|
||||
"MODEL_PROTOCOL_UNSUPPORTED",
|
||||
"responses unsupported",
|
||||
transient=False,
|
||||
http_status=404,
|
||||
)
|
||||
if '{"ok":true}' in prompt:
|
||||
return _result('{"ok":true}', "chat_completions")
|
||||
return _result("OK" if not stream else "OK streamed", "chat_completions")
|
||||
|
||||
client.generate = generate # type: ignore[method-assign]
|
||||
result = asyncio.run(client.probe(test_both=True))
|
||||
|
||||
assert result["responses"]["status"] == "unsupported"
|
||||
assert result["chat_completions"]["status"] == "supported"
|
||||
assert result["structured_output"]["status"] == "supported"
|
||||
assert result["stream"]["status"] == "supported"
|
||||
assert result["effective_protocol"] == "chat_completions"
|
||||
assert result["dependency_status"] == "available"
|
||||
|
||||
|
||||
def test_stream_assembly_and_structured_json_parsing() -> None:
|
||||
client = _client()
|
||||
responses = client._assemble_stream(
|
||||
"responses",
|
||||
[
|
||||
{"type": "response.output_text.delta", "delta": "hel"},
|
||||
{"type": "response.output_text.delta", "delta": "lo"},
|
||||
],
|
||||
)
|
||||
chat = client._assemble_stream(
|
||||
"chat_completions",
|
||||
[
|
||||
{"choices": [{"delta": {"content": "hel"}}]},
|
||||
{"choices": [{"delta": {"content": "lo"}}]},
|
||||
],
|
||||
)
|
||||
assert responses["output_text"] == "hello"
|
||||
assert chat["choices"][0]["message"]["content"] == "hello"
|
||||
assert parse_structured_json('```json\n{"ok": true}\n```') == {"ok": True}
|
||||
|
||||
with pytest.raises(ModelGatewayError, match="完成标记") as incomplete:
|
||||
client._assert_stream_complete("chat_completions", False, False)
|
||||
assert incomplete.value.code == "MODEL_STREAM_INCOMPLETE"
|
||||
client._assert_stream_complete("responses", False, True)
|
||||
|
||||
|
||||
def test_payloads_cover_both_protocols_native_schema_and_vision() -> None:
|
||||
client = _client()
|
||||
schema = {"type": "object", "properties": {"ok": {"type": "boolean"}}}
|
||||
responses = client._request_payload(
|
||||
"responses",
|
||||
"inspect",
|
||||
32,
|
||||
model="vision-model",
|
||||
image_data_url="data:image/png;base64,AAAA",
|
||||
json_schema=schema,
|
||||
)
|
||||
assert responses["model"] == "vision-model"
|
||||
assert responses["input"][0]["content"][1]["type"] == "input_image"
|
||||
assert responses["text"]["format"]["schema"] == schema
|
||||
|
||||
chat = client._request_payload(
|
||||
"chat_completions",
|
||||
"inspect",
|
||||
32,
|
||||
model="vision-model",
|
||||
image_data_url="data:image/png;base64,AAAA",
|
||||
json_schema=schema,
|
||||
)
|
||||
assert chat["messages"][0]["content"][1]["type"] == "image_url"
|
||||
assert chat["response_format"]["json_schema"]["schema"] == schema
|
||||
|
||||
|
||||
def test_generic_bad_request_is_not_reported_as_protocol_unsupported() -> None:
|
||||
client = _client()
|
||||
request = httpx.Request("POST", "https://model.example.test/v1/responses")
|
||||
error = client._error(
|
||||
httpx.Response(400, request=request, json={"error": {"message": "invalid model"}})
|
||||
)
|
||||
assert error.code == "MODEL_REQUEST_INVALID"
|
||||
assert client._capability_error(error)["status"] == "error"
|
||||
|
||||
|
||||
def test_native_schema_failure_falls_back_to_strict_json_prompt() -> None:
|
||||
client = _client("responses")
|
||||
client.list_models = AsyncMock(return_value=[])
|
||||
|
||||
async def generate(
|
||||
prompt: str,
|
||||
*,
|
||||
protocol: str | None = None,
|
||||
stream: bool | None = None,
|
||||
json_schema: dict | None = None,
|
||||
**_kwargs: object,
|
||||
) -> GenerationResult:
|
||||
if json_schema:
|
||||
raise ModelGatewayError(
|
||||
"MODEL_REQUEST_INVALID", "schema unsupported", transient=False, http_status=400
|
||||
)
|
||||
if '{"ok":true}' in prompt:
|
||||
return _result('{"ok":true}', protocol or "responses")
|
||||
return _result("OK streamed" if stream else "OK", protocol or "responses")
|
||||
|
||||
client.generate = generate # type: ignore[method-assign]
|
||||
result = asyncio.run(client.probe(test_both=False))
|
||||
assert result["native_structured_output"]["status"] == "error"
|
||||
assert result["structured_output"]["status"] == "supported"
|
||||
assert result["effective_protocol"] == "responses"
|
||||
|
||||
|
||||
def _protocol_transport(supported: set[str]) -> httpx.MockTransport:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
if request.url.path.endswith("/models"):
|
||||
return httpx.Response(200, json={"data": [{"id": "test-model"}]})
|
||||
protocol = (
|
||||
"chat_completions" if request.url.path.endswith("/chat/completions") else "responses"
|
||||
)
|
||||
if protocol not in supported:
|
||||
return httpx.Response(404, json={"error": {"message": "not found"}})
|
||||
payload = json.loads(request.content)
|
||||
structured = "text" in payload or "response_format" in payload
|
||||
text = '{"ok":true}' if structured else "OK"
|
||||
if payload.get("stream"):
|
||||
if protocol == "responses":
|
||||
event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "response-stream",
|
||||
"model": "test-model",
|
||||
"output_text": text,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 1},
|
||||
},
|
||||
}
|
||||
else:
|
||||
event = {
|
||||
"id": "chat-stream",
|
||||
"model": "test-model",
|
||||
"choices": [{"delta": {"content": text}}],
|
||||
}
|
||||
body = f"data: {json.dumps(event)}\n\ndata: [DONE]\n\n"
|
||||
return httpx.Response(200, text=body, headers={"content-type": "text/event-stream"})
|
||||
if protocol == "responses":
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "response-json",
|
||||
"model": "test-model",
|
||||
"output_text": text,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 1},
|
||||
},
|
||||
)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "chat-json",
|
||||
"model": "test-model",
|
||||
"choices": [{"message": {"content": text}}],
|
||||
"usage": {"prompt_tokens": 2, "completion_tokens": 1},
|
||||
},
|
||||
)
|
||||
|
||||
return httpx.MockTransport(handler)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("supported", "expected"),
|
||||
[
|
||||
({"responses"}, "responses"),
|
||||
({"chat_completions"}, "chat_completions"),
|
||||
({"responses", "chat_completions"}, "responses"),
|
||||
],
|
||||
)
|
||||
def test_probe_selects_supported_openai_protocols(
|
||||
supported: set[str],
|
||||
expected: str,
|
||||
) -> None:
|
||||
client = OpenAICompatibleClient(
|
||||
_client().config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=_protocol_transport(supported),
|
||||
)
|
||||
capabilities = asyncio.run(client.probe(test_both=True))
|
||||
assert capabilities["effective_protocol"] == expected
|
||||
assert capabilities["available"] is True
|
||||
assert capabilities["structured_output"]["status"] == "supported"
|
||||
assert capabilities["stream"]["status"] == "supported"
|
||||
for protocol in {"responses", "chat_completions"}:
|
||||
expected_status = "supported" if protocol in supported else "unsupported"
|
||||
assert capabilities[protocol]["status"] == expected_status
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("status", "code", "retry_after"),
|
||||
[(429, "MODEL_RATE_LIMITED", 17), (503, "MODEL_SERVER_ERROR", None)],
|
||||
)
|
||||
def test_transient_http_errors_keep_the_selected_protocol(
|
||||
status: int,
|
||||
code: str,
|
||||
retry_after: int | None,
|
||||
) -> None:
|
||||
calls: list[str] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
calls.append(request.url.path)
|
||||
headers = {"Retry-After": str(retry_after)} if retry_after else None
|
||||
return httpx.Response(status, json={"error": {"message": "temporary"}}, headers=headers)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
with pytest.raises(ModelGatewayError) as raised:
|
||||
asyncio.run(client.generate("test", protocol="responses", stream=False))
|
||||
assert raised.value.code == code
|
||||
assert raised.value.transient is True
|
||||
assert raised.value.retry_after == retry_after
|
||||
assert calls == ["/v1/responses"] * 3
|
||||
|
||||
|
||||
def test_transient_model_request_retries_until_the_provider_recovers() -> None:
|
||||
calls = 0
|
||||
|
||||
def handler(_request: httpx.Request) -> httpx.Response:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
if calls < 3:
|
||||
return httpx.Response(503, json={"error": {"message": "temporary"}})
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={"id": "response-ok", "model": "test-model", "output_text": "OK"},
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
result = asyncio.run(client.generate("test", protocol="responses", stream=False))
|
||||
|
||||
assert result.text == "OK"
|
||||
assert result.attempts == 3
|
||||
assert calls == 3
|
||||
|
||||
|
||||
def test_remote_protocol_disconnect_is_retried_as_a_transient_outage() -> None:
|
||||
calls = 0
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
if calls < 3:
|
||||
raise httpx.RemoteProtocolError(
|
||||
"Server disconnected without sending a response.",
|
||||
request=request,
|
||||
)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={"id": "response-ok", "model": "test-model", "output_text": "OK"},
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
result = asyncio.run(client.generate("test", protocol="responses", stream=False))
|
||||
|
||||
assert result.text == "OK"
|
||||
assert calls == 3
|
||||
|
||||
|
||||
def test_repeated_remote_protocol_disconnect_is_reported_as_model_unavailable() -> None:
|
||||
calls = 0
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
raise httpx.RemoteProtocolError(
|
||||
"Server disconnected without sending a response.",
|
||||
request=request,
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
with pytest.raises(ModelGatewayError) as raised:
|
||||
asyncio.run(client.generate("test", protocol="responses", stream=False))
|
||||
|
||||
assert raised.value.code == "MODEL_UNAVAILABLE"
|
||||
assert raised.value.transient is True
|
||||
assert calls == 3
|
||||
|
||||
|
||||
def test_stream_transport_failure_does_not_fallback_to_non_stream() -> None:
|
||||
streamed: list[bool] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
streamed.append(bool(json.loads(request.content).get("stream")))
|
||||
raise httpx.RemoteProtocolError(
|
||||
"Server disconnected during a streamed response.",
|
||||
request=request,
|
||||
)
|
||||
|
||||
config = _client("responses", stream=True).config
|
||||
client = OpenAICompatibleClient(
|
||||
config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
with pytest.raises(ModelGatewayError) as raised:
|
||||
asyncio.run(client.generate("test"))
|
||||
|
||||
assert raised.value.code == "MODEL_UNAVAILABLE"
|
||||
assert streamed == [True, True, True]
|
||||
|
||||
|
||||
def test_stream_read_timeout_matches_the_configured_request_timeout() -> None:
|
||||
client = _client("responses", stream=True)
|
||||
client.config.timeout_seconds = 300
|
||||
|
||||
timeout = client._timeout(stream=True)
|
||||
|
||||
assert timeout.connect == 30
|
||||
assert timeout.read == 300
|
||||
assert timeout.write == 120
|
||||
assert timeout.pool == 30
|
||||
|
||||
|
||||
def test_model_discovery_maps_remote_protocol_disconnect_to_model_unavailable() -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
raise httpx.RemoteProtocolError(
|
||||
"Server disconnected without sending a response.",
|
||||
request=request,
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
with pytest.raises(ModelGatewayError) as raised:
|
||||
asyncio.run(client.list_models())
|
||||
|
||||
assert raised.value.code == "MODEL_UNAVAILABLE"
|
||||
assert raised.value.transient is True
|
||||
|
||||
|
||||
def test_model_discovery_rejects_invalid_json_response() -> None:
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(
|
||||
lambda _request: httpx.Response(200, text="not-json")
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(ModelGatewayError) as raised:
|
||||
asyncio.run(client.list_models())
|
||||
|
||||
assert raised.value.code == "MODEL_RESPONSE_INVALID"
|
||||
assert raised.value.transient is False
|
||||
|
||||
|
||||
def test_auto_protocol_falls_back_after_repeated_responses_failures() -> None:
|
||||
calls: list[str] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
calls.append(request.url.path)
|
||||
if request.url.path.endswith("/responses"):
|
||||
return httpx.Response(503, json={"error": {"message": "temporary"}})
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "chat-ok",
|
||||
"model": "test-model",
|
||||
"choices": [{"message": {"content": "OK"}}],
|
||||
},
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("auto").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
result = asyncio.run(client.generate("test", stream=False))
|
||||
|
||||
assert result.protocol == "chat_completions"
|
||||
assert calls == ["/v1/responses"] * 3 + ["/v1/chat/completions"]
|
||||
|
||||
|
||||
def test_rejected_native_schema_retries_with_plain_json_request() -> None:
|
||||
payloads: list[dict] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
payload = json.loads(request.content)
|
||||
payloads.append(payload)
|
||||
if "text" in payload:
|
||||
return httpx.Response(400, json={"error": {"message": "schema unsupported"}})
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "response-ok",
|
||||
"model": "test-model",
|
||||
"output_text": '{"ok":true}',
|
||||
},
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
client.generate(
|
||||
"return json",
|
||||
protocol="responses",
|
||||
stream=False,
|
||||
json_schema={"type": "object"},
|
||||
)
|
||||
)
|
||||
|
||||
assert parse_structured_json(result.text) == {"ok": True}
|
||||
assert len(payloads) == 2
|
||||
assert "text" in payloads[0]
|
||||
assert "text" not in payloads[1]
|
||||
|
||||
|
||||
def test_incomplete_sse_stream_is_rejected_without_partial_result() -> None:
|
||||
def handler(_request: httpx.Request) -> httpx.Response:
|
||||
body = 'data: {"type":"response.output_text.delta","delta":"partial"}\n\n'
|
||||
return httpx.Response(200, text=body, headers={"content-type": "text/event-stream"})
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses", stream=True).config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
with pytest.raises(ModelGatewayError) as raised:
|
||||
asyncio.run(client.generate("test", protocol="responses", stream=True))
|
||||
assert raised.value.code == "MODEL_STREAM_INCOMPLETE"
|
||||
assert raised.value.transient is True
|
||||
|
||||
|
||||
def test_real_responses_endpoint_when_runtime_credentials_are_supplied() -> None:
|
||||
base_url = os.getenv("MEMRELAY_TEST_MODEL_BASE_URL")
|
||||
token = os.getenv("MEMRELAY_TEST_MODEL_TOKEN")
|
||||
model = os.getenv("MEMRELAY_TEST_MODEL_NAME", "gpt-5.6-sol")
|
||||
if not base_url or not token:
|
||||
pytest.skip("真实模型连接只在运行时注入凭证后执行")
|
||||
secret_box = SecretBox(b"1" * 32)
|
||||
config = CurationModelConfig(
|
||||
revision=1,
|
||||
name="runtime-responses-test",
|
||||
base_url=base_url,
|
||||
encrypted_token=secret_box.encrypt(token, "curation-model-token"),
|
||||
text_model=model,
|
||||
protocol="responses",
|
||||
stream=True,
|
||||
timeout_seconds=180,
|
||||
active=True,
|
||||
)
|
||||
client = OpenAICompatibleClient(config, secret_box)
|
||||
|
||||
result = asyncio.run(
|
||||
client.generate(
|
||||
'Return only this JSON object: {"ok":true}',
|
||||
protocol="responses",
|
||||
stream=False,
|
||||
max_output_tokens=64,
|
||||
)
|
||||
)
|
||||
assert parse_structured_json(result.text)["ok"] is True
|
||||
assert result.protocol == "responses"
|
||||
assert result.request_id
|
||||
assert result.response_model
|
||||
|
||||
capabilities = asyncio.run(client.probe(test_both=False))
|
||||
assert capabilities["responses"]["status"] == "supported"
|
||||
assert capabilities["non_stream"]["status"] == "supported"
|
||||
assert capabilities["structured_output"]["status"] == "supported"
|
||||
assert capabilities["stream"]["status"] == "supported"
|
||||
assert capabilities["effective_protocol"] == "responses"
|
||||
assert capabilities["available"] is True
|
||||
|
||||
|
||||
def _truncated_chat_transport() -> httpx.MockTransport:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "chat-truncated",
|
||||
"model": "test-model",
|
||||
"choices": [
|
||||
{"message": {"content": '{"documents": [{"docu'}, "finish_reason": "length"}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 4096},
|
||||
},
|
||||
)
|
||||
|
||||
return httpx.MockTransport(handler)
|
||||
|
||||
|
||||
def test_chat_finish_reason_length_raises_dedicated_truncation_error() -> None:
|
||||
# 生产回归:截断的 JSON 之前被误报为 MODEL_OUTPUT_INVALID,掩盖了真实原因。
|
||||
client = OpenAICompatibleClient(
|
||||
_client("chat_completions").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=_truncated_chat_transport(),
|
||||
)
|
||||
with pytest.raises(ModelGatewayError) as excinfo:
|
||||
asyncio.run(client.generate("prompt", protocol="chat_completions", stream=False))
|
||||
assert excinfo.value.code == "MODEL_OUTPUT_TRUNCATED"
|
||||
assert excinfo.value.transient is False
|
||||
|
||||
|
||||
def test_responses_incomplete_status_raises_dedicated_truncation_error() -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "response-truncated",
|
||||
"model": "test-model",
|
||||
"status": "incomplete",
|
||||
"incomplete_details": {"reason": "max_output_tokens"},
|
||||
"output_text": '{"documents": [{"docu',
|
||||
"usage": {"input_tokens": 10, "output_tokens": 4096},
|
||||
},
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
with pytest.raises(ModelGatewayError) as excinfo:
|
||||
asyncio.run(client.generate("prompt", protocol="responses", stream=False))
|
||||
assert excinfo.value.code == "MODEL_OUTPUT_TRUNCATED"
|
||||
assert excinfo.value.transient is False
|
||||
|
||||
|
||||
def test_responses_stream_incomplete_terminal_event_reports_truncation() -> None:
|
||||
# response.incomplete 是正常终止事件:不能按流中断(transient)无限重试,
|
||||
# 而应报告确定性的输出截断。
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
event = {
|
||||
"type": "response.incomplete",
|
||||
"response": {
|
||||
"id": "response-stream-truncated",
|
||||
"model": "test-model",
|
||||
"status": "incomplete",
|
||||
"incomplete_details": {"reason": "max_output_tokens"},
|
||||
"output_text": '{"documents": [{"docu',
|
||||
},
|
||||
}
|
||||
body = f"data: {json.dumps(event)}\n\ndata: [DONE]\n\n"
|
||||
return httpx.Response(200, text=body, headers={"content-type": "text/event-stream"})
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses", stream=True).config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
with pytest.raises(ModelGatewayError) as excinfo:
|
||||
asyncio.run(client.generate("prompt", protocol="responses", stream=True))
|
||||
assert excinfo.value.code == "MODEL_OUTPUT_TRUNCATED"
|
||||
|
||||
|
||||
def test_probe_ignores_truncation_flags_when_usable_text_is_present() -> None:
|
||||
# 探测使用较小的输出预算,推理模型可能带 length 标记仍返回可用文本;
|
||||
# 探测不能因此把依赖判成配置错误。
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
if request.url.path.endswith("/models"):
|
||||
return httpx.Response(200, json={"data": [{"id": "test-model"}]})
|
||||
payload = json.loads(request.content)
|
||||
structured = "text" in payload or "response_format" in payload
|
||||
text = '{"ok":true}' if structured else "OK"
|
||||
if payload.get("stream"):
|
||||
event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "probe-stream",
|
||||
"model": "test-model",
|
||||
"status": "incomplete",
|
||||
"incomplete_details": {"reason": "max_output_tokens"},
|
||||
"output_text": text,
|
||||
},
|
||||
}
|
||||
body = f"data: {json.dumps(event)}\n\ndata: [DONE]\n\n"
|
||||
return httpx.Response(200, text=body, headers={"content-type": "text/event-stream"})
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "probe-json",
|
||||
"model": "test-model",
|
||||
"status": "incomplete",
|
||||
"incomplete_details": {"reason": "max_output_tokens"},
|
||||
"output_text": text,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 16},
|
||||
},
|
||||
)
|
||||
|
||||
client = OpenAICompatibleClient(
|
||||
_client("responses").config,
|
||||
SecretBox(b"0" * 32),
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
capabilities = asyncio.run(client.probe(test_both=False))
|
||||
assert capabilities["responses"]["status"] == "supported"
|
||||
assert capabilities["structured_output"]["status"] == "supported"
|
||||
assert capabilities["available"] is True
|
||||
|
||||
|
||||
def test_chat_stream_assembly_preserves_finish_reason() -> None:
|
||||
client = _client()
|
||||
assembled = client._assemble_stream(
|
||||
"chat_completions",
|
||||
[
|
||||
{"choices": [{"delta": {"content": "hel"}}]},
|
||||
{"choices": [{"delta": {"content": "lo"}, "finish_reason": "length"}]},
|
||||
],
|
||||
)
|
||||
assert assembled["choices"][0]["message"]["content"] == "hello"
|
||||
assert assembled["choices"][0]["finish_reason"] == "length"
|
||||
Reference in New Issue
Block a user