refactor(stream): 将不完整流 token 估算逻辑收敛到 StreamContext

- 新增 has_partial_response / ensure_estimated_output_tokens / should_estimate_incomplete_tokens 方法
- CancelledError 路径在归因前即补充 output_tokens,确保日志包含估算值
- CLI Handler 和 Chat Handler 的兜底估算统一使用 should_estimate_incomplete_tokens
- 移除 cli_monitor_mixin 和 stream_telemetry 中重复的条件判断
- 新增对应单元测试
This commit is contained in:
fawney19
2026-03-17 03:37:06 +08:00
parent 40e0b82fa0
commit 438f16094f
6 changed files with 159 additions and 22 deletions

View File

@@ -33,6 +33,12 @@ async def _cancel_immediately() -> AsyncGenerator[bytes, None]:
raise asyncio.CancelledError()
async def _yield_once_then_cancel(ctx: StreamContext) -> AsyncGenerator[bytes, None]:
ctx.append_text("partial output")
yield b"data: chunk\n\n"
raise asyncio.CancelledError()
@pytest.mark.asyncio
async def test_create_monitored_stream_marks_client_disconnected_when_confirmed() -> None:
monitor = _DummyMonitor()
@@ -70,7 +76,9 @@ async def test_create_monitored_stream_marks_server_cancelled_when_confirmed_con
@pytest.mark.asyncio
async def test_create_monitored_stream_marks_cancelled_unknown_when_disconnect_check_uncertain() -> None:
async def test_create_monitored_stream_marks_cancelled_unknown_when_disconnect_check_uncertain() -> (
None
):
monitor = _DummyMonitor()
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-unknown")
@@ -85,3 +93,23 @@ async def test_create_monitored_stream_marks_cancelled_unknown_when_disconnect_c
assert ctx.status_code == 503
assert ctx.error_message == "cancelled_unknown"
assert "cancel_origin=cancelled_unknown" in (ctx.upstream_response or "")
@pytest.mark.asyncio
async def test_create_monitored_stream_estimates_output_tokens_before_unknown_cancel_log() -> None:
monitor = _DummyMonitor()
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-estimate")
request = _RequestStub([asyncio.TimeoutError()])
monitored = monitor._create_monitored_stream(ctx, _yield_once_then_cancel(ctx), request)
with pytest.raises(asyncio.CancelledError):
async for _ in monitored:
pass
expected_output_tokens = max(1, len("partial output") // 4)
assert ctx.status_code == 503
assert ctx.error_message == "cancelled_unknown"
assert ctx.output_tokens == expected_output_tokens
assert f"output_tokens={expected_output_tokens}" in (ctx.upstream_response or "")

View File

@@ -129,3 +129,34 @@ def test_get_log_summary_without_first_byte_time() -> None:
assert "TTFB:" not in summary
assert "Total: 456ms" in summary
assert "in:100 out:50" in summary
def test_ensure_estimated_output_tokens_uses_collected_text() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.append_text("partial output")
changed = ctx.ensure_estimated_output_tokens()
assert changed is True
assert ctx.output_tokens == max(1, len("partial output") // 4)
def test_should_estimate_incomplete_tokens_for_interrupted_partial_stream() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.status_code = 503
ctx.chunk_count = 3
assert ctx.should_estimate_incomplete_tokens() is True
def test_should_estimate_incomplete_tokens_when_output_already_estimated() -> None:
"""ensure_estimated_output_tokens 已补了 output但 input 仍为 0 时仍需估算。"""
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.status_code = 503
ctx.chunk_count = 3
ctx.append_text("partial")
ctx.ensure_estimated_output_tokens()
assert ctx.output_tokens > 0
assert ctx.input_tokens == 0
assert ctx.should_estimate_incomplete_tokens() is True

View File

@@ -0,0 +1,66 @@
from __future__ import annotations
import time
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from src.api.handlers.base import stream_telemetry as stream_telemetry_module
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
class _DummyDb:
def close(self) -> None:
pass
@pytest.mark.asyncio
async def test_record_stream_stats_estimates_tokens_for_failed_partial_stream(
monkeypatch: pytest.MonkeyPatch,
) -> None:
recorder = StreamTelemetryRecorder(
request_id="req-telemetry",
user_id="1",
api_key_id="2",
client_ip="127.0.0.1",
format_id="openai:chat",
)
recorder._get_telemetry_writer = AsyncMock( # type: ignore[method-assign]
return_value=SimpleNamespace(include_bodies=False)
)
recorder._dispatch_record = AsyncMock() # type: ignore[method-assign]
recorder._update_candidate_status = AsyncMock() # type: ignore[method-assign]
ctx = StreamContext(
model="test-model",
api_format="openai:chat",
request_id="req-telemetry",
user_id=1,
api_key_id=2,
)
ctx.provider_name = "test-provider"
ctx.status_code = 503
ctx.data_count = 2
ctx.chunk_count = 4
ctx.append_text("partial output")
monkeypatch.setattr(stream_telemetry_module, "get_db", lambda: iter([_DummyDb()]))
monkeypatch.setattr(
stream_telemetry_module.SystemConfigService,
"should_log_body",
lambda _db: False,
)
monkeypatch.setattr(stream_telemetry_module.config, "stream_stats_delay", 0)
await recorder.record_stream_stats(
ctx,
original_headers={},
original_request_body={"input": [{"content": "hello world"}]},
start_time=time.time(),
)
assert ctx.input_tokens > 0
assert ctx.output_tokens == max(1, len("partial output") // 4)
recorder._dispatch_record.assert_awaited_once() # type: ignore[attr-defined]