From 9defe45eea4a932258a048b9f8b5a11ac6da1215 Mon Sep 17 00:00:00 2001 From: huaiyj <8699003+huaiyj@user.noreply.gitee.com> Date: Mon, 24 Aug 2026 20:53:07 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=94=AF=E6=8C=81=E4=BC=9A=E8=AF=9D?= =?UTF-8?q?=E6=8A=A5=E5=91=8A=E5=AF=BC=E5=87=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ocai-service/ocai/api/v2_route.py | 32 ++++++- ocai-service/ocai/main.py | 1 + ocai-service/ocai/services/chat_service.py | 86 ++++++++++++++++++- .../tests/api/test_v2_route_endpoints.py | 46 ++++++++++ .../tests/services/test_session_export.py | 62 +++++++++++++ 5 files changed, 225 insertions(+), 2 deletions(-) create mode 100644 ocai-service/tests/services/test_session_export.py diff --git a/ocai-service/ocai/api/v2_route.py b/ocai-service/ocai/api/v2_route.py index 6280b17..7a7ef3c 100644 --- a/ocai-service/ocai/api/v2_route.py +++ b/ocai-service/ocai/api/v2_route.py @@ -30,12 +30,13 @@ P6 起所有 service 通过 :mod:`ocai.api.deps` 的 FastAPI ``Depends`` 提供 from __future__ import annotations +import json import logging from collections.abc import AsyncGenerator, Iterable from typing import TYPE_CHECKING, Annotated, Any from fastapi import APIRouter, Body, Depends, Query, Request -from fastapi.responses import JSONResponse, StreamingResponse +from fastapi.responses import JSONResponse, Response, StreamingResponse from starlette.concurrency import iterate_in_threadpool from ocai.api.deps import get_app_config, get_chat_service, get_user_preferences_service @@ -168,6 +169,35 @@ async def get_session_summary( return result +@router.get("/sessions/{session_id}/export") +async def export_session( + session_id: str, + chat_service: ChatSvc, + user: CurrentUser, + format: str = Query("markdown", pattern="^(markdown|json)$"), +) -> Any: + result = chat_service.export_session(session_id.strip(), user.login_name) + if result is None: + return JSONResponse( + {"error": {"code": "NOT_FOUND", "message": "会话不存在"}}, + status_code=404, + ) + safe_filename = "".join(char if char.isalnum() or char in "-_" else "_" for char in session_id) + if format == "json": + return Response( + content=json.dumps(result, ensure_ascii=False, indent=2, default=str), + media_type="application/json; charset=utf-8", + headers={"Content-Disposition": f'attachment; filename="{safe_filename}.json"'}, + ) + + markdown = chat_service.session_export_to_markdown(result) + return Response( + content=markdown, + media_type="text/markdown; charset=utf-8", + headers={"Content-Disposition": f'attachment; filename="{safe_filename}.md"'}, + ) + + @router.post("/sessions/delete") async def delete_session( chat_service: ChatSvc, diff --git a/ocai-service/ocai/main.py b/ocai-service/ocai/main.py index 09e8a99..30d7f34 100644 --- a/ocai-service/ocai/main.py +++ b/ocai-service/ocai/main.py @@ -226,6 +226,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: cache_repo=cache_repo, message_converter=message_converter, metrics_collector=metrics_collector, + trace_loader=agent_repo.query_conversation_traces, ) user_preferences_service = UserPreferencesService( diff --git a/ocai-service/ocai/services/chat_service.py b/ocai-service/ocai/services/chat_service.py index 0d9d45f..43e1956 100644 --- a/ocai-service/ocai/services/chat_service.py +++ b/ocai-service/ocai/services/chat_service.py @@ -15,13 +15,14 @@ from __future__ import annotations +import json import logging import queue as _queue import re import threading import time import uuid -from collections.abc import Generator +from collections.abc import Callable, Generator from typing import TYPE_CHECKING, Any from ocai.ai.agents.base import AgentContext, AgentEventType @@ -77,6 +78,7 @@ class ChatService(MessagePersistenceMixin): cache_repo: CacheRepository, message_converter: MessageConverter | None = None, metrics_collector: MetricsCollector | None = None, + trace_loader: Callable[..., list[dict]] | None = None, ) -> None: self._agent_manager = agent_manager self._v2_repo = v2_session_repo @@ -84,6 +86,7 @@ class ChatService(MessagePersistenceMixin): self._converter: MessageConverter = message_converter or MessageConverter() self._writer = SSEWriter() self._metrics: MetricsCollector = metrics_collector or MetricsCollector() + self._trace_loader = trace_loader @property def metrics_collector(self) -> MetricsCollector: @@ -172,6 +175,87 @@ class ChatService(MessagePersistenceMixin): "context": {}, } + def export_session(self, session_id: str, user_id: str) -> dict | None: + """Build a portable session document after enforcing ownership.""" + doc = self._v2_repo.get_session(session_id, user_id=user_id) + if doc is None: + return None + + history = self._format_messages_for_history(doc.get("messages", [])) + traces = [] + if self._trace_loader is not None: + try: + traces = self._trace_loader( + session_id=session_id, + user_id=user_id, + limit=100, + ) + except Exception: + logger.warning("Failed to load traces for session export", exc_info=True) + return { + "session_id": session_id, + "summary": doc.get("summary", ""), + "created_at": self._isoformat(doc.get("created_at")), + "updated_at": self._isoformat(doc.get("updated_at")), + "messages": history, + "traces": traces, + } + + @staticmethod + def session_export_to_markdown(export: dict) -> str: + """Render a structured session export as an archival Markdown report.""" + lines = ["# OCAI 会话报告", "", f"- 会话 ID:`{export['session_id']}`"] + if export.get("created_at"): + lines.append(f"- 创建时间:{export['created_at']}") + if export.get("updated_at"): + lines.append(f"- 更新时间:{export['updated_at']}") + if export.get("summary"): + lines.extend(["", "## 摘要", "", str(export["summary"])]) + + lines.extend(["", "## 对话记录"]) + role_names = {"user": "用户", "assistant": "助手", "cancel": "取消"} + for message in export.get("messages", []): + role = role_names.get(message.get("role"), str(message.get("role", "消息"))) + timestamp = message.get("timestamp") or "" + heading = f"### {role}" + if timestamp: + heading += f" ({timestamp})" + lines.extend(["", heading, ""]) + for part in message.get("parts", []): + part_type = part.get("type", "unknown") + content = part.get("content") + if part_type == "text" and content: + lines.append(str(content)) + elif part_type == "thinking" and content: + lines.extend(["
", "推理过程", "", str(content), "", "
"]) + else: + lines.extend( + [ + f"**{part_type}**", + "", + "```json", + json.dumps(part, ensure_ascii=False, indent=2, default=str), + "```", + ] + ) + traces = export.get("traces") or [] + if traces: + lines.extend(["", "## 执行 Trace", ""]) + for trace in traces: + lines.extend( + [ + "```json", + json.dumps(trace, ensure_ascii=False, indent=2, default=str), + "```", + "", + ] + ) + return "\n".join(lines).rstrip() + "\n" + + @staticmethod + def _isoformat(value: Any) -> str: + return value.isoformat() if hasattr(value, "isoformat") else str(value or "") + @staticmethod def _format_messages_for_history(messages: list[dict]) -> list[dict]: """将存储的消息列表转换为前端所需的 parts 格式(兼容新旧格式) diff --git a/ocai-service/tests/api/test_v2_route_endpoints.py b/ocai-service/tests/api/test_v2_route_endpoints.py index a7c5508..df1574e 100644 --- a/ocai-service/tests/api/test_v2_route_endpoints.py +++ b/ocai-service/tests/api/test_v2_route_endpoints.py @@ -92,6 +92,16 @@ class _StubChatService: self._record("handle_user_confirm", **kwargs) return {"status": "confirmed", "confirm_id": kwargs["confirm_id"]} + def export_session(self, session_id: str, user_id: str) -> dict | None: + self._record("export_session", session_id, user_id) + if session_id == "missing": + return None + return {"session_id": session_id, "summary": "summary", "messages": []} + + @staticmethod + def session_export_to_markdown(export: dict) -> str: + return f"# {export['session_id']}\n" + class _StubUserPreferencesService: def __init__(self): @@ -142,6 +152,42 @@ def _build_test_app( return app, TestClient(app) +class SessionExportEndpointTests(unittest.TestCase): + def test_json_export(self): + service = _StubChatService() + _, client = _build_test_app(chat_service=service) + + response = client.get("/api/v2/ocai/sessions/s1/export?format=json") + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json()["session_id"], "s1") + self.assertIn("s1.json", response.headers["content-disposition"]) + self.assertIn(("export_session", ("s1", "alice"), {}), service.calls) + + def test_markdown_export_is_attachment(self): + _, client = _build_test_app() + + response = client.get("/api/v2/ocai/sessions/s1/export") + + self.assertEqual(response.status_code, 200) + self.assertTrue(response.headers["content-type"].startswith("text/markdown")) + self.assertIn("s1.md", response.headers["content-disposition"]) + + def test_missing_session_returns_404(self): + _, client = _build_test_app() + + response = client.get("/api/v2/ocai/sessions/missing/export") + + self.assertEqual(response.status_code, 404) + + def test_invalid_format_returns_422(self): + _, client = _build_test_app() + + response = client.get("/api/v2/ocai/sessions/s1/export?format=pdf") + + self.assertEqual(response.status_code, 422) + + # --------------------------------------------------------------------------- # /health # --------------------------------------------------------------------------- diff --git a/ocai-service/tests/services/test_session_export.py b/ocai-service/tests/services/test_session_export.py new file mode 100644 index 0000000..6cca72e --- /dev/null +++ b/ocai-service/tests/services/test_session_export.py @@ -0,0 +1,62 @@ +from unittest.mock import MagicMock + +from ocai.services.chat_service import ChatService + + +class _Repo: + def get_session(self, session_id, user_id=None): + if session_id != "s1" or user_id != "alice": + return None + return { + "session_id": "s1", + "summary": "磁盘告警排查", + "created_at": "2026-08-24T10:00:00", + "updated_at": "2026-08-24T10:05:00", + "messages": [ + { + "message_id": "m1", + "role": "user", + "content": "检查磁盘", + "timestamp": "2026-08-24T10:00:00", + }, + { + "message_id": "m2", + "role": "assistant", + "timestamp": "2026-08-24T10:01:00", + "parts": [ + {"type": "thinking", "content": "需要查看容量"}, + {"type": "tool_call", "name": "df", "arguments": {"path": "/"}}, + {"type": "text", "content": "根分区已满"}, + ], + }, + ], + } + + +def _service(trace_loader=None): + return ChatService(MagicMock(), _Repo(), MagicMock(), trace_loader=trace_loader) + + +def test_export_session_keeps_structured_parts_and_checks_owner(): + service = _service() + + result = service.export_session("s1", "alice") + + assert result["summary"] == "磁盘告警排查" + assert result["messages"][1]["parts"][1]["type"] == "tool_call" + assert service.export_session("s1", "bob") is None + + +def test_markdown_export_contains_reasoning_and_tool_call(): + trace_loader = MagicMock(return_value=[{"trace_id": "t1", "events": [{"type": "tool_call"}]}]) + service = _service(trace_loader=trace_loader) + exported = service.export_session("s1", "alice") + markdown = service.session_export_to_markdown(exported) + + assert "# OCAI 会话报告" in markdown + assert "磁盘告警排查" in markdown + assert "推理过程" in markdown + assert '"type": "tool_call"' in markdown + assert exported["traces"][0]["trace_id"] == "t1" + assert "## 执行 Trace" in markdown + trace_loader.assert_called_once_with(session_id="s1", user_id="alice", limit=100) -- Gitee