更新
This commit is contained in:
@@ -8,7 +8,13 @@ from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from doctor_workstation.core.errors import ApiBusinessError, ApiHttpError, ApiProtocolError
|
||||
from doctor_workstation.core.errors import (
|
||||
ApiBusinessError,
|
||||
ApiHttpError,
|
||||
ApiProtocolError,
|
||||
ApiTimeoutError,
|
||||
ApiTransportError,
|
||||
)
|
||||
from doctor_workstation.core.models import Appointment, Consultation, PageResult, Prescription
|
||||
from doctor_workstation.services.mock_repository import DemoDoctorRepository
|
||||
from doctor_workstation.services.repository import (
|
||||
@@ -774,10 +780,11 @@ def test_remote_diagnosis_ai_stream_normalises_chunks_in_order() -> None:
|
||||
assert client.post_calls == []
|
||||
|
||||
|
||||
def test_remote_diagnosis_ai_stream_falls_back_once_but_not_for_error_event() -> None:
|
||||
@pytest.mark.parametrize("status_code", [404, 405])
|
||||
def test_remote_diagnosis_ai_stream_falls_back_for_missing_route(status_code: int) -> None:
|
||||
class MissingStreamClient(RecordingClient):
|
||||
def post_event_stream(self, *args: Any, **kwargs: Any):
|
||||
raise ApiHttpError("missing", status_code=404)
|
||||
raise ApiHttpError("missing", status_code=status_code)
|
||||
|
||||
missing_client = MissingStreamClient()
|
||||
events = list(
|
||||
@@ -789,6 +796,57 @@ def test_remote_diagnosis_ai_stream_falls_back_once_but_not_for_error_event() ->
|
||||
"tcm.diagnosis/aiAssistant"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"stream_error",
|
||||
[
|
||||
ApiHttpError("server error", status_code=500),
|
||||
ApiTimeoutError("timed out"),
|
||||
ApiTransportError("disconnected"),
|
||||
ApiProtocolError("invalid event stream"),
|
||||
],
|
||||
)
|
||||
def test_remote_diagnosis_ai_stream_failures_do_not_issue_full_request(
|
||||
stream_error: Exception,
|
||||
) -> None:
|
||||
class FailedStreamClient(RecordingClient):
|
||||
def post_event_stream(self, *args: Any, **kwargs: Any):
|
||||
raise stream_error
|
||||
|
||||
client = FailedStreamClient()
|
||||
with pytest.raises(type(stream_error)):
|
||||
list(RemoteDoctorRepository(client).stream_diagnosis_ai(501, "请分析"))
|
||||
assert client.post_calls == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("after_start", [False, True])
|
||||
def test_remote_diagnosis_ai_stream_incomplete_response_does_not_fall_back(
|
||||
after_start: bool,
|
||||
) -> None:
|
||||
class IncompleteStreamClient(RecordingClient):
|
||||
def post_event_stream(self, *args: Any, **kwargs: Any):
|
||||
if after_start:
|
||||
yield {"event": "start", "data": {"model_key": "qwen"}}
|
||||
|
||||
client = IncompleteStreamClient()
|
||||
with pytest.raises(ApiProtocolError, match="ended before done"):
|
||||
list(RemoteDoctorRepository(client).stream_diagnosis_ai(501, "请分析"))
|
||||
assert client.post_calls == []
|
||||
|
||||
|
||||
def test_remote_diagnosis_ai_stream_does_not_fall_back_after_start() -> None:
|
||||
class FailedAfterStartClient(RecordingClient):
|
||||
def post_event_stream(self, *args: Any, **kwargs: Any):
|
||||
yield {"event": "start", "data": {"model_key": "qwen"}}
|
||||
raise ApiHttpError("missing", status_code=404)
|
||||
|
||||
client = FailedAfterStartClient()
|
||||
with pytest.raises(ApiHttpError):
|
||||
list(RemoteDoctorRepository(client).stream_diagnosis_ai(501, "请分析"))
|
||||
assert client.post_calls == []
|
||||
|
||||
|
||||
def test_remote_diagnosis_ai_stream_error_event_does_not_fall_back() -> None:
|
||||
class ErrorStreamClient(RecordingClient):
|
||||
def post_event_stream(self, *args: Any, **kwargs: Any):
|
||||
yield {"event": "error", "data": {"message": "模型繁忙"}}
|
||||
|
||||
Reference in New Issue
Block a user