feat(chat): add configurable streaming output
This commit is contained in:
@@ -1,4 +1,6 @@
|
||||
import asyncio
|
||||
from decimal import Decimal
|
||||
from unittest.mock import patch
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -9,7 +11,7 @@ from app.models import Base
|
||||
from app.models.admin import Admin
|
||||
from app.models.ai_config import ModelConfig
|
||||
from app.schemas.admin import AgentRuntimeConfigSaveRequest
|
||||
from app.services.model_stream_service import _openai_stream_payload
|
||||
from app.services.model_stream_service import _openai_stream_payload, _stream_configured_model_async
|
||||
from app.services.model_service import ModelClientService, _max_output_tokens
|
||||
from app.services.rag_service import RagResult
|
||||
|
||||
@@ -58,6 +60,7 @@ def test_runtime_config_defaults_to_long_answer_safe_max_tokens():
|
||||
assert result["modelId"] == model.id
|
||||
assert result["modelName"] == "正式模型"
|
||||
assert result["maxToken"] == 8192
|
||||
assert result["streamEnabled"] == 1
|
||||
assert _max_output_tokens(model) == 8192
|
||||
|
||||
|
||||
@@ -76,6 +79,7 @@ def test_saved_runtime_config_is_persisted_on_enabled_model():
|
||||
presencePenalty=0.2,
|
||||
frequencyPenalty=0.4,
|
||||
maxToken=12000,
|
||||
streamEnabled=0,
|
||||
),
|
||||
db=db,
|
||||
current_admin=admin,
|
||||
@@ -89,6 +93,8 @@ def test_saved_runtime_config_is_persisted_on_enabled_model():
|
||||
assert model.top_k == 40
|
||||
assert model.presence_penalty == Decimal("0.20")
|
||||
assert model.frequency_penalty == Decimal("0.40")
|
||||
assert model.stream_enabled == 0
|
||||
assert result["streamEnabled"] == 0
|
||||
|
||||
payload = _openai_stream_payload(
|
||||
model,
|
||||
@@ -101,6 +107,27 @@ def test_saved_runtime_config_is_persisted_on_enabled_model():
|
||||
assert payload["frequency_penalty"] == 0.4
|
||||
|
||||
|
||||
def test_disabled_stream_returns_one_complete_chunk():
|
||||
model = _model()
|
||||
model.stream_enabled = 0
|
||||
answer = "完整回答" * 20
|
||||
rag_result = RagResult(
|
||||
question="测试",
|
||||
knowledge_scopes=[],
|
||||
chunks=[],
|
||||
prompt="测试",
|
||||
allow_general_knowledge=True,
|
||||
)
|
||||
|
||||
async def collect():
|
||||
return [item async for item in _stream_configured_model_async(model, rag_result)]
|
||||
|
||||
with patch("app.services.model_stream_service._call_configured_model", return_value=answer):
|
||||
chunks = asyncio.run(collect())
|
||||
|
||||
assert chunks == [answer]
|
||||
|
||||
|
||||
def test_agent_debug_does_not_truncate_long_answer(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"app.services.model_service._call_configured_model",
|
||||
|
||||
@@ -113,6 +113,15 @@ def test_queued_chat_reports_position_and_completes(monkeypatch):
|
||||
assert released is True
|
||||
|
||||
|
||||
def test_chat_streaming_response_disables_proxy_buffering():
|
||||
payload = type("Payload", (), {"sessionId": 1, "message": "问题"})()
|
||||
response = chat.completions(payload, object(), object())
|
||||
|
||||
assert response.headers["cache-control"] == "no-cache, no-transform"
|
||||
assert response.headers["x-accel-buffering"] == "no"
|
||||
assert response.media_type == "text/event-stream"
|
||||
|
||||
|
||||
def test_sensitive_system_config_never_returns_plaintext():
|
||||
from app.api.admin_settings import _config_dict
|
||||
|
||||
|
||||
Reference in New Issue
Block a user