变量函数重命名
This commit is contained in:
@@ -13,7 +13,7 @@ from models import (
|
|||||||
SendMessageRequest,
|
SendMessageRequest,
|
||||||
SendMessageResponse,
|
SendMessageResponse,
|
||||||
)
|
)
|
||||||
from service import ChatService, DeepSeekUpstreamError
|
from service import ChatProviderError, ChatService
|
||||||
from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError
|
from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError
|
||||||
|
|
||||||
|
|
||||||
@@ -23,7 +23,7 @@ async def encode_sse(
|
|||||||
try:
|
try:
|
||||||
async for event in events:
|
async for event in events:
|
||||||
yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
|
yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
|
||||||
except (DeepSeekUpstreamError, SessionStorageError) as exc:
|
except (ChatProviderError, SessionStorageError) as exc:
|
||||||
error = {"type": "error", "detail": str(exc)}
|
error = {"type": "error", "detail": str(exc)}
|
||||||
yield f"data: {json.dumps(error, ensure_ascii=False)}\n\n"
|
yield f"data: {json.dumps(error, ensure_ascii=False)}\n\n"
|
||||||
yield "data: [DONE]\n\n"
|
yield "data: [DONE]\n\n"
|
||||||
@@ -109,7 +109,7 @@ def create_app(
|
|||||||
return JSONResponse(content=response.model_dump(mode="json"))
|
return JSONResponse(content=response.model_dump(mode="json"))
|
||||||
except SessionNotFoundError as exc:
|
except SessionNotFoundError as exc:
|
||||||
raise HTTPException(status_code=404, detail="session not found") from exc
|
raise HTTPException(status_code=404, detail="session not found") from exc
|
||||||
except DeepSeekUpstreamError as exc:
|
except ChatProviderError as exc:
|
||||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||||
except SessionStorageError as exc:
|
except SessionStorageError as exc:
|
||||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||||
|
|||||||
+5
-5
@@ -7,7 +7,7 @@ from models import Message, SendMessageResponse, Session, TokenUsage
|
|||||||
from storage import JsonSessionStorage
|
from storage import JsonSessionStorage
|
||||||
|
|
||||||
|
|
||||||
class DeepSeekUpstreamError(Exception):
|
class ChatProviderError(Exception):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
@@ -41,10 +41,10 @@ class ChatService:
|
|||||||
)
|
)
|
||||||
assistant_content = completion.choices[0].message.content
|
assistant_content = completion.choices[0].message.content
|
||||||
except (APIError, IndexError, AttributeError) as exc:
|
except (APIError, IndexError, AttributeError) as exc:
|
||||||
raise DeepSeekUpstreamError("DeepSeek request failed") from exc
|
raise ChatProviderError("Upstream chat request failed") from exc
|
||||||
|
|
||||||
if assistant_content is None:
|
if assistant_content is None:
|
||||||
raise DeepSeekUpstreamError("DeepSeek returned an empty response")
|
raise ChatProviderError("Upstream chat provider returned an empty response")
|
||||||
|
|
||||||
assistant_message = await self._save_exchange(
|
assistant_message = await self._save_exchange(
|
||||||
session, content, user_created_at, assistant_content
|
session, content, user_created_at, assistant_content
|
||||||
@@ -94,11 +94,11 @@ class ChatService:
|
|||||||
finally:
|
finally:
|
||||||
await stream.close()
|
await stream.close()
|
||||||
except (APIError, AttributeError, TypeError) as exc:
|
except (APIError, AttributeError, TypeError) as exc:
|
||||||
raise DeepSeekUpstreamError("DeepSeek stream failed") from exc
|
raise ChatProviderError("Upstream chat stream failed") from exc
|
||||||
|
|
||||||
assistant_content = "".join(parts)
|
assistant_content = "".join(parts)
|
||||||
if not assistant_content:
|
if not assistant_content:
|
||||||
raise DeepSeekUpstreamError("DeepSeek returned an empty response")
|
raise ChatProviderError("Upstream chat provider returned an empty response")
|
||||||
|
|
||||||
assistant_message = await self._save_exchange(
|
assistant_message = await self._save_exchange(
|
||||||
session, content, user_created_at, assistant_content
|
session, content, user_created_at, assistant_content
|
||||||
|
|||||||
+33
-33
@@ -80,8 +80,8 @@ class FakeClient:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def client_and_deepseek(tmp_path: Path) -> tuple[TestClient, FakeClient]:
|
def client_and_provider(tmp_path: Path) -> tuple[TestClient, FakeClient]:
|
||||||
deepseek = FakeClient()
|
provider = FakeClient()
|
||||||
settings = Settings(
|
settings = Settings(
|
||||||
api_key=None,
|
api_key=None,
|
||||||
base_url="https://api.deepseek.com",
|
base_url="https://api.deepseek.com",
|
||||||
@@ -89,9 +89,9 @@ def client_and_deepseek(tmp_path: Path) -> tuple[TestClient, FakeClient]:
|
|||||||
default_system_prompt="default prompt",
|
default_system_prompt="default prompt",
|
||||||
data_dir=tmp_path,
|
data_dir=tmp_path,
|
||||||
)
|
)
|
||||||
app = create_app(settings=settings, client=deepseek) # type: ignore[arg-type]
|
app = create_app(settings=settings, client=provider) # type: ignore[arg-type]
|
||||||
with TestClient(app) as client:
|
with TestClient(app) as client:
|
||||||
yield client, deepseek
|
yield client, provider
|
||||||
|
|
||||||
|
|
||||||
def create_session(client: TestClient, **body: str) -> str:
|
def create_session(client: TestClient, **body: str) -> str:
|
||||||
@@ -101,9 +101,9 @@ def create_session(client: TestClient, **body: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def test_create_session_uses_default_prompt(
|
def test_create_session_uses_default_prompt(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
|
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
client, _ = client_and_deepseek
|
client, _ = client_and_provider
|
||||||
session_id = create_session(client)
|
session_id = create_session(client)
|
||||||
|
|
||||||
data = json.loads((tmp_path / f"{session_id}.json").read_text())
|
data = json.loads((tmp_path / f"{session_id}.json").read_text())
|
||||||
@@ -112,9 +112,9 @@ def test_create_session_uses_default_prompt(
|
|||||||
|
|
||||||
|
|
||||||
def test_create_session_without_body(
|
def test_create_session_without_body(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient]
|
client_and_provider: tuple[TestClient, FakeClient]
|
||||||
) -> None:
|
) -> None:
|
||||||
client, _ = client_and_deepseek
|
client, _ = client_and_provider
|
||||||
|
|
||||||
response = client.post("/sessions")
|
response = client.post("/sessions")
|
||||||
|
|
||||||
@@ -123,9 +123,9 @@ def test_create_session_without_body(
|
|||||||
|
|
||||||
|
|
||||||
def test_create_session_accepts_custom_prompt(
|
def test_create_session_accepts_custom_prompt(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient]
|
client_and_provider: tuple[TestClient, FakeClient]
|
||||||
) -> None:
|
) -> None:
|
||||||
client, _ = client_and_deepseek
|
client, _ = client_and_provider
|
||||||
response = client.post("/sessions", json={"system_prompt": "回答中文"})
|
response = client.post("/sessions", json={"system_prompt": "回答中文"})
|
||||||
|
|
||||||
assert response.status_code == 201
|
assert response.status_code == 201
|
||||||
@@ -133,9 +133,9 @@ def test_create_session_accepts_custom_prompt(
|
|||||||
|
|
||||||
|
|
||||||
def test_multi_turn_request_contains_full_history(
|
def test_multi_turn_request_contains_full_history(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
|
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
client, deepseek = client_and_deepseek
|
client, provider = client_and_provider
|
||||||
session_id = create_session(client, system_prompt="system")
|
session_id = create_session(client, system_prompt="system")
|
||||||
|
|
||||||
first = client.post(
|
first = client.post(
|
||||||
@@ -146,7 +146,7 @@ def test_multi_turn_request_contains_full_history(
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert first.status_code == second.status_code == 200
|
assert first.status_code == second.status_code == 200
|
||||||
assert deepseek.completions.calls[1]["messages"] == [
|
assert provider.completions.calls[1]["messages"] == [
|
||||||
{"role": "system", "content": "system"},
|
{"role": "system", "content": "system"},
|
||||||
{"role": "user", "content": "first"},
|
{"role": "user", "content": "first"},
|
||||||
{"role": "assistant", "content": "reply-1"},
|
{"role": "assistant", "content": "reply-1"},
|
||||||
@@ -169,9 +169,9 @@ def test_multi_turn_request_contains_full_history(
|
|||||||
|
|
||||||
|
|
||||||
def test_invalid_session_and_blank_message(
|
def test_invalid_session_and_blank_message(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient]
|
client_and_provider: tuple[TestClient, FakeClient]
|
||||||
) -> None:
|
) -> None:
|
||||||
client, _ = client_and_deepseek
|
client, _ = client_and_provider
|
||||||
|
|
||||||
missing = client.post(
|
missing = client.post(
|
||||||
"/sessions/00000000-0000-0000-0000-000000000000/messages",
|
"/sessions/00000000-0000-0000-0000-000000000000/messages",
|
||||||
@@ -187,12 +187,12 @@ def test_invalid_session_and_blank_message(
|
|||||||
|
|
||||||
|
|
||||||
def test_upstream_failure_does_not_change_history(
|
def test_upstream_failure_does_not_change_history(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
|
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
client, deepseek = client_and_deepseek
|
client, provider = client_and_provider
|
||||||
session_id = create_session(client)
|
session_id = create_session(client)
|
||||||
before = (tmp_path / f"{session_id}.json").read_text()
|
before = (tmp_path / f"{session_id}.json").read_text()
|
||||||
deepseek.completions.error = APIConnectionError(request=object()) # type: ignore[arg-type]
|
provider.completions.error = APIConnectionError(request=object()) # type: ignore[arg-type]
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
f"/sessions/{session_id}/messages", json={"content": "hello"}
|
f"/sessions/{session_id}/messages", json={"content": "hello"}
|
||||||
@@ -203,9 +203,9 @@ def test_upstream_failure_does_not_change_history(
|
|||||||
|
|
||||||
|
|
||||||
def test_corrupt_session_returns_500(
|
def test_corrupt_session_returns_500(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
|
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
client, _ = client_and_deepseek
|
client, _ = client_and_provider
|
||||||
session_id = create_session(client)
|
session_id = create_session(client)
|
||||||
(tmp_path / f"{session_id}.json").write_text("not-json")
|
(tmp_path / f"{session_id}.json").write_text("not-json")
|
||||||
|
|
||||||
@@ -217,9 +217,9 @@ def test_corrupt_session_returns_500(
|
|||||||
|
|
||||||
|
|
||||||
def test_streaming_response_is_sse_and_persists_complete_message(
|
def test_streaming_response_is_sse_and_persists_complete_message(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
|
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
client, deepseek = client_and_deepseek
|
client, provider = client_and_provider
|
||||||
session_id = create_session(client, system_prompt="system")
|
session_id = create_session(client, system_prompt="system")
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
@@ -241,8 +241,8 @@ def test_streaming_response_is_sse_and_persists_complete_message(
|
|||||||
]
|
]
|
||||||
assert json.loads(payloads[-2])["usage"]["total_tokens"] == 14
|
assert json.loads(payloads[-2])["usage"]["total_tokens"] == 14
|
||||||
assert payloads[-1] == "[DONE]"
|
assert payloads[-1] == "[DONE]"
|
||||||
assert deepseek.completions.calls[0]["stream"] is True
|
assert provider.completions.calls[0]["stream"] is True
|
||||||
assert deepseek.completions.calls[0]["stream_options"] == {
|
assert provider.completions.calls[0]["stream_options"] == {
|
||||||
"include_usage": True
|
"include_usage": True
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -255,12 +255,12 @@ def test_streaming_response_is_sse_and_persists_complete_message(
|
|||||||
|
|
||||||
|
|
||||||
def test_streaming_failure_emits_error_and_does_not_persist(
|
def test_streaming_failure_emits_error_and_does_not_persist(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
|
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
client, deepseek = client_and_deepseek
|
client, provider = client_and_provider
|
||||||
session_id = create_session(client)
|
session_id = create_session(client)
|
||||||
before = (tmp_path / f"{session_id}.json").read_text()
|
before = (tmp_path / f"{session_id}.json").read_text()
|
||||||
deepseek.completions.stream_error = APIConnectionError( # type: ignore[arg-type]
|
provider.completions.stream_error = APIConnectionError( # type: ignore[arg-type]
|
||||||
request=object()
|
request=object()
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -276,9 +276,9 @@ def test_streaming_failure_emits_error_and_does_not_persist(
|
|||||||
|
|
||||||
|
|
||||||
def test_legacy_messages_without_timestamp_remain_usable(
|
def test_legacy_messages_without_timestamp_remain_usable(
|
||||||
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
|
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
client, _ = client_and_deepseek
|
client, _ = client_and_provider
|
||||||
session_id = create_session(client)
|
session_id = create_session(client)
|
||||||
path = tmp_path / f"{session_id}.json"
|
path = tmp_path / f"{session_id}.json"
|
||||||
data = json.loads(path.read_text())
|
data = json.loads(path.read_text())
|
||||||
@@ -296,13 +296,13 @@ def test_legacy_messages_without_timestamp_remain_usable(
|
|||||||
|
|
||||||
def test_same_session_concurrent_messages_are_serialized(tmp_path: Path) -> None:
|
def test_same_session_concurrent_messages_are_serialized(tmp_path: Path) -> None:
|
||||||
async def scenario() -> None:
|
async def scenario() -> None:
|
||||||
deepseek = FakeClient()
|
provider = FakeClient()
|
||||||
deepseek.completions.delay = 0.01
|
provider.completions.delay = 0.01
|
||||||
storage = JsonSessionStorage(tmp_path)
|
storage = JsonSessionStorage(tmp_path)
|
||||||
session = await storage.create("system")
|
session = await storage.create("system")
|
||||||
service = ChatService(
|
service = ChatService(
|
||||||
storage=storage,
|
storage=storage,
|
||||||
client=deepseek, # type: ignore[arg-type]
|
client=provider, # type: ignore[arg-type]
|
||||||
model="deepseek-v4-flash",
|
model="deepseek-v4-flash",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -311,7 +311,7 @@ def test_same_session_concurrent_messages_are_serialized(tmp_path: Path) -> None
|
|||||||
service.generate_response(session.session_id, "second"),
|
service.generate_response(session.session_id, "second"),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert deepseek.completions.calls[1]["messages"] == [
|
assert provider.completions.calls[1]["messages"] == [
|
||||||
{"role": "system", "content": "system"},
|
{"role": "system", "content": "system"},
|
||||||
{"role": "user", "content": "first"},
|
{"role": "user", "content": "first"},
|
||||||
{"role": "assistant", "content": "reply-1"},
|
{"role": "assistant", "content": "reply-1"},
|
||||||
|
|||||||
Reference in New Issue
Block a user