feat: 导入 MemRelay 初始源码

This commit is contained in:
Nixevol
2026-09-23 21:40:46 +08:00
commit 0ade3bb318
167 changed files with 53511 additions and 0 deletions
+756
View File
@@ -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"