feat(agent): stream admin debug preview

This commit is contained in:
2026-07-17 13:47:05 +08:00
parent 13fed0467a
commit bfaf2ebf67
11 changed files with 620 additions and 75 deletions

View File

@@ -1,6 +1,11 @@
from __future__ import annotations
import json
from collections.abc import AsyncIterator
from fastapi import APIRouter, Depends, HTTPException, Query, status
from fastapi.encoders import jsonable_encoder
from fastapi.responses import StreamingResponse
from sqlalchemy import func, select
from sqlalchemy.orm import Session
@@ -20,11 +25,10 @@ from app.schemas.admin import (
SystemConfigSaveRequest,
)
from app.services.admin_service import OperationLogService
from app.services.agent_debug_service import AgentDebugService
from app.services.feishu_service import FeishuKnowledgeService
from app.services.knowledge_agent_service import KnowledgeAgentService
from app.services.knowledge_service import KnowledgeScope
from app.services.model_service import ModelClientService
from app.services.rag_service import PromptService, RagResult
from app.services.secret_service import MASKED_SECRET, SENSITIVE_CONFIG_KEYS, SecretService
router = APIRouter()
@@ -146,37 +150,11 @@ async def debug_agent(
if model is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模型不存在")
rag_result = await KnowledgeAgentService.build_result(
db,
question=payload.question,
version_overrides=payload.knowledgeVersions,
preview_knowledge_ids=payload.knowledgeIds,
)
debug_messages = [
{"role": "system", "content": payload.promptContent.strip()},
*(rag_result.messages or []),
]
rag_result = RagResult(
question=rag_result.question,
knowledge_scopes=rag_result.knowledge_scopes,
chunks=rag_result.chunks,
prompt=PromptService.render_messages(debug_messages),
allow_general_knowledge=rag_result.allow_general_knowledge,
retrieval_log_id=rag_result.retrieval_log_id,
tool_trace=rag_result.tool_trace,
messages=debug_messages,
)
rag_result = await AgentDebugService.build_result(db, payload)
result = await ModelClientService.debug_model_async(
model,
rag_result,
{
"temperature": payload.temperature,
"top_p": payload.topP,
"top_k": payload.topK,
"presence_penalty": payload.presencePenalty,
"frequency_penalty": payload.frequencyPenalty,
"max_token": payload.maxToken,
},
AgentDebugService.overrides(payload),
)
result["retrievalTrace"] = rag_result.tool_trace or []
result["retrievalLogId"] = rag_result.retrieval_log_id
@@ -185,6 +163,42 @@ async def debug_agent(
return api_success(result)
@router.post("/agent/debug/stream")
def debug_agent_stream(
payload: AgentDebugRequest,
db: Session = Depends(get_db),
current_admin: Admin = Depends(get_current_admin),
) -> StreamingResponse:
return StreamingResponse(
_debug_agent_stream(payload, db, current_admin),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache, no-transform",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
},
)
async def _debug_agent_stream(
payload: AgentDebugRequest,
db: Session,
current_admin: Admin,
) -> AsyncIterator[str]:
async for event in AgentDebugService.stream(payload, db, current_admin):
yield _admin_sse_event(event)
yield _admin_sse_done()
def _admin_sse_event(event: dict) -> str:
event = jsonable_encoder(event)
return f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
def _admin_sse_done() -> str:
return "data: [DONE]\n\n"
@router.get("/agent/runtime-config")
def get_agent_runtime_config(
db: Session = Depends(get_db),

View File

@@ -0,0 +1,96 @@
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from sqlalchemy.orm import Session
from app.models.admin import Admin
from app.models.ai_config import ModelConfig
from app.schemas.admin import AgentDebugRequest
from app.services.admin_service import OperationLogService
from app.services.knowledge_agent_service import KnowledgeAgentService
from app.services.model_stream_service import ModelStreamService
from app.services.rag_service import PromptService, RagResult
class AgentDebugService:
@staticmethod
async def build_result(db: Session, payload: AgentDebugRequest) -> RagResult:
rag_result = await KnowledgeAgentService.build_result(
db,
question=payload.question,
version_overrides=payload.knowledgeVersions,
preview_knowledge_ids=payload.knowledgeIds,
)
debug_messages = [
{"role": "system", "content": payload.promptContent.strip()},
*(rag_result.messages or []),
]
return RagResult(
question=rag_result.question,
knowledge_scopes=rag_result.knowledge_scopes,
chunks=rag_result.chunks,
prompt=PromptService.render_messages(debug_messages),
allow_general_knowledge=rag_result.allow_general_knowledge,
retrieval_log_id=rag_result.retrieval_log_id,
tool_trace=rag_result.tool_trace,
messages=debug_messages,
)
@staticmethod
def overrides(payload: AgentDebugRequest) -> dict:
return {
"temperature": payload.temperature,
"top_p": payload.topP,
"top_k": payload.topK,
"presence_penalty": payload.presencePenalty,
"frequency_penalty": payload.frequencyPenalty,
"max_token": payload.maxToken,
}
@classmethod
async def stream(
cls,
payload: AgentDebugRequest,
db: Session,
current_admin: Admin,
) -> AsyncIterator[dict]:
model = db.get(ModelConfig, payload.modelId)
if model is None:
yield {"type": "error", "message": "模型不存在"}
return
try:
yield {"type": "status", "stage": "retrieving", "message": "思考中"}
rag_result = await cls.build_result(db, payload)
model_response = ModelStreamService.debug_stream_async(
model,
rag_result,
cls.overrides(payload),
)
async for chunk in model_response.chunks:
if chunk:
yield {"type": "content", "content": chunk}
OperationLogService.write(
db,
admin_id=current_admin.id,
module="agent",
action="debug_stream",
target_id=model.id,
)
db.commit()
yield {
"type": "complete",
"message": "Agent 调试完成",
"modelName": model_response.model_name,
"retrieveCount": len(rag_result.chunks),
"knowledgeIds": rag_result.knowledge_ids,
"retrievalTrace": rag_result.tool_trace or [],
"retrievalLogId": rag_result.retrieval_log_id,
}
except asyncio.CancelledError:
db.rollback()
raise
except Exception as exc:
db.rollback()
yield {"type": "error", "message": str(exc) or "Agent 调试失败"}

View File

@@ -17,6 +17,7 @@ from app.services.model_service import (
_anthropic_headers,
_auth_headers,
_call_configured_model,
_copy_model_with_overrides,
_decimal_to_float,
_load_extra_params,
_mock_answer,
@@ -117,6 +118,22 @@ class ModelStreamService:
chunks=_stream_configured_model_async(model, rag_result),
)
@staticmethod
def debug_stream_async(
model: ModelConfig,
rag_result: RagResult,
overrides: dict[str, Any],
) -> AsyncStreamingModelResponse:
debug_model = _copy_model_with_overrides(model, overrides)
if not (debug_model.api_url or debug_model.base_url) or not debug_model.api_key:
raise ExternalServiceError("模型 Base URL/API URL 或 API Key 未配置", provider="model")
return AsyncStreamingModelResponse(
model_id=model.id,
model_name=model.model_name,
input_token=_rough_token_count(rag_result.prompt),
chunks=_stream_configured_model_async(debug_model, rag_result),
)
def _get_enabled_model(db: Session) -> ModelConfig | None:
return db.scalar(

View File

@@ -1,17 +1,22 @@
import asyncio
from decimal import Decimal
import json
from unittest.mock import patch
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool
from app.api.admin_settings import get_agent_runtime_config, save_agent_runtime_config
from app.api.admin_settings import _debug_agent_stream, get_agent_runtime_config, save_agent_runtime_config
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, _stream_configured_model_async
from app.schemas.admin import AgentDebugRequest, AgentRuntimeConfigSaveRequest
from app.services.model_stream_service import (
AsyncStreamingModelResponse,
_openai_stream_payload,
_stream_configured_model_async,
)
from app.services.model_service import ModelClientService, _max_output_tokens
from app.services.rag_service import RagResult
@@ -128,6 +133,66 @@ def test_disabled_stream_returns_one_complete_chunk():
assert chunks == [answer]
def test_agent_debug_stream_emits_status_content_and_trace():
async def chunks():
yield "<think>内部思考</think>"
yield "- **第一步**:停一下"
rag_result = RagResult(
question="怎么冷静",
knowledge_scopes=[],
chunks=[],
prompt="怎么冷静",
allow_general_knowledge=True,
retrieval_log_id=9,
tool_trace=[{"tool": "KnowledgeSearch", "status": "success"}],
messages=[{"role": "user", "content": "怎么冷静"}],
)
model_response = AsyncStreamingModelResponse(
model_id=1,
model_name="production-model",
input_token=10,
chunks=chunks(),
)
async def collect(payload, db, admin):
return [event async for event in _debug_agent_stream(payload, db, admin)]
with _database() as db:
admin = _admin()
model = _model()
db.add_all([admin, model])
db.commit()
payload = AgentDebugRequest(
promptContent="你是测试助手",
modelId=model.id,
knowledgeIds=[],
question="怎么冷静",
maxToken=8192,
)
with (
patch(
"app.services.agent_debug_service.KnowledgeAgentService.build_result",
return_value=rag_result,
),
patch(
"app.services.agent_debug_service.ModelStreamService.debug_stream_async",
return_value=model_response,
),
):
events = asyncio.run(collect(payload, db, admin))
decoded = [
json.loads(event.removeprefix("data: "))
for event in events
if event != "data: [DONE]\n\n"
]
assert [event["type"] for event in decoded] == ["status", "content", "content", "complete"]
assert decoded[1]["content"].startswith("<think>")
assert decoded[2]["content"] == "- **第一步**:停一下"
assert decoded[3]["retrievalTrace"][0]["tool"] == "KnowledgeSearch"
def test_agent_debug_does_not_truncate_long_answer(monkeypatch):
monkeypatch.setattr(
"app.services.model_service._call_configured_model",