diff --git a/.env.example b/.env.example index 46882b4..f630023 100644 --- a/.env.example +++ b/.env.example @@ -8,6 +8,9 @@ # ── 基础设施 (Infrastructure) ──────────────────────────────────────── # NATS_URL=nats://127.0.0.1:4222 # REDIS_ADDR=127.0.0.1:6379 +# REDIS_URL=redis://127.0.0.1:6379/0 +# AGENTHUB_DISTRIBUTED_CACHE_ENABLED=true +# AGENTHUB_MEMORY_EVENTS_ENABLED=true # DATABASE_DSN=postgres://agenthub:agenthub@localhost:5432/agenthub?sslmode=disable # ── 安全 (Security) ────────────────────────────────────────────────── @@ -35,6 +38,13 @@ # OPENAI_COMPATIBLE_BASE_URL=http://127.0.0.1:11434/v1 # OPENAI_COMPATIBLE_API_KEY=not-needed +# 国产模型精确 token 统计(指向本地 tokenizer.json 或其目录) +# AGENTHUB_TOKENIZER_QWEN_PATH=/models/qwen/tokenizer.json +# AGENTHUB_TOKENIZER_DEEPSEEK_PATH=/models/deepseek/tokenizer.json +# AGENTHUB_TOKENIZER_DOUBAO_PATH=/models/doubao/tokenizer.json +# AGENTHUB_TOKENIZER_ZHIPU_PATH=/models/glm/tokenizer.json +# AGENTHUB_TOKENIZER_MOONSHOT_PATH=/models/kimi/tokenizer.json + # vLLM 本地推理(优先级高于 OPENAI_COMPATIBLE_*) # 启用 vLLM 服务:docker compose -f deploy/docker-compose.platform.yml --profile vllm up # 调用方式:model 字段传 vllm-,如 vllm-Qwen/Qwen2.5-7B-Instruct diff --git a/.gitignore b/.gitignore index 407be88..38d45f7 100644 --- a/.gitignore +++ b/.gitignore @@ -9,6 +9,7 @@ __pycache__/ # ─── Frontend ───────────────────────────────────────────────────── frontend/node_modules/ frontend/.next/ +frontend/.next-*/ frontend/out/ frontend/.env*.local diff --git a/app/api/admin/workflows.py b/app/api/admin/workflows.py index c364d64..7a9b545 100644 --- a/app/api/admin/workflows.py +++ b/app/api/admin/workflows.py @@ -14,17 +14,17 @@ from __future__ import annotations -import json - from fastapi import APIRouter, Depends, HTTPException -from app.db.init_db import now from app.db.session import aexecute -from app.schemas.common import AgentRouteActiveRequest, AgentRouteRequest -from app.schemas.dag import DAGConfig +from app.schemas.common import AgentRouteActiveRequest +from app.schemas.workflow import AgentRouteRequest, WorkflowDraftRequest, WorkflowValidationRequest from app.services.agent_route_service import agent_route_service from app.services.auth_service import get_current_user, require_admin, write_audit -from app.services.template_engine import template_engine +from app.services.context_summary_cache import context_summary_cache +from app.services.workflow_contract import validate_workflow_contract +from app.services.workflow_draft_service import workflow_draft_service +from app.services.workflow_errors import WorkflowVersionConflict router = APIRouter(prefix="/workflows", tags=["admin-workflows"]) @@ -44,6 +44,59 @@ async def list_workflows(user: dict = Depends(get_current_user)) -> list[dict]: return await agent_route_service.list_routes(_uid(user)) +@router.post("/validate") +async def validate_workflow(data: WorkflowValidationRequest, user: dict = Depends(get_current_user)) -> dict: + """Validate and normalize an editor graph without persisting it.""" + require_admin(user) + result = validate_workflow_contract(data.nodes, data.edges, schema_version=data.schemaVersion) + return result.model_dump(mode="json") + + +@router.get("/drafts") +async def list_workflow_drafts(user: dict = Depends(get_current_user)) -> list[dict]: + require_admin(user) + return await workflow_draft_service.list_drafts(_uid(user)) + + +@router.get("/drafts/{draft_key}") +async def get_workflow_draft(draft_key: str, user: dict = Depends(get_current_user)) -> dict: + require_admin(user) + draft = await workflow_draft_service.get_draft(_uid(user), draft_key) + if not draft: + raise HTTPException(status_code=404, detail="Workflow draft not found") + return draft + + +@router.put("/drafts/{draft_key}") +async def save_workflow_draft( + draft_key: str, data: WorkflowDraftRequest, user: dict = Depends(get_current_user), +) -> dict: + require_admin(user) + try: + return await workflow_draft_service.save_draft(_uid(user), draft_key, data) + except WorkflowVersionConflict as exc: + raise _version_conflict(exc, "workflow_draft_version_conflict") from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + +@router.delete("/drafts/{draft_key}") +async def delete_workflow_draft(draft_key: str, user: dict = Depends(get_current_user)) -> dict: + require_admin(user) + if not await workflow_draft_service.delete_draft(_uid(user), draft_key): + raise HTTPException(status_code=404, detail="Workflow draft not found") + return {"status": "success", "draftKey": draft_key} + + +@router.get("/{route_id}") +async def get_workflow(route_id: int, user: dict = Depends(get_current_user)) -> dict: + require_admin(user) + route = await agent_route_service.get_route(route_id, _uid(user)) + if not route: + raise HTTPException(status_code=404, detail="Workflow not found") + return route + + # ── CREATE ──────────────────────────────────────────────────────────────── @@ -54,7 +107,15 @@ async def create_workflow(data: AgentRouteRequest, user: dict = Depends(get_curr uid = _uid(user) try: route = await agent_route_service.create_route( - uid, data.name, data.description, data.triggerKeywords, data.nodes, data.isDefault, + uid, + data.name, + data.description, + data.triggerKeywords, + data.nodes, + edges=data.edges, + is_default=data.isDefault, + active=data.active, + schema_version=data.schemaVersion, ) except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc @@ -75,29 +136,28 @@ async def update_workflow(route_id: int, data: AgentRouteRequest, user: dict = D require_admin(user) uid = _uid(user) - existing = await agent_route_service.get_route(route_id, uid) - if not existing: - raise HTTPException(status_code=404, detail="Workflow not found") - - dag = DAGConfig(total=len(data.nodes), completed=0, nodes=data.nodes) - template_engine.validate(dag) - - if data.isDefault: - await aexecute("UPDATE agent_routes SET is_default = 0 WHERE user_id = $1", uid) - await aexecute( - "UPDATE agent_routes SET name = $1, description = $2, trigger_keywords = $3, " - "nodes_json = $4, is_default = $5, updated_at = $6 WHERE id = $7 AND user_id = $8", - data.name, - data.description, - json.dumps(data.triggerKeywords, ensure_ascii=False), - json.dumps(data.nodes, ensure_ascii=False), - 1 if data.isDefault else 0, - now(), - route_id, - uid, - ) - - route = await agent_route_service.get_route(route_id, uid) + if data.version < 1: + raise HTTPException(status_code=428, detail="Workflow version is required for updates") + try: + route = await agent_route_service.update_route( + route_id, + uid, + name=data.name, + description=data.description, + trigger_keywords=data.triggerKeywords, + nodes=data.nodes, + edges=data.edges, + is_default=data.isDefault, + active=data.active, + schema_version=data.schemaVersion, + expected_version=data.version, + ) + except WorkflowVersionConflict as exc: + raise _version_conflict(exc, "workflow_version_conflict") from exc + except LookupError as exc: + raise HTTPException(status_code=404, detail=str(exc)) from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc audit_id = write_audit( user["id"], "admin", "workflow_update", "L2", "approve", {"routeId": route_id, "name": data.name}, @@ -105,6 +165,18 @@ async def update_workflow(route_id: int, data: AgentRouteRequest, user: dict = D return {"status": "success", "route": route, "auditId": audit_id} +def _version_conflict(exc: WorkflowVersionConflict, code: str) -> HTTPException: + return HTTPException( + status_code=409, + detail={ + "code": code, + "message": str(exc), + "expectedVersion": exc.expected_version, + "currentVersion": exc.current_version, + }, + ) + + # ── DELETE ──────────────────────────────────────────────────────────────── @@ -119,6 +191,7 @@ async def delete_workflow(route_id: int, user: dict = Depends(get_current_user)) raise HTTPException(status_code=404, detail="Workflow not found") await aexecute("DELETE FROM agent_routes WHERE id = $1 AND user_id = $2", route_id, uid) + context_summary_cache.invalidate("route", uid) audit_id = write_audit( user["id"], "admin", "workflow_delete", "L2", "approve", diff --git a/app/api/agent.py b/app/api/agent.py index a711a95..24e19d7 100644 --- a/app/api/agent.py +++ b/app/api/agent.py @@ -306,6 +306,9 @@ async def create_agent(data: AgentCreateRequest, user: dict = Depends(get_curren ) except Exception as exc: raise HTTPException(status_code=400, detail="Agent 已存在或参数无效") from exc + from app.services.context_summary_cache import context_summary_cache + context_summary_cache.invalidate("agent", str(user["id"])) + context_summary_cache.invalidate("agent", "shared") audit_id = write_audit(user["id"], agent_id, "agent_create", "L2", "approve", {**data.model_dump(), "apiKey": "***" if data.apiKey else ""}) return {"status": "success", "agentId": agent_id, "auditId": audit_id} @@ -319,6 +322,9 @@ async def delete_agent(agent_id: str, user: dict = Depends(get_current_user)) -> ) if not deleted: raise HTTPException(status_code=404, detail="Agent 不存在") + from app.services.context_summary_cache import context_summary_cache + context_summary_cache.invalidate("agent", str(user["id"])) + context_summary_cache.invalidate("agent", "shared") audit_id = write_audit(user["id"], agent_id, "agent_delete", "L2", "approve", {"agentId": agent_id}) return {"status": "success", "agentId": agent_id, "auditId": audit_id} @@ -368,6 +374,9 @@ async def update_agent(agent_id: str, data: AgentUpdateRequest, user: dict = Dep avatar_url, tags_json, data.baseUrl.strip(), config_json, agent_id, user["id"], ) + from app.services.context_summary_cache import context_summary_cache + context_summary_cache.invalidate("agent", str(user["id"])) + context_summary_cache.invalidate("agent", "shared") audit_id = write_audit(user["id"], agent_id, "agent_update", "L2", "approve", {**data.model_dump(), "apiKey": "***" if data.apiKey else ""}) return {"status": "success", "agentId": agent_id, "auditId": audit_id} diff --git a/app/api/memory.py b/app/api/memory.py index 189bc1c..4170a6e 100644 --- a/app/api/memory.py +++ b/app/api/memory.py @@ -8,10 +8,20 @@ from pydantic import BaseModel, Field from app.config import AUTO_MEMORY_ENABLED, MEMORY_DIR -from app.services.memory import MemoryDocument, MemoryHeader, MemoryScanner, MemoryStorage, MemoryType +from app.services.memory import ( + CognitiveMemoryType, + MemoryDocument, + MemoryHeader, + MemoryScanner, + MemoryScope, + MemoryStorage, + MemoryType, +) from app.services.memory.consolidator import MemoryConsolidator from app.services.memory.extractor import MemoryExtractor from app.services.memory.session_memory import SessionMemoryManager +from app.services.memory.semantic_memory import SemanticMemoryStore +from app.services.memory.procedural_memory import ProceduralMemoryCatalog from app.services.auth.service import get_current_user from app.services.auth.session_guard import check_session_access from app.db.session import afetch_all @@ -91,6 +101,14 @@ def _get_session_mgr(user_id: str = "") -> SessionMemoryManager: return _session_mgr_shared +def _get_semantic_store(user_id: str) -> SemanticMemoryStore: + return SemanticMemoryStore(_get_user_memory_dir(user_id or "local-admin")) + + +def _get_procedural_catalog(user_id: str) -> ProceduralMemoryCatalog: + return ProceduralMemoryCatalog(user_id, _get_storage(user_id)) + + # -- Pydantic request/response models ----------------------------------- @@ -100,6 +118,9 @@ class MemoryCreateRequest(BaseModel): type: MemoryType = Field(MemoryType.REFERENCE, description="记忆类型") body: str = Field("", description="记忆内容正文(Markdown)") filename: Optional[str] = Field(None, description="可选:指定文件名,默认从 name 自动生成") + memory_type: CognitiveMemoryType = CognitiveMemoryType.SEMANTIC + scope: MemoryScope = MemoryScope.USER + source: str = Field("manual", min_length=1, max_length=128) class MemoryUpdateRequest(BaseModel): @@ -107,6 +128,9 @@ class MemoryUpdateRequest(BaseModel): description: Optional[str] = Field(None, max_length=512) type: Optional[MemoryType] = None body: Optional[str] = None + memory_type: Optional[CognitiveMemoryType] = None + scope: Optional[MemoryScope] = None + source: Optional[str] = Field(None, min_length=1, max_length=128) class MemoryFileInfo(BaseModel): @@ -117,6 +141,10 @@ class MemoryFileInfo(BaseModel): mtime: float created_at: str = "" updated_at: str = "" + memory_type: str = CognitiveMemoryType.SEMANTIC.value + scope: str = MemoryScope.USER.value + source: str = "legacy-file" + version: int = 1 class MemoryDetail(BaseModel): @@ -150,12 +178,43 @@ class SessionTransferRequest(BaseModel): mode: str = Field("append", description="append 或 overwrite") +@router.get("/semantic") +async def list_semantic_memories( + query: str = Query("", max_length=1000), + active_only: bool = Query(True), + user: dict = Depends(get_current_user), +) -> dict: + """List structured semantic memories extracted from episodic summaries.""" + store = _get_semantic_store(str(user["id"])) + records = ( + await store.search(query, limit=50) + if query else await store.list_records(active_only=active_only) + ) + return { + "count": len(records), + "records": [record.__dict__ for record in records], + } + + +@router.get("/procedural") +async def list_procedural_memories( + query: str = Query("", max_length=1000), + user: dict = Depends(get_current_user), +) -> dict: + """List the unified read-through catalog for skills, workflows and policies.""" + user_id = str(user["id"]) + catalog = _get_procedural_catalog(user_id) + records = await catalog.search(query, limit=100) if query else await catalog.list_records() + return {"count": len(records), "records": [record.to_dict() for record in records]} + + # -- endpoints ----------------------------------------------------------- @router.get("/files", response_model=list[MemoryFileInfo]) async def list_memories( type_filter: Optional[str] = Query(None, alias="type"), + memory_type_filter: Optional[str] = Query(None, alias="memory_type"), user: dict = Depends(get_current_user), ): """List all memory files with headers, optionally filtered by type.""" @@ -168,6 +227,12 @@ async def list_memories( headers = [h for h in headers if h.type == mt] except ValueError: raise HTTPException(status_code=400, detail=f"无效的记忆类型: {type_filter}") + if memory_type_filter: + try: + cognitive_type = CognitiveMemoryType(memory_type_filter) + headers = [h for h in headers if h.memory_type == cognitive_type] + except ValueError: + raise HTTPException(status_code=400, detail=f"无效的认知记忆类型: {memory_type_filter}") return [ MemoryFileInfo( filename=h.filename, @@ -177,6 +242,10 @@ async def list_memories( mtime=h.mtime, created_at=h.created_at, updated_at=h.updated_at, + memory_type=h.memory_type.value, + scope=h.scope.value, + source=h.source, + version=h.version, ) for h in headers ] @@ -198,6 +267,10 @@ async def read_memory(filename: str, user: dict = Depends(get_current_user)): "type": doc.meta.type.value, "created_at": doc.meta.created_at, "updated_at": doc.meta.updated_at, + "memory_type": doc.meta.memory_type.value, + "scope": doc.meta.scope.value, + "source": doc.meta.source, + "version": doc.meta.version, }, body=doc.body, ) @@ -369,6 +442,10 @@ async def import_memories( type_=doc.meta.type, body=doc.body, filename=filename, + memory_type=doc.meta.memory_type, + scope=doc.meta.scope, + source=doc.meta.source, + version=doc.meta.version, ) imported.append(filename) except Exception as exc: @@ -396,6 +473,9 @@ async def create_memory(req: MemoryCreateRequest, user: dict = Depends(get_curre type_=req.type, body=req.body, filename=req.filename, + memory_type=req.memory_type, + scope=req.scope, + source=req.source, ) fname = Path(doc.file_path).name return MemoryFileInfo( @@ -406,6 +486,10 @@ async def create_memory(req: MemoryCreateRequest, user: dict = Depends(get_curre mtime=0, created_at=doc.meta.created_at, updated_at=doc.meta.updated_at, + memory_type=doc.meta.memory_type.value, + scope=doc.meta.scope.value, + source=doc.meta.source, + version=doc.meta.version, ) @@ -422,6 +506,9 @@ async def update_memory(filename: str, req: MemoryUpdateRequest, user: dict = De new_desc = req.description if req.description is not None else doc.meta.description new_type = req.type if req.type is not None else doc.meta.type new_body = req.body if req.body is not None else doc.body + new_memory_type = req.memory_type if req.memory_type is not None else doc.meta.memory_type + new_scope = req.scope if req.scope is not None else doc.meta.scope + new_source = req.source if req.source is not None else doc.meta.source await storage.save( name=new_name, @@ -429,6 +516,9 @@ async def update_memory(filename: str, req: MemoryUpdateRequest, user: dict = De type_=new_type, body=new_body, filename=filename, + memory_type=new_memory_type, + scope=new_scope, + source=new_source, ) # Re-read to get fresh metadata updated = await storage.get(filename) @@ -442,6 +532,10 @@ async def update_memory(filename: str, req: MemoryUpdateRequest, user: dict = De mtime=0, created_at=updated.meta.created_at, updated_at=updated.meta.updated_at, + memory_type=updated.meta.memory_type.value, + scope=updated.meta.scope.value, + source=updated.meta.source, + version=updated.meta.version, ) diff --git a/app/api/websocket.py b/app/api/websocket.py index 94dfdc5..734fda1 100644 --- a/app/api/websocket.py +++ b/app/api/websocket.py @@ -706,13 +706,21 @@ async def _append_turn_to_session_memory( """ try: store = _get_session_store(user_id) - await store.append_turn( + turn_count = await store.append_turn( session_id=session_id, user_message=user_message, agent_response=agent_response, sender=sender or "user", agent_name=agent_name or "assistant", ) + if turn_count >= 10 and turn_count % 10 == 0: + try: + from app.services.memory_summary_consumer import memory_summary_consumer + await memory_summary_consumer.request_compaction( + session_id, user_id or "local-admin", + ) + except Exception: + logger.debug("memory compaction request failed", exc_info=True) # Invalidate the memory context cache so the next agent call # picks up the updated session memory. try: diff --git a/app/db/init_db.py b/app/db/init_db.py index e624ec7..30178cd 100644 --- a/app/db/init_db.py +++ b/app/db/init_db.py @@ -97,9 +97,18 @@ def _default_password_hash() -> str: retry_count INTEGER DEFAULT 0, error_type TEXT, session_id TEXT, - created_at TEXT NOT NULL + created_at TEXT NOT NULL, + memory_type TEXT NOT NULL DEFAULT 'episodic', + memory_scope TEXT NOT NULL DEFAULT 'session', + memory_source TEXT NOT NULL DEFAULT 'task_execution', + memory_version INTEGER NOT NULL DEFAULT 1 )""", + """ALTER TABLE task_execution_history ADD COLUMN IF NOT EXISTS memory_type TEXT NOT NULL DEFAULT 'episodic'""", + """ALTER TABLE task_execution_history ADD COLUMN IF NOT EXISTS memory_scope TEXT NOT NULL DEFAULT 'session'""", + """ALTER TABLE task_execution_history ADD COLUMN IF NOT EXISTS memory_source TEXT NOT NULL DEFAULT 'task_execution'""", + """ALTER TABLE task_execution_history ADD COLUMN IF NOT EXISTS memory_version INTEGER NOT NULL DEFAULT 1""", """CREATE INDEX IF NOT EXISTS idx_teh_agent_type ON task_execution_history(assigned_agent, task_type)""", + """CREATE INDEX IF NOT EXISTS idx_teh_memory_type_scope ON task_execution_history(memory_type, memory_scope, session_id)""", """CREATE TABLE IF NOT EXISTS dag_templates ( id SERIAL PRIMARY KEY, name TEXT NOT NULL, @@ -132,12 +141,29 @@ def _default_password_hash() -> str: description TEXT NOT NULL DEFAULT '', trigger_keywords TEXT NOT NULL DEFAULT '[]', nodes_json TEXT NOT NULL, + edges_json TEXT NOT NULL DEFAULT '[]', is_default INTEGER NOT NULL DEFAULT 0, active INTEGER NOT NULL DEFAULT 1, + version INTEGER NOT NULL DEFAULT 1, + schema_version INTEGER NOT NULL DEFAULT 1, created_at TEXT NOT NULL, updated_at TEXT NOT NULL )""", """CREATE UNIQUE INDEX IF NOT EXISTS idx_agent_routes_name_user ON agent_routes(name, user_id)""", + """CREATE TABLE IF NOT EXISTS workflow_drafts ( + id SERIAL PRIMARY KEY, + user_id TEXT NOT NULL, + workflow_id INTEGER, + draft_key TEXT NOT NULL, + name TEXT NOT NULL DEFAULT '', + payload_json TEXT NOT NULL, + base_version INTEGER NOT NULL DEFAULT 0, + version INTEGER NOT NULL DEFAULT 1, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + UNIQUE(user_id, draft_key) + )""", + """CREATE INDEX IF NOT EXISTS idx_workflow_drafts_user_updated ON workflow_drafts(user_id, updated_at DESC)""", """CREATE TABLE IF NOT EXISTS audit_log ( id TEXT PRIMARY KEY, user_id TEXT NOT NULL, @@ -368,7 +394,7 @@ async def _apply_alembic_migrations(conn) -> None: current = row["version_num"] if row else None # Current head revision (must match migrations/versions/) - head = "ff209a40779d" + head = "c1a7d4e82b6f" if current == head: logger.info("init_db: Alembic already at head (%s)", head) @@ -553,6 +579,39 @@ async def _migrate_agent_routes_pg(conn) -> None: except Exception as exc: logger.warning("agent_routes user_id migration skipped: %s", exc) + for migration in ( + "ALTER TABLE agent_routes ADD COLUMN IF NOT EXISTS edges_json TEXT NOT NULL DEFAULT '[]'", + "ALTER TABLE agent_routes ADD COLUMN IF NOT EXISTS version INTEGER NOT NULL DEFAULT 1", + "ALTER TABLE agent_routes ADD COLUMN IF NOT EXISTS schema_version INTEGER NOT NULL DEFAULT 1", + ): + try: + await conn.execute(migration) + except Exception as exc: + logger.warning("agent_routes editor migration skipped: %s", exc) + + try: + await conn.execute( + """CREATE TABLE IF NOT EXISTS workflow_drafts ( + id SERIAL PRIMARY KEY, + user_id TEXT NOT NULL, + workflow_id INTEGER, + draft_key TEXT NOT NULL, + name TEXT NOT NULL DEFAULT '', + payload_json TEXT NOT NULL, + base_version INTEGER NOT NULL DEFAULT 0, + version INTEGER NOT NULL DEFAULT 1, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + UNIQUE(user_id, draft_key) + )""" + ) + await conn.execute( + "CREATE INDEX IF NOT EXISTS idx_workflow_drafts_user_updated " + "ON workflow_drafts(user_id, updated_at DESC)" + ) + except Exception as exc: + logger.warning("workflow_drafts migration skipped: %s", exc) + # 2. Drop old unique constraint on name (single-column) try: await conn.execute("ALTER TABLE agent_routes DROP CONSTRAINT IF EXISTS agent_routes_name_key") diff --git a/app/schemas/common.py b/app/schemas/common.py index 9ff8c6a..623b22a 100644 --- a/app/schemas/common.py +++ b/app/schemas/common.py @@ -2,6 +2,8 @@ from pydantic import BaseModel, Field +from app.schemas.workflow import AgentRouteRequest + class LoginRequest(BaseModel): name: str @@ -74,14 +76,6 @@ class AuditConfirmRequest(BaseModel): payload: dict = Field(default_factory=dict) -class AgentRouteRequest(BaseModel): - name: str - description: str = "" - triggerKeywords: list[str] = Field(default_factory=list) - nodes: list[dict] - isDefault: bool = False - - class AgentRouteActiveRequest(BaseModel): active: bool diff --git a/app/schemas/dag.py b/app/schemas/dag.py index 9128bc8..73d3a3d 100644 --- a/app/schemas/dag.py +++ b/app/schemas/dag.py @@ -1,23 +1,54 @@ from __future__ import annotations -from pydantic import BaseModel, Field +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field + + +DAG_NODE_TYPES = Literal[ + "start", "agent", "tool", "ifelse", "end", "code", "http", "knowledge", "human" +] + + +class DAGEdge(BaseModel): + """Stable editor edge contract; aliases preserve the existing React shape.""" + + model_config = ConfigDict(populate_by_name=True) + + id: str = "" + source: str = Field(alias="from") + target: str = Field(alias="to") + label: str = "" + condition: str = "" class DAGNode(BaseModel): + model_config = ConfigDict(extra="allow") + id: str + type: DAG_NODE_TYPES = "agent" + name: str = "" domain: str = "" # canvas flow nodes may not have this agent: str = "" # canvas flow nodes may not have this description: str = "" # canvas flow nodes may not have this dependencies: list[str] = Field(default_factory=list) + x: float = 0 + y: float = 0 + config: dict[str, Any] = Field(default_factory=dict) status: str = "PENDING" priority: int = 1 # 1=highest, 3=lowest estimated_effort: str = "medium" # low|medium|high class DAGConfig(BaseModel): - total: int + model_config = ConfigDict(populate_by_name=True) + + schema_version: int = Field(default=1, alias="schemaVersion") + version: int = 1 + total: int = 0 completed: int = 0 nodes: list[DAGNode] + edges: list[DAGEdge] = Field(default_factory=list) templateId: int | None = None templateName: str | None = None similarity: float | None = None diff --git a/app/schemas/workflow.py b/app/schemas/workflow.py new file mode 100644 index 0000000..5bed39e --- /dev/null +++ b/app/schemas/workflow.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +from typing import Any, Literal + +from pydantic import BaseModel, Field + + +class AgentRouteRequest(BaseModel): + name: str = Field(min_length=1, max_length=128) + description: str = Field(default="", max_length=4000) + triggerKeywords: list[str] = Field(default_factory=list, max_length=100) + nodes: list[dict[str, Any]] = Field(default_factory=list) + edges: list[dict[str, Any]] = Field(default_factory=list) + isDefault: bool = False + active: bool = True + schemaVersion: int = Field(default=1, ge=1) + version: int = Field(default=0, ge=0) + + +class WorkflowValidationRequest(BaseModel): + nodes: list[dict[str, Any]] = Field(default_factory=list) + edges: list[dict[str, Any]] = Field(default_factory=list) + schemaVersion: int = Field(default=1, ge=1) + + +class DAGValidationIssue(BaseModel): + code: str + message: str + severity: Literal["error", "warning"] = "error" + nodeId: str | None = None + edgeId: str | None = None + + +class DAGValidationResult(BaseModel): + valid: bool + normalized: dict[str, Any] | None = None + issues: list[DAGValidationIssue] = Field(default_factory=list) + + +class WorkflowDraftRequest(AgentRouteRequest): + name: str = Field(default="", max_length=128) + workflowId: int | None = None + baseVersion: int = Field(default=0, ge=0) + draftVersion: int = Field(default=0, ge=0) diff --git a/app/services/adapter_manager.py b/app/services/adapter_manager.py index 239be2e..666a87a 100644 --- a/app/services/adapter_manager.py +++ b/app/services/adapter_manager.py @@ -984,6 +984,9 @@ class OllamaAdapter(BaseAdapter): async def execute_prompt(self, prompt: str, model: str, api_key: str = "", base_url: str = "", **kwargs: Any) -> str: url = (base_url.rstrip("/") if base_url else OLLAMA_BASE_URL) + "/api/generate" payload = {"model": model or "llama3", "prompt": prompt, "stream": False} + system_prompt = str(kwargs.get("system_prompt") or "") + if system_prompt: + payload["system"] = system_prompt try: response = await _retry_request("POST", url, json_body=payload) if response.status_code >= 400: @@ -1002,6 +1005,8 @@ async def execute_prompt(self, prompt: str, model: str, api_key: str = "", base_ async def stream_prompt(self, prompt: str, model: str, api_key: str = "", base_url: str = "", *, system_prompt: str = "") -> AsyncGenerator[str, None]: url = (base_url.rstrip("/") if base_url else OLLAMA_BASE_URL) + "/api/generate" payload = {"model": model or "llama3", "prompt": prompt, "stream": True} + if system_prompt: + payload["system"] = system_prompt self.last_usage = {} full_text = "" try: diff --git a/app/services/agent_route_service.py b/app/services/agent_route_service.py index f46bd77..ae0392b 100644 --- a/app/services/agent_route_service.py +++ b/app/services/agent_route_service.py @@ -4,76 +4,184 @@ from typing import Any from app.db.init_db import now -from app.db.session import afetch_all, afetch_one, aexecute, aexecute_insert +from app.db.session import afetch_all, afetch_one, aexecute, aexecute_insert, atransaction from app.schemas.dag import DAGConfig +from app.services.context_summary_cache import context_summary_cache from app.services.template_engine import template_engine +from app.services.workflow_contract import require_valid_workflow +from app.services.workflow_errors import WorkflowVersionConflict + + +_ROUTE_COLUMNS = ( + "id,name,description,trigger_keywords,nodes_json,edges_json,is_default,active," + "version,schema_version,created_at,updated_at" +) class AgentRouteService: async def list_routes(self, user_id: str, active_only: bool = False) -> list[dict[str, Any]]: if active_only: - sql = "SELECT id,name,description,trigger_keywords,nodes_json,is_default,active,created_at,updated_at FROM agent_routes WHERE user_id=$1 AND active=1 ORDER BY is_default DESC,id DESC" - rows = await afetch_all(sql, user_id) - else: - sql = "SELECT id,name,description,trigger_keywords,nodes_json,is_default,active,created_at,updated_at FROM agent_routes WHERE user_id=$1 ORDER BY is_default DESC,id DESC" - rows = await afetch_all(sql, user_id) - for row in rows: - row["triggerKeywords"] = json.loads(row.pop("trigger_keywords") or "[]") - row["nodes"] = json.loads(row.pop("nodes_json") or "[]") - row["isDefault"] = bool(row.pop("is_default")) - row["active"] = bool(row["active"]) - return rows - - async def create_route(self, user_id: str, name: str, description: str, trigger_keywords: list[str], nodes: list[dict[str, Any]], is_default: bool = False) -> dict[str, Any]: - dag = DAGConfig(total=len(nodes), completed=0, nodes=nodes) - template_engine.validate(dag) + cached = context_summary_cache.get("route", user_id, "active-routes") + if cached is not None: + return json.loads(cached) + active_clause = " AND active=1" if active_only else "" + rows = await afetch_all( + f"SELECT {_ROUTE_COLUMNS} FROM agent_routes WHERE user_id=$1{active_clause} " + "ORDER BY is_default DESC,id DESC", + user_id, + ) + routes = [self._deserialize_route(row) for row in rows] + if not active_only: + return routes + serialized = json.dumps(routes, ensure_ascii=False, sort_keys=True, default=str) + context_summary_cache.set("route", user_id, "active-routes", serialized) + return json.loads(serialized) + + async def create_route( + self, + user_id: str, + name: str, + description: str, + trigger_keywords: list[str], + nodes: list[dict[str, Any]], + edges: list[dict[str, Any]] | None = None, + is_default: bool = False, + active: bool = True, + schema_version: int = 1, + ) -> dict[str, Any]: + normalized = require_valid_workflow(nodes, edges, schema_version=schema_version) existing = await afetch_one("SELECT id FROM agent_routes WHERE name=$1 AND user_id=$2", name, user_id) if existing: raise ValueError(f"Agent route with name '{name}' already exists") if is_default: - await aexecute("UPDATE agent_routes SET is_default=0 WHERE user_id=$1", user_id) + await aexecute( + "UPDATE agent_routes SET is_default=0,version=version+1,updated_at=$1 " + "WHERE user_id=$2 AND is_default=1", + now(), + user_id, + ) route_id = await aexecute_insert( - "INSERT INTO agent_routes(name,user_id,description,trigger_keywords,nodes_json,is_default,active,created_at,updated_at) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9) RETURNING id", - name, user_id, description, json.dumps(trigger_keywords, ensure_ascii=False), json.dumps(nodes, ensure_ascii=False), 1 if is_default else 0, 1, now(), now(), + "INSERT INTO agent_routes(name,user_id,description,trigger_keywords,nodes_json,edges_json," + "is_default,active,version,schema_version,created_at,updated_at) " + "VALUES($1,$2,$3,$4,$5,$6,$7,$8,1,$9,$10,$11) RETURNING id", + name, + user_id, + description, + json.dumps(trigger_keywords, ensure_ascii=False), + json.dumps(normalized["nodes"], ensure_ascii=False), + json.dumps(normalized["edges"], ensure_ascii=False), + 1 if is_default else 0, + 1 if active else 0, + schema_version, + now(), + now(), ) route = await self.get_route(int(route_id), user_id) if not route: raise ValueError("Agent route create failed") + context_summary_cache.invalidate("route", user_id) + return route + + async def update_route( + self, + route_id: int, + user_id: str, + *, + name: str, + description: str, + trigger_keywords: list[str], + nodes: list[dict[str, Any]], + edges: list[dict[str, Any]], + is_default: bool, + active: bool, + schema_version: int, + expected_version: int, + ) -> dict[str, Any]: + normalized = require_valid_workflow(nodes, edges, schema_version=schema_version) + async with atransaction() as conn: + updated = await conn.fetchrow( + "UPDATE agent_routes SET name=$1,description=$2,trigger_keywords=$3,nodes_json=$4," + "edges_json=$5,is_default=$6,active=$7,schema_version=$8,version=version+1,updated_at=$9 " + "WHERE id=$10 AND user_id=$11 AND version=$12 RETURNING id,version", + name, + description, + json.dumps(trigger_keywords, ensure_ascii=False), + json.dumps(normalized["nodes"], ensure_ascii=False), + json.dumps(normalized["edges"], ensure_ascii=False), + 1 if is_default else 0, + 1 if active else 0, + schema_version, + now(), + route_id, + user_id, + expected_version, + ) + if updated and is_default: + await conn.execute( + "UPDATE agent_routes SET is_default=0,version=version+1,updated_at=$1 " + "WHERE user_id=$2 AND id<>$3 AND is_default=1", + now(), + user_id, + route_id, + ) + if not updated: + current = await afetch_one( + "SELECT version FROM agent_routes WHERE id=$1 AND user_id=$2", route_id, user_id, + ) + if not current: + raise LookupError("Workflow not found") + raise WorkflowVersionConflict(expected_version, int(current["version"])) + route = await self.get_route(route_id, user_id) + if not route: + raise LookupError("Workflow not found") + context_summary_cache.invalidate("route", user_id) return route async def get_route(self, route_id: int, user_id: str) -> dict[str, Any] | None: - row = await afetch_one("SELECT id,name,description,trigger_keywords,nodes_json,is_default,active,created_at,updated_at FROM agent_routes WHERE id=$1 AND user_id=$2", route_id, user_id) - if not row: - return None - row["triggerKeywords"] = json.loads(row.pop("trigger_keywords") or "[]") - row["nodes"] = json.loads(row.pop("nodes_json") or "[]") - row["isDefault"] = bool(row.pop("is_default")) - row["active"] = bool(row["active"]) - return row + row = await afetch_one( + f"SELECT {_ROUTE_COLUMNS} FROM agent_routes WHERE id=$1 AND user_id=$2", route_id, user_id, + ) + return self._deserialize_route(row) if row else None async def set_default(self, route_id: int, user_id: str) -> dict[str, Any]: if not await self.get_route(route_id, user_id): raise ValueError("Agent route not found") - await aexecute("UPDATE agent_routes SET is_default=0 WHERE user_id=$1", user_id) - await aexecute("UPDATE agent_routes SET is_default=1,active=1,updated_at=$1 WHERE id=$2 AND user_id=$3", now(), route_id, user_id) + await aexecute( + "UPDATE agent_routes SET is_default=0,version=version+1,updated_at=$1 " + "WHERE user_id=$2 AND id<>$3 AND is_default=1", + now(), user_id, route_id, + ) + await aexecute( + "UPDATE agent_routes SET is_default=1,active=1,version=version+1,updated_at=$1 " + "WHERE id=$2 AND user_id=$3", + now(), route_id, user_id, + ) route = await self.get_route(route_id, user_id) if not route: raise ValueError("Agent route not found") + context_summary_cache.invalidate("route", user_id) return route async def set_active(self, route_id: int, user_id: str, active: bool) -> dict[str, Any]: if not await self.get_route(route_id, user_id): raise ValueError("Agent route not found") - await aexecute("UPDATE agent_routes SET active=$1,updated_at=$2 WHERE id=$3 AND user_id=$4", 1 if active else 0, now(), route_id, user_id) + await aexecute( + "UPDATE agent_routes SET active=$1,version=version+1,updated_at=$2 WHERE id=$3 AND user_id=$4", + 1 if active else 0, now(), route_id, user_id, + ) route = await self.get_route(route_id, user_id) if not route: raise ValueError("Agent route not found") + context_summary_cache.invalidate("route", user_id) return route async def resolve_dag(self, intent: str, user_id: str = "") -> tuple[DAGConfig, int | None, dict[str, Any] | None]: route = await self._match_route(intent, user_id) if route: - dag = DAGConfig(total=len(route["nodes"]), completed=0, nodes=route["nodes"]) + dag = DAGConfig( + total=len(route["nodes"]), nodes=route["nodes"], edges=route.get("edges", []), + version=route.get("version", 1), schema_version=route.get("schemaVersion", 1), + ) template_engine.validate(dag) return dag, None, route dag, template_id = await template_engine.match_template(intent) @@ -82,12 +190,10 @@ async def resolve_dag(self, intent: str, user_id: str = "") -> tuple[DAGConfig, async def _match_route(self, intent: str, user_id: str = "") -> dict[str, Any] | None: routes = await self.list_routes(user_id, active_only=True) intent_lower = intent.lower() - for route in routes: explicit_tokens = [f"#route:{route['name'].lower()}", f"#路线:{route['name']}", f"@路线:{route['name']}"] if any(token in intent_lower or token in intent for token in explicit_tokens): return route - best_score = 0.0 best_route: dict[str, Any] | None = None for route in routes: @@ -97,49 +203,37 @@ async def _match_route(self, intent: str, user_id: str = "") -> dict[str, Any] | hits = sum(1 for keyword in keywords if keyword.lower() in intent_lower or keyword in intent) score = hits / max(len(keywords), 1) if score > best_score: - best_score = score - best_route = route + best_score, best_route = score, route if best_route and best_score >= 0.25: return best_route - return next((route for route in routes if route.get("isDefault")), None) - async def extract_route_ref(self, content: str, user_id: str = "") -> tuple[dict[str, Any] | None, str]: - """Extract ``#route:name`` / ``#路线:name`` from content and match against routes. - - Returns ``(matched_route, stripped_content)`` where ``matched_route`` is - the route dict if a match was found and ``stripped_content`` has the - route token removed. Returns ``(None, content)`` if no match. - """ import re - # Match patterns like: #route:标准研发闭环 or #路线:快速代码生成 - # Also handles leading/trailing whitespace and full-width chars in names - pattern = r'(?:^|\s)(?:#route:|#路线:|@路线:)\s*([^\s]+)' - match = re.search(pattern, content) + match = re.search(r"(?:^|\s)(?:#route:|#路线:|@路线:)\s*([^\s]+)", content) if not match: return None, content - - route_name = match.group(1).strip() - routes = await self.list_routes(user_id, active_only=True) - - # Exact match first, then case-insensitive - matched = None - route_name_lower = route_name.lower() - for route in routes: - if route["name"] == route_name or route["name"].lower() == route_name_lower: - matched = route - break - + route_name = match.group(1).strip().lower() + matched = next( + (route for route in await self.list_routes(user_id, active_only=True) if route["name"].lower() == route_name), + None, + ) if not matched: return None, content - - # Strip the matched token from content - stripped = content[:match.start()] + content[match.end():] - stripped = stripped.strip() - - return matched, stripped + return matched, (content[:match.start()] + content[match.end():]).strip() + + @staticmethod + def _deserialize_route(row: dict[str, Any]) -> dict[str, Any]: + route = dict(row) + route["triggerKeywords"] = json.loads(route.pop("trigger_keywords") or "[]") + route["nodes"] = json.loads(route.pop("nodes_json") or "[]") + route["edges"] = json.loads(route.pop("edges_json") or "[]") + route["isDefault"] = bool(route.pop("is_default")) + route["active"] = bool(route["active"]) + route["schemaVersion"] = int(route.pop("schema_version")) + route["version"] = int(route["version"]) + return route agent_route_service = AgentRouteService() diff --git a/app/services/agent_service.py b/app/services/agent_service.py index eb4c8c3..5e0bec5 100644 --- a/app/services/agent_service.py +++ b/app/services/agent_service.py @@ -26,6 +26,15 @@ ) from app.services import prompt_sections from app.services.prompt_cache import prompt_cache +from app.services.prompt_messages import split_prompt_for_adapter +from app.services.response_quality import estimate_response_quality +from app.services.token_budget import ( + TokenBudget, + cognitive_memory_budgets, + count_tokens, + fit_prompt, + truncate_to_tokens, +) from app.services.secret_service import decrypt_secret from app.services.text_processing import ( filter_streaming_chunk, @@ -159,7 +168,7 @@ async def seed_default_agents_for_user(user_id: str) -> None: # window, reducing disk I/O while staying fresh enough for cross-session # memory. Invalidation is explicit via _invalidate_memory_cache() when # new memories are written (extraction, /memory commands). -_MEMORY_CACHE: dict[str, Any] = {"context": "", "ts": 0.0, "ttl": 300.0, "key": ""} +_MEMORY_CONTEXT_CACHE: dict[str, tuple[float, str]] = {} _SESSION_MGRS: dict[str, Any] = {} @@ -1122,7 +1131,34 @@ async def call_agent(session_id: str, content: str, user_id: str, attachments: l ) models = choose_models(await candidate_models_for_role(agent["agent_id"], user_id)) history = await _build_conversation_history(session_id) - memory_ctx = await _build_memory_context(user_id=user_id, session_id=session_id) + primary_model = models[0] if models else {} + provider = primary_model.get("provider", "") + model_name = primary_model.get("model_name", "") + budget = TokenBudget.for_model(provider, model_name) + memory_budgets = cognitive_memory_budgets( + budget.section_limit("history") + budget.section_limit("memory"), content, domain, + ) + history_tokens_before = count_tokens(history, provider, model_name) + history, history_truncated = truncate_to_tokens( + history, memory_budgets["working"], provider, model_name, + ) + from app.services.performance_monitor import monitor + monitor.record_token_compaction( + "history", + history_tokens_before, + count_tokens(history, provider, model_name), + history_truncated, + ) + memory_ctx = await _build_memory_context( + user_id=user_id, + session_id=session_id, + history=history, + provider=provider, + model=model_name, + max_tokens=sum(value for key, value in memory_budgets.items() if key != "working"), + section_budgets=memory_budgets, + query=content, + ) # Use tool-enabled call loop (handles tool detection, execution, synthesis) result, usage, selected = await _run_tool_call_loop( @@ -1134,6 +1170,8 @@ async def call_agent(session_id: str, content: str, user_id: str, attachments: l ) content_out = normalize_agent_output(agent["agent_id"], result, content) + from app.services.performance_monitor import monitor + monitor.record_answer_quality(estimate_response_quality(content, content_out)) adapter = adapter_manager.get_adapter(selected.get("provider", "mock")) if usage and usage.get("total_tokens", 0) > 0: prompt_tokens, completion_tokens, total_tokens = usage["prompt_tokens"], usage["completion_tokens"], usage["total_tokens"] @@ -1241,9 +1279,40 @@ async def on_chunk(chunk: str) -> None: async def _run_loop(): try: + history = await _build_conversation_history(session_id) + primary_model = models[0] if models else {} + provider = primary_model.get("provider", "") + model_name = primary_model.get("model_name", "") + budget = TokenBudget.for_model(provider, model_name) + memory_budgets = cognitive_memory_budgets( + budget.section_limit("history") + budget.section_limit("memory"), + content, + agent["domain"], + ) + history_tokens_before = count_tokens(history, provider, model_name) + history, history_truncated = truncate_to_tokens( + history, memory_budgets["working"], provider, model_name, + ) + from app.services.performance_monitor import monitor + monitor.record_token_compaction( + "history", + history_tokens_before, + count_tokens(history, provider, model_name), + history_truncated, + ) + memory_context = await _build_memory_context( + user_id=user_id, + session_id=session_id, + history=history, + provider=provider, + model=model_name, + max_tokens=sum(value for key, value in memory_budgets.items() if key != "working"), + section_budgets=memory_budgets, + query=content, + ) r, u, s = await _run_tool_call_loop( session_id, content, user_id, agent, agent["domain"], llm_input, - models, await _build_conversation_history(session_id), await _build_memory_context(), + models, history, memory_context, collab_ctx, token=token, on_tool_event=on_tool_event, streaming_executor=_get_streaming_executor(), stream_callback=on_chunk, @@ -1305,6 +1374,9 @@ async def _run_loop(): _, result, usage, selected = loop_result content_out = normalize_agent_output(agent["agent_id"], result, content) + from app.services.performance_monitor import monitor + monitor.record_answer_quality(estimate_response_quality(content, content_out)) + # ── Stream performance metrics ──────────────────────────── t_end = time.perf_counter() total_ms = (t_end - t_start) * 1000 @@ -1391,93 +1463,126 @@ async def _build_conversation_history(session_id: str, max_chars: int = 3600) -> return result -async def _build_memory_context(user_id: str = "", session_id: str = "", max_chars: int = 2200, force: bool = False) -> str: - """Load current-session conversation memory and global summary. - - Only loads what is strictly needed for conversational continuity: - 1. Current session's raw conversation memory (from session_store) - 2. Current session's LLM-generated summary - 3. Global cross-session summary - - File-backed memories (MEMORY.md etc.) are deliberately NOT loaded here — - they were adding 200+ file I/O operations per message for marginal value. - Agents access them on-demand via the ``memory_search`` tool instead. +async def _build_memory_context( + user_id: str = "", + session_id: str = "", + *, + history: str = "", + provider: str = "", + model: str = "", + max_tokens: int = 3000, + query: str = "", + section_budgets: dict[str, int] | None = None, + force: bool = False, +) -> str: + """Build a budgeted, deduplicated L0/L1/L3 memory projection. - Cached with a 60 s TTL. Call with force=True to bypass the cache. + DB history is treated as L0 working context and excluded from the memory + projection. Session summary is L1, file conversation is the recent durable + tail, and global summary is L3 cross-session memory. L2 knowledge remains + retrieval-only through memory_search. """ - uid = user_id or "local-admin" - cache_key = f"{uid}:{session_id}" if session_id else uid - now_ts = time.monotonic() - if not force and _MEMORY_CACHE.get("key") == cache_key and _MEMORY_CACHE["context"] and (now_ts - _MEMORY_CACHE["ts"]) < _MEMORY_CACHE["ttl"]: - return _MEMORY_CACHE["context"] + import hashlib from app.config import MEMORY_DIR from app.services.memory.storage import MemoryStorage from app.services.memory.session_memory import SessionMemoryManager from app.services.memory.session_store import SessionMemoryStore + from app.services.memory.semantic_memory import SemanticMemoryStore + from app.services.memory.procedural_memory import ProceduralMemoryCatalog + from app.services.memory_context import MemoryContextSection, build_memory_context + from app.services.performance_monitor import monitor + + uid = user_id or "local-admin" + history_fp = hashlib.sha256(history.encode("utf-8")).hexdigest()[:12] + query_fp = hashlib.sha256(query.encode("utf-8")).hexdigest()[:12] + budget_fp = hashlib.sha256( + json.dumps(section_budgets or {}, sort_keys=True).encode("utf-8") + ).hexdigest()[:8] + cache_key = f"{uid}:{session_id}:{provider}:{model}:{max_tokens}:{history_fp}:{query_fp}:{budget_fp}" + now_ts = time.monotonic() + cached = _MEMORY_CONTEXT_CACHE.get(cache_key) + if not force and cached and now_ts - cached[0] < 60.0: + return cached[1] user_memory_dir = MEMORY_DIR / "users" / uid storage = MemoryStorage(user_memory_dir) session_mgr = SessionMemoryManager(storage) session_store = SessionMemoryStore(user_memory_dir) + semantic_store = SemanticMemoryStore(user_memory_dir) + procedural_catalog = ProceduralMemoryCatalog(uid, storage) + sections: list[MemoryContextSection] = [] - sections: list[str] = [] - - # ── Current session conversation memory (highest priority) ───── if session_id: try: - conv = await session_store.get_conversation( - session_id, max_chars=max_chars, - ) - if conv and len(conv) > 100: - sections.append( - "【当前会话对话记忆】\n" - "以下是你与用户在本会话中的对话记录:\n\n" - f"{conv}" - ) + session_summary = await session_mgr.get_session_summary(session_id) + if session_summary: + sections.append(MemoryContextSection("session-summary", session_summary, 1, "episodic")) except Exception: - pass + logger.debug("memory context: session summary unavailable", exc_info=True) - # ── Current session summary (LLM-generated) ──────────────────── - if session_id: try: - sess_summary = await session_mgr.get_session_summary(session_id) - if sess_summary: - sections.append( - "【当前会话摘要】\n" - f"{sess_summary}" - ) + conversation = await session_store.get_conversation( + session_id, max_chars=max(2200, max_tokens * 4), recent_turns=6, + ) + if conversation: + sections.append(MemoryContextSection("recent-durable-memory", conversation, 3, "episodic")) except Exception: - pass + logger.debug("memory context: durable conversation unavailable", exc_info=True) - # ── Global summary (cross-session aggregated) ───────────────── try: - global_summary = await session_mgr.get_global_summary() - if global_summary: - sections.append( - "【全局记忆 — 跨会话聚合摘要】\n" - f"{global_summary}" + semantic_records = await semantic_store.search(query, limit=6) + if semantic_records: + semantic_text = "\n".join( + f"- [{record.category}] {record.value} " + f"(confidence={record.confidence:.2f}, source={record.source})" + for record in semantic_records ) + sections.append(MemoryContextSection("semantic-memory", semantic_text, 2, "semantic")) except Exception: - pass + logger.debug("memory context: semantic memory unavailable", exc_info=True) - if not sections: - _MEMORY_CACHE["context"] = "" - _MEMORY_CACHE["ts"] = now_ts - _MEMORY_CACHE["key"] = cache_key - return "" + try: + global_summary = await session_mgr.get_global_summary() + if global_summary: + sections.append(MemoryContextSection("global-summary", global_summary, 4, "semantic")) + except Exception: + logger.debug("memory context: global summary unavailable", exc_info=True) - result = "\n\n".join(sections) + "\n─── 以上为记忆上下文,以下是当前对话 ───\n" - _MEMORY_CACHE["context"] = result - _MEMORY_CACHE["ts"] = now_ts - _MEMORY_CACHE["key"] = cache_key + if not section_budgets or section_budgets.get("procedural", 0) > 128: + try: + procedures = await procedural_catalog.search(query, limit=6) + if procedures: + procedure_text = "\n".join( + f"- [{record.kind}] {record.name}: {record.description} " + f"(source={record.source}, version={record.source_version})" + for record in procedures + ) + sections.append(MemoryContextSection("procedural-memory", procedure_text, 2, "procedural")) + except Exception: + logger.debug("memory context: procedural memory unavailable", exc_info=True) + + result, stats = build_memory_context( + sections, + exclude_texts=[history], + max_tokens=max_tokens, + provider=provider, + model=model, + section_budgets=section_budgets, + ) + monitor.record_token_compaction( + "memory", + int(stats["tokens_before"]), + int(stats["tokens_after"]), + bool(stats["truncated"]), + ) + _MEMORY_CONTEXT_CACHE[cache_key] = (now_ts, result) return result def _invalidate_memory_cache() -> None: """Clear the memory context cache so next call rebuilds it.""" - _MEMORY_CACHE["context"] = "" - _MEMORY_CACHE["ts"] = 0.0 + _MEMORY_CONTEXT_CACHE.clear() async def _load_settings() -> dict[str, Any]: @@ -1971,18 +2076,10 @@ async def _run_tool_call_loop( # reuse on iterations 1-4, appending only the dynamic conversation # content. This saves ~3-6KB of repeated string construction per # iteration. - _loop_prefix_key = prompt_cache.make_prefix_key( - agent["agent_id"], domain, not simple_mode, - tuple(sorted(available_tools)) if available_tools else None, - session_id, preprocess_context, - models[0].get("provider", "") if models else "", - models[0].get("model_name", "") if models else "", - ) _loop_prefix_cached: str | None = None # populated on iteration 0 - _loop_prefix_cached_mode: str | None = None # 记录缓存时使用的 mode,mode 切换时失效 async def _loop_body() -> tuple[str, dict, dict]: - nonlocal final_text, usage, selected, adapter, _loop_prefix_cached, _loop_prefix_cached_mode + nonlocal final_text, usage, selected, adapter, _loop_prefix_cached # ── Circuit-breaker counters (reset when a tool succeeds) ──────── # Three-tier detection: # 1. Same-tool-same-error: the deadliest loop — agent calls the same @@ -2048,14 +2145,28 @@ async def _loop_body() -> tuple[str, dict, dict]: preprocess_context=preprocess_context, permission_mode=get_permission_mode_for_session(session_id).value, ) - # Split and cache the system prefix for subsequent iterations. - # The anchor "符号消息:" is present in every build_prompt() - # return path (CodeGen, Orchestrator, Architect, General). - anchor = "符号消息:" - split_idx = prompt.rfind(anchor) - if split_idx > 0: - _loop_prefix_cached = prompt[:split_idx] - _loop_prefix_cached_mode = get_permission_mode_for_session(session_id).value + + primary_model = models[0] if models else {} + prompt, prompt_budget_stats = fit_prompt( + prompt, + primary_model.get("provider", ""), + primary_model.get("model_name", ""), + anchor="符号消息:", + ) + from app.services.performance_monitor import monitor + monitor.record_token_compaction( + "prompt", + int(prompt_budget_stats["tokens_before"]), + int(prompt_budget_stats["tokens_after"]), + bool(prompt_budget_stats["truncated"]), + ) + + # Send the static prefix once as the system message. Only the + # dynamic symbolic message and conversation are user content. + system_prefix, user_prompt = split_prompt_for_adapter(prompt) + if system_prefix: + _loop_prefix_cached = system_prefix + prompt = user_prompt logger.info( "tool_loop iter=%d: prompt_len=%d has_tool_section=%s", diff --git a/app/services/context_summary_cache.py b/app/services/context_summary_cache.py new file mode 100644 index 0000000..2091ef9 --- /dev/null +++ b/app/services/context_summary_cache.py @@ -0,0 +1,112 @@ +from __future__ import annotations + +import hashlib +import threading +import time +from dataclasses import dataclass +from typing import Any, Callable + + +@dataclass +class _Entry: + value: Any + version: int + expires_at: float + + +class ContextSummaryCache: + """Shared versioned cache for compact route and agent prompt summaries.""" + + def __init__(self, ttl_seconds: float = 300.0) -> None: + self._ttl = ttl_seconds + self._entries: dict[str, _Entry] = {} + self._versions: dict[str, int] = {} + self._hits = 0 + self._misses = 0 + self._lock = threading.Lock() + + @staticmethod + def _owner_key(kind: str, owner_id: str) -> str: + return f"{kind}:{owner_id}" + + def get_or_build( + self, + kind: str, + owner_id: str, + fingerprint_source: str, + builder: Callable[[], Any], + ) -> Any: + owner_key = self._owner_key(kind, owner_id) + digest = hashlib.sha256(fingerprint_source.encode("utf-8")).hexdigest()[:24] + now = time.monotonic() + with self._lock: + version = self._versions.get(owner_key, 0) + cache_key = f"{owner_key}:{version}:{digest}" + entry = self._entries.get(cache_key) + if entry and entry.version == version and entry.expires_at > now: + self._hits += 1 + return entry.value + self._misses += 1 + + value = builder() + with self._lock: + self._entries[cache_key] = _Entry(value, version, now + self._ttl) + return value + + def get(self, kind: str, owner_id: str, slot: str) -> Any | None: + owner_key = self._owner_key(kind, owner_id) + now = time.monotonic() + with self._lock: + version = self._versions.get(owner_key, 0) + entry = self._entries.get(f"{owner_key}:{version}:slot:{slot}") + if entry and entry.expires_at > now: + self._hits += 1 + return entry.value + self._misses += 1 + return None + + def set(self, kind: str, owner_id: str, slot: str, value: Any) -> None: + owner_key = self._owner_key(kind, owner_id) + with self._lock: + version = self._versions.get(owner_key, 0) + self._entries[f"{owner_key}:{version}:slot:{slot}"] = _Entry( + value, version, time.monotonic() + self._ttl, + ) + + def invalidate(self, kind: str, owner_id: str, *, propagate: bool = True) -> int: + owner_key = self._owner_key(kind, owner_id) + with self._lock: + self._versions[owner_key] = self._versions.get(owner_key, 0) + 1 + stale = [key for key in self._entries if key.startswith(owner_key + ":")] + for key in stale: + self._entries.pop(key, None) + version = self._versions[owner_key] + if propagate: + try: + from app.services.distributed_cache_versions import distributed_cache_version_bus + + distributed_cache_version_bus.schedule(kind, owner_id) + except Exception: + pass + return version + + def stats(self) -> dict[str, int | float]: + with self._lock: + total = self._hits + self._misses + return { + "entries": len(self._entries), + "hits": self._hits, + "misses": self._misses, + "hit_ratio": round(self._hits / total, 4) if total else 0.0, + "versions": len(self._versions), + } + + def clear(self) -> None: + with self._lock: + self._entries.clear() + self._versions.clear() + self._hits = 0 + self._misses = 0 + + +context_summary_cache = ContextSummaryCache() diff --git a/app/services/distributed_cache_versions.py b/app/services/distributed_cache_versions.py new file mode 100644 index 0000000..1c718e6 --- /dev/null +++ b/app/services/distributed_cache_versions.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +import asyncio +import json +import logging +import os +import uuid +from typing import Any + + +logger = logging.getLogger("agenthub.cache.version_bus") + + +class DistributedCacheVersionBus: + CHANNEL = "agenthub:cache:versions" + HASH_KEY = "agenthub:cache:version-counters" + + def __init__(self) -> None: + self._redis: Any = None + self._pubsub: Any = None + self._listener: asyncio.Task | None = None + self._instance_id = uuid.uuid4().hex[:12] + + async def start(self) -> bool: + if os.getenv("AGENTHUB_DISTRIBUTED_CACHE_ENABLED", "true").lower() not in {"1", "true", "yes"}: + return False + try: + import redis.asyncio as redis + + redis_url = os.getenv("REDIS_URL", "").strip() + if not redis_url: + redis_url = f"redis://{os.getenv('REDIS_ADDR', '127.0.0.1:6379')}/0" + self._redis = redis.from_url( + redis_url, + encoding="utf-8", + decode_responses=True, + socket_connect_timeout=2, + socket_timeout=2, + ) + await asyncio.wait_for(self._redis.ping(), timeout=3) + self._pubsub = self._redis.pubsub() + await self._pubsub.subscribe(self.CHANNEL) + self._listener = asyncio.create_task(self._listen(), name="cache-version-listener") + logger.info("distributed cache version bus connected") + return True + except Exception as exc: + logger.warning("distributed cache version bus disabled: %s", exc) + await self.close() + return False + + def schedule(self, kind: str, owner_id: str) -> None: + if self._redis is None: + return + try: + asyncio.get_running_loop().create_task(self.publish(kind, owner_id)) + except RuntimeError: + return + + async def publish(self, kind: str, owner_id: str) -> int: + if self._redis is None: + return 0 + key = f"{kind}:{owner_id}" + version = int(await self._redis.hincrby(self.HASH_KEY, key, 1)) + await self._redis.publish( + self.CHANNEL, + json.dumps({ + "kind": kind, + "owner_id": owner_id, + "version": version, + "instance_id": self._instance_id, + }), + ) + return version + + async def close(self) -> None: + if self._listener is not None: + self._listener.cancel() + try: + await self._listener + except asyncio.CancelledError: + pass + self._listener = None + if self._pubsub is not None: + await self._pubsub.aclose() + self._pubsub = None + if self._redis is not None: + await self._redis.aclose() + self._redis = None + + async def _listen(self) -> None: + assert self._pubsub is not None + async for message in self._pubsub.listen(): + if message.get("type") != "message": + continue + try: + payload = json.loads(message.get("data") or "{}") + if payload.get("instance_id") == self._instance_id: + continue + kind = str(payload["kind"]) + owner_id = str(payload["owner_id"]) + from app.services.context_summary_cache import context_summary_cache + + context_summary_cache.invalidate(kind, owner_id, propagate=False) + except Exception: + logger.warning("invalid cache version event", exc_info=True) + + +distributed_cache_version_bus = DistributedCacheVersionBus() diff --git a/app/services/memory/__init__.py b/app/services/memory/__init__.py index 33cdc22..273f6f3 100644 --- a/app/services/memory/__init__.py +++ b/app/services/memory/__init__.py @@ -1,14 +1,70 @@ -from app.services.memory.models import MemoryHeader, MemoryMeta, MemoryType, MemoryDocument -from app.services.memory.storage import MemoryStorage -from app.services.memory.scanner import MemoryScanner -from app.services.memory.extractor import MemoryExtractor -from app.services.memory.session_memory import SessionMemoryManager -from app.services.memory.session_store import SessionMemoryStore, SessionMemoryInfo -from app.services.memory.consolidator import MemoryConsolidator +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any + +from app.services.memory.models import ( + CognitiveMemoryType, + MemoryDocument, + MemoryHeader, + MemoryMeta, + MemoryScope, + MemoryType, +) + + +_LAZY_EXPORTS = { + "MemoryStorage": ("app.services.memory.storage", "MemoryStorage"), + "MemoryScanner": ("app.services.memory.scanner", "MemoryScanner"), + "MemoryExtractor": ("app.services.memory.extractor", "MemoryExtractor"), + "SessionMemoryManager": ("app.services.memory.session_memory", "SessionMemoryManager"), + "SessionMemoryStore": ("app.services.memory.session_store", "SessionMemoryStore"), + "SessionMemoryInfo": ("app.services.memory.session_store", "SessionMemoryInfo"), + "MemoryConsolidator": ("app.services.memory.consolidator", "MemoryConsolidator"), + "SemanticMemoryStore": ("app.services.memory.semantic_memory", "SemanticMemoryStore"), + "SemanticMemoryRecord": ("app.services.memory.semantic_memory", "SemanticMemoryRecord"), + "ProceduralMemoryCatalog": ("app.services.memory.procedural_memory", "ProceduralMemoryCatalog"), + "ProceduralMemoryRecord": ("app.services.memory.procedural_memory", "ProceduralMemoryRecord"), +} + + +def __getattr__(name: str) -> Any: + target = _LAZY_EXPORTS.get(name) + if target is None: + raise AttributeError(name) + module_name, attribute = target + value = getattr(import_module(module_name), attribute) + globals()[name] = value + return value + + +if TYPE_CHECKING: + from app.services.memory.consolidator import MemoryConsolidator + from app.services.memory.extractor import MemoryExtractor + from app.services.memory.scanner import MemoryScanner + from app.services.memory.session_memory import SessionMemoryManager + from app.services.memory.session_store import SessionMemoryInfo, SessionMemoryStore + from app.services.memory.storage import MemoryStorage + from app.services.memory.semantic_memory import SemanticMemoryRecord, SemanticMemoryStore + from app.services.memory.procedural_memory import ProceduralMemoryCatalog, ProceduralMemoryRecord + __all__ = [ - "MemoryHeader", "MemoryMeta", "MemoryType", "MemoryDocument", - "MemoryStorage", "MemoryScanner", "MemoryExtractor", - "SessionMemoryManager", "SessionMemoryStore", "SessionMemoryInfo", + "CognitiveMemoryType", "MemoryConsolidator", + "MemoryDocument", + "MemoryExtractor", + "MemoryHeader", + "MemoryMeta", + "MemoryScanner", + "MemoryScope", + "MemoryStorage", + "MemoryType", + "SessionMemoryInfo", + "SessionMemoryManager", + "SessionMemoryStore", + "SemanticMemoryRecord", + "SemanticMemoryStore", + "ProceduralMemoryCatalog", + "ProceduralMemoryRecord", ] diff --git a/app/services/memory/models.py b/app/services/memory/models.py index b299e15..4f363b5 100644 --- a/app/services/memory/models.py +++ b/app/services/memory/models.py @@ -9,7 +9,7 @@ class MemoryType(str, Enum): - """Four strictly-defined memory types matching the architecture document.""" + """Legacy content category retained for API and file compatibility.""" USER = "user" FEEDBACK = "feedback" @@ -25,6 +25,23 @@ class MemoryType(str, Enum): } +class CognitiveMemoryType(str, Enum): + WORKING = "working" + EPISODIC = "episodic" + SEMANTIC = "semantic" + PROCEDURAL = "procedural" + + +class MemoryScope(str, Enum): + REQUEST = "request" + SESSION = "session" + USER = "user" + TEAM = "team" + AGENT = "agent" + TENANT = "tenant" + GLOBAL = "global" + + @dataclass class MemoryMeta: """YAML frontmatter of a memory file.""" @@ -32,6 +49,10 @@ class MemoryMeta: name: str description: str type: MemoryType + memory_type: CognitiveMemoryType = CognitiveMemoryType.SEMANTIC + scope: MemoryScope = MemoryScope.USER + source: str = "manual" + version: int = 1 created_at: str = "" updated_at: str = "" @@ -40,6 +61,10 @@ def to_frontmatter(self) -> str: lines.append(f'name: {self.name}') lines.append(f'description: {self.description}') lines.append(f'type: {self.type.value}') + lines.append(f'memory_type: {self.memory_type.value}') + lines.append(f'scope: {self.scope.value}') + lines.append(f'source: {self.source}') + lines.append(f'version: {max(1, self.version)}') if self.created_at: lines.append(f'created_at: {self.created_at}') if self.updated_at: @@ -75,10 +100,26 @@ def from_frontmatter(cls, text: str, filename: str = "") -> Optional["MemoryMeta mem_type = MemoryType(raw_type) except ValueError: mem_type = MemoryType.REFERENCE + try: + cognitive_type = CognitiveMemoryType(fields.get("memory_type", "semantic")) + except ValueError: + cognitive_type = CognitiveMemoryType.SEMANTIC + try: + scope = MemoryScope(fields.get("scope", "user")) + except ValueError: + scope = MemoryScope.USER + try: + version = max(1, int(fields.get("version", "1"))) + except ValueError: + version = 1 return cls( name=name, description=desc, type=mem_type, + memory_type=cognitive_type, + scope=scope, + source=fields.get("source", "legacy-file"), + version=version, created_at=fields.get("created_at", ""), updated_at=fields.get("updated_at", ""), ) @@ -93,6 +134,10 @@ class MemoryHeader: mtime: float description: str type: MemoryType + memory_type: CognitiveMemoryType = CognitiveMemoryType.SEMANTIC + scope: MemoryScope = MemoryScope.USER + source: str = "legacy-file" + version: int = 1 name: str = "" created_at: str = "" updated_at: str = "" diff --git a/app/services/memory/procedural_memory.py b/app/services/memory/procedural_memory.py new file mode 100644 index 0000000..fb8d6cd --- /dev/null +++ b/app/services/memory/procedural_memory.py @@ -0,0 +1,198 @@ +from __future__ import annotations + +import hashlib +import json +import re +from dataclasses import asdict, dataclass +from typing import Any, Iterable + +from app.services.memory.models import CognitiveMemoryType, MemoryScope +from app.services.memory.storage import MemoryStorage +from app.services.tool_registry import ToolDefinition, tool_registry + + +@dataclass(frozen=True) +class ProceduralMemoryRecord: + id: str + kind: str + name: str + description: str + source: str + source_version: str + content_hash: str + scope: str + risk_level: str = "" + memory_type: str = CognitiveMemoryType.PROCEDURAL.value + + def to_dict(self) -> dict[str, str]: + return asdict(self) + + +class ProceduralMemoryCatalog: + """Read-through catalog over existing procedural sources of truth.""" + + def __init__(self, user_id: str, memory_storage: MemoryStorage) -> None: + self._user_id = user_id + self._memory_storage = memory_storage + + async def list_records(self) -> list[ProceduralMemoryRecord]: + groups = await self._load_groups() + records = [record for group in groups for record in group] + unique = {record.id: record for record in records} + return sorted(unique.values(), key=lambda record: (record.kind, record.name.lower())) + + async def search(self, query: str, limit: int = 8) -> list[ProceduralMemoryRecord]: + records = await self.list_records() + terms = _terms(query) + scored: list[tuple[int, ProceduralMemoryRecord]] = [] + for record in records: + haystack = _terms(f"{record.kind} {record.name} {record.description}") + score = len(terms & haystack) + if score or not terms: + scored.append((score, record)) + scored.sort(key=lambda item: (item[0], item[1].kind, item[1].name), reverse=True) + return [record for _, record in scored[:limit]] + + async def _load_groups(self) -> list[list[ProceduralMemoryRecord]]: + async def safe(loader) -> list[ProceduralMemoryRecord]: + try: + return await loader() + except Exception: + return [] + + return [ + await safe(self._skills), + await safe(self._tools), + await safe(self._dag_templates), + await safe(self._agent_routes), + await safe(self._sops), + await safe(self._tool_policies), + ] + + async def _skills(self) -> list[ProceduralMemoryRecord]: + from app.services.tools.skill_tools import skill_list_handler + + response = await skill_list_handler("all") + skills = (response.get("result") or {}).get("skills") if response.get("success") else [] + return [ + _record( + "skill", + str(item.get("name") or "unknown"), + str(item.get("description") or ""), + f"skill:{item.get('source', 'unknown')}:{item.get('path', '')}", + str(item.get("version") or "1"), + MemoryScope.USER.value if item.get("source") == "user" else MemoryScope.TEAM.value, + json.dumps(item, ensure_ascii=False, sort_keys=True, default=str), + ) + for item in (skills or []) + ] + + async def _tools(self) -> list[ProceduralMemoryRecord]: + tools = tool_registry.list_all() + if not tools: + from app.services.tools.definitions import BUILTIN_TOOLS + + tools = BUILTIN_TOOLS + return records_from_tools(tools) + + async def _dag_templates(self) -> list[ProceduralMemoryRecord]: + from app.services.template_engine import template_engine + + return [ + _record( + "dag-template", str(item["name"]), str(item.get("category") or ""), + f"dag-template:{item['id']}", "1", MemoryScope.GLOBAL.value, + json.dumps(item, ensure_ascii=False, sort_keys=True, default=str), + ) + for item in await template_engine.list_templates() + ] + + async def _agent_routes(self) -> list[ProceduralMemoryRecord]: + from app.services.agent_route_service import agent_route_service + + routes = await agent_route_service.list_routes(self._user_id) + return [ + _record( + "agent-route", str(item["name"]), str(item.get("description") or ""), + f"agent-route:{item['id']}", str(item.get("updated_at") or "1"), + MemoryScope.USER.value, + json.dumps(item, ensure_ascii=False, sort_keys=True, default=str), + ) + for item in routes + ] + + async def _sops(self) -> list[ProceduralMemoryRecord]: + headers = await self._memory_storage.list_headers() + return [ + _record( + "sop", header.name or header.filename, header.description, + f"memory-file:{header.filename}", str(header.version), header.scope.value, + f"{header.path}|{header.mtime}|{header.version}", + ) + for header in headers + if header.memory_type == CognitiveMemoryType.PROCEDURAL + ] + + async def _tool_policies(self) -> list[ProceduralMemoryRecord]: + from app.db.session import afetch_all + + rows = await afetch_all( + "SELECT id,agent_id,tool_pattern,path_pattern,behavior,source,priority,enabled " + "FROM tool_permission_rules WHERE enabled=1 ORDER BY priority DESC" + ) + return [ + _record( + "tool-policy", + f"{row.get('agent_id', '*')}:{row.get('tool_pattern', '*')}", + f"{row.get('behavior', 'ask')} path={row.get('path_pattern', '*')}", + f"tool-policy:{row['id']}", str(row.get("priority", 0)), + MemoryScope.TENANT.value, + json.dumps(row, ensure_ascii=False, sort_keys=True, default=str), + ) + for row in rows + ] + + +def records_from_tools(tools: Iterable[ToolDefinition]) -> list[ProceduralMemoryRecord]: + return [ + _record( + "tool", tool.name, tool.description, f"tool-registry:{tool.name}", "1", + MemoryScope.GLOBAL.value, json.dumps(tool.to_dict(), ensure_ascii=False, sort_keys=True), + risk_level=tool.risk_level, + ) + for tool in tools + ] + + +def _record( + kind: str, + name: str, + description: str, + source: str, + version: str, + scope: str, + fingerprint: str, + *, + risk_level: str = "", +) -> ProceduralMemoryRecord: + content_hash = hashlib.sha256(fingerprint.encode("utf-8")).hexdigest()[:24] + record_id = hashlib.sha256(f"{kind}|{source}".encode("utf-8")).hexdigest()[:24] + return ProceduralMemoryRecord( + id=record_id, + kind=kind, + name=name, + description=description[:500], + source=source, + source_version=version, + content_hash=content_hash, + scope=scope, + risk_level=risk_level, + ) + + +def _terms(text: str) -> set[str]: + lowered = text.lower() + latin = set(re.findall(r"[a-z0-9_-]{2,}", lowered)) + cjk_text = "".join(re.findall(r"[\u3400-\u9fff]", lowered)) + cjk = {cjk_text[index:index + 2] for index in range(max(0, len(cjk_text) - 1))} + return latin | cjk diff --git a/app/services/memory/scanner.py b/app/services/memory/scanner.py index 7f11f6b..515aa42 100644 --- a/app/services/memory/scanner.py +++ b/app/services/memory/scanner.py @@ -4,7 +4,7 @@ from datetime import datetime, timedelta from typing import Optional -from app.services.memory.models import MemoryHeader, MemoryType +from app.services.memory.models import CognitiveMemoryType, MemoryHeader, MemoryType from app.services.memory.storage import MemoryStorage @@ -27,6 +27,14 @@ async def filter_by_type(self, type_: MemoryType, max_files: int = 200) -> list[ """Return only memories of a specific type.""" return [h for h in await self.scan(max_files=max_files) if h.type == type_] + async def filter_by_memory_type( + self, memory_type: CognitiveMemoryType, max_files: int = 200, + ) -> list[MemoryHeader]: + return [ + header for header in await self.scan(max_files=max_files) + if header.memory_type == memory_type + ] + async def format_manifest(self, headers: Optional[list[MemoryHeader]] = None) -> str: """Format scan results as a text manifest (formatMemoryManifest equivalent).""" if headers is None: @@ -39,7 +47,8 @@ async def format_manifest(self, headers: Optional[list[MemoryHeader]] = None) -> freshness = self.freshness_text(h.mtime) lines.append( f" - {h.filename} | {h.name} | type={h.type.value} | " - f"{h.description[:50]}{freshness}" + f"memory_type={h.memory_type.value} | scope={h.scope.value} | " + f"v{h.version} | {h.description[:50]}{freshness}" ) return "\n".join(lines) diff --git a/app/services/memory/semantic_memory.py b/app/services/memory/semantic_memory.py new file mode 100644 index 0000000..9a86028 --- /dev/null +++ b/app/services/memory/semantic_memory.py @@ -0,0 +1,187 @@ +from __future__ import annotations + +import asyncio +import hashlib +import re +from dataclasses import asdict, dataclass +from datetime import UTC, datetime +from pathlib import Path +from typing import Any + +from app.services.memory.models import CognitiveMemoryType, MemoryScope +from app.utils.async_file import aexists, aread_json, awrite_json, amkdir + + +@dataclass(frozen=True) +class SemanticCandidate: + key: str + value: str + category: str + confidence: float + + +@dataclass +class SemanticMemoryRecord: + id: str + key: str + value: str + category: str + confidence: float + source: str + source_session_id: str + source_event_id: str + version: int + status: str + created_at: str + updated_at: str + memory_type: str = CognitiveMemoryType.SEMANTIC.value + scope: str = MemoryScope.USER.value + expires_at: str = "" + superseded_by: str = "" + + +_CATEGORY_PATTERNS: tuple[tuple[str, re.Pattern[str]], ...] = ( + ("preference", re.compile(r"(?:用户|团队)?(?:偏好|习惯|倾向)[::]?\s*(.+)", re.I)), + ("decision", re.compile(r"(?:关键)?(?:决定|决策|结论)[::]?\s*(.+)", re.I)), + ("constraint", re.compile(r"(?:约束|必须|禁止|不得|需要遵守)[::]?\s*(.+)", re.I)), + ("fact", re.compile(r"(?:事实|已确认|确认信息)[::]?\s*(.+)", re.I)), +) + + +def extract_semantic_candidates(summary: str) -> list[SemanticCandidate]: + """Extract only explicit durable-memory signals from an episodic summary.""" + candidates: list[SemanticCandidate] = [] + seen: set[str] = set() + for sentence in re.split(r"[。!?!?\n]+", summary): + sentence = re.sub(r"\s+", " ", sentence).strip(" -\t") + if len(sentence) < 4: + continue + for category, pattern in _CATEGORY_PATTERNS: + match = pattern.search(sentence) + if not match: + continue + value = match.group(1).strip(" ::")[:500] + if len(value) < 2: + continue + label = sentence[:match.start(1)].strip(" ::") or category + normalized_label = re.sub(r"[^\w\u3400-\u9fff]+", "", label.lower())[:48] + if normalized_label in {"用户偏好", "团队偏好", "偏好", category}: + normalized_label += ":" + re.sub(r"[^\w\u3400-\u9fff]+", "", value.lower())[:16] + key = f"{category}:{normalized_label}" + if key in seen: + break + seen.add(key) + candidates.append(SemanticCandidate(key, value, category, 0.78)) + break + return candidates[:12] + + +class SemanticMemoryStore: + """Structured semantic sidecar that preserves the existing memory storage.""" + + _locks: dict[str, asyncio.Lock] = {} + + def __init__(self, user_memory_dir: str | Path) -> None: + self._dir = Path(user_memory_dir).resolve() / "semantic" + self._path = self._dir / "records.json" + self._lock = self._locks.setdefault(str(self._path), asyncio.Lock()) + + async def list_records(self, *, active_only: bool = True) -> list[SemanticMemoryRecord]: + raw = await self._read_records() + records = [SemanticMemoryRecord(**item) for item in raw] + if active_only: + records = [record for record in records if record.status == "active"] + return sorted(records, key=lambda record: record.updated_at, reverse=True) + + async def upsert_candidates( + self, + candidates: list[SemanticCandidate], + *, + source: str, + source_session_id: str, + source_event_id: str, + ) -> list[SemanticMemoryRecord]: + if not candidates: + return [] + async with self._lock: + raw = await self._read_records() + records = [SemanticMemoryRecord(**item) for item in raw] + changed: list[SemanticMemoryRecord] = [] + now = datetime.now(UTC).isoformat() + + for candidate in candidates: + current = next( + (record for record in records if record.key == candidate.key and record.status == "active"), + None, + ) + if current and _normalize(current.value) == _normalize(candidate.value): + current.confidence = max(current.confidence, candidate.confidence) + current.updated_at = now + current.version += 1 + current.source = source + current.source_session_id = source_session_id + current.source_event_id = source_event_id + changed.append(current) + continue + + version = (current.version + 1) if current else 1 + record_id = hashlib.sha256( + f"{candidate.key}|{candidate.value}|{version}".encode("utf-8") + ).hexdigest()[:24] + new_record = SemanticMemoryRecord( + id=record_id, + key=candidate.key, + value=candidate.value, + category=candidate.category, + confidence=candidate.confidence, + source=source, + source_session_id=source_session_id, + source_event_id=source_event_id, + version=version, + status="active", + created_at=now, + updated_at=now, + ) + if current: + current.status = "superseded" + current.superseded_by = record_id + current.updated_at = now + records.append(new_record) + changed.append(new_record) + + await amkdir(self._dir) + await awrite_json(self._path, [asdict(record) for record in records]) + return changed + + async def search(self, query: str, limit: int = 6) -> list[SemanticMemoryRecord]: + records = await self.list_records(active_only=True) + query_terms = _terms(query) + scored: list[tuple[float, SemanticMemoryRecord]] = [] + for record in records: + record_terms = _terms(record.key + " " + record.value) + overlap = len(query_terms & record_terms) + relevance = overlap / max(1, len(query_terms)) if query_terms else 0.0 + if relevance > 0 or record.category == "preference": + scored.append((relevance + record.confidence * 0.2, record)) + scored.sort(key=lambda item: (item[0], item[1].updated_at), reverse=True) + return [record for _, record in scored[:limit]] + + async def _read_records(self) -> list[dict[str, Any]]: + if not await aexists(self._path): + return [] + try: + data = await aread_json(self._path) + return data if isinstance(data, list) else [] + except (OSError, ValueError, TypeError): + return [] + + +def _normalize(text: str) -> str: + return re.sub(r"[^\w\u3400-\u9fff]+", "", text).lower() + + +def _terms(text: str) -> set[str]: + normalized = _normalize(text) + latin = set(re.findall(r"[a-z0-9_]{2,}", text.lower())) + cjk = {normalized[index:index + 2] for index in range(max(0, len(normalized) - 1))} + return latin | cjk diff --git a/app/services/memory/session_memory.py b/app/services/memory/session_memory.py index 4bf15ef..16814ae 100644 --- a/app/services/memory/session_memory.py +++ b/app/services/memory/session_memory.py @@ -2,15 +2,16 @@ import logging import os +import asyncio from datetime import datetime from pathlib import Path from typing import Any, Optional from app.config import MEMORY_DIR -from app.db.session import afetch_all from app.services.adapter_manager import adapter_manager -from app.services.memory.models import MemoryType, sanitize_filename +from app.services.memory.models import CognitiveMemoryType, MemoryScope, MemoryType, sanitize_filename from app.services.memory.storage import MemoryStorage +from app.services.memory.summary_version import SummaryVersion, should_accept_summary from app.utils.async_file import ( aexists, aread_text, @@ -61,6 +62,8 @@ class SessionMemoryManager: State file: .claude/memory/sessions/.session_state.json """ + _summary_locks: dict[str, asyncio.Lock] = {} + def __init__(self, storage: Optional[MemoryStorage] = None) -> None: self._storage = storage or MemoryStorage(MEMORY_DIR) self._sessions_dir = self._storage.base / "sessions" @@ -172,7 +175,12 @@ async def update_global_summary(self) -> str: global_path = self._storage.base / "总体系统记忆文档.md" await awrite_text(global_path, result) - self._state.setdefault("global", {})["updated_at"] = datetime.now().isoformat(timespec="seconds") + global_state = self._state.setdefault("global", {}) + global_state["updated_at"] = datetime.now().isoformat(timespec="seconds") + global_state["memory_type"] = CognitiveMemoryType.SEMANTIC.value + global_state["scope"] = MemoryScope.USER.value + global_state["source"] = "session-summary-aggregation" + global_state["version"] = max(1, int(global_state.get("version", 0)) + 1) await self._save_state() # Invalidate memory context cache so the next agent call picks up fresh data @@ -205,22 +213,66 @@ async def get_global_summary(self) -> str: pass return "" - async def write_session_summary(self, session_id: str, summary: str) -> None: + async def write_session_summary( + self, + session_id: str, + summary: str, + *, + covered_sequence_start: int = 0, + covered_sequence_end: int = 0, + generated_at: float = 0.0, + source_event_id: str = "", + force: bool = True, + ) -> bool: """Write summary for a session and refresh cursor timestamp.""" await amkdir(self._sessions_dir) summary_path = self._sessions_dir / f"{sanitize_filename(session_id)}" + lock_key = str(summary_path) + lock = self._summary_locks.setdefault(lock_key, asyncio.Lock()) try: - await awrite_text(summary_path, summary or "") - await self._ensure_state_loaded() - sessions = self._state.setdefault("sessions", {}) - current = sessions.get(session_id, {}) - sessions[session_id] = { - "last_msg_id": current.get("last_msg_id", ""), - "updated_at": datetime.now().isoformat(timespec="seconds"), - } - await self._save_state() + async with lock: + await self._ensure_state_loaded() + sessions = self._state.setdefault("sessions", {}) + current = sessions.get(session_id, {}) + current_version = SummaryVersion( + covered_sequence_start=int(current.get("covered_sequence_start", 0)), + covered_sequence_end=int(current.get("covered_sequence_end", 0)), + generated_at=float(current.get("summary_generated_at", 0.0)), + source_event_id=str(current.get("source_event_id", "")), + ) + incoming_version = SummaryVersion( + covered_sequence_start=covered_sequence_start, + covered_sequence_end=covered_sequence_end, + generated_at=generated_at, + source_event_id=source_event_id, + ) + if not force and not should_accept_summary(current_version, incoming_version): + logger.info( + "stale session summary rejected session=%s incoming_end=%d current_end=%d event=%s", + session_id, covered_sequence_end, current_version.covered_sequence_end, source_event_id, + ) + return False + await awrite_text(summary_path, summary or "") + stored_sequence_start = covered_sequence_start or current_version.covered_sequence_start + stored_sequence_end = covered_sequence_end or current_version.covered_sequence_end + sessions[session_id] = { + **current, + "last_msg_id": current.get("last_msg_id", ""), + "updated_at": datetime.now().isoformat(timespec="seconds"), + "memory_type": CognitiveMemoryType.EPISODIC.value, + "scope": MemoryScope.SESSION.value, + "source": "session-summary", + "version": max(1, int(current.get("version", 0)) + 1), + "covered_sequence_start": stored_sequence_start, + "covered_sequence_end": stored_sequence_end, + "summary_generated_at": generated_at or current_version.generated_at, + "source_event_id": source_event_id or current_version.source_event_id, + } + await self._save_state() + return True except OSError as exc: logger.error("failed to write session summary for session=%s: %s", session_id, exc) + return False async def list_session_summaries(self) -> list[dict[str, Any]]: """List all session summaries with metadata.""" @@ -251,6 +303,8 @@ async def list_session_summaries(self) -> list[dict[str, Any]]: # ── internal helpers ──────────────────────────────────────────── async def _get_session_messages(self, session_id: str) -> list[dict[str, Any]]: + from app.db.session import afetch_all + try: return await afetch_all( "SELECT id, sender, content, type, created_at " @@ -330,6 +384,7 @@ async def _call_llm_raw(self, prompt: str) -> str | None: async def _list_summarization_models(self) -> list[dict[str, str]]: """Return all available models for summarization in priority order.""" + from app.db.session import afetch_all from app.services.secret_service import decrypt_secret candidates: list[dict[str, str]] = [] @@ -402,9 +457,15 @@ async def _save_state(self) -> None: async def _update_session_cursor(self, session_id: str, message_id: str) -> None: await self._ensure_state_loaded() sessions = self._state.setdefault("sessions", {}) + current = sessions.get(session_id, {}) sessions[session_id] = { + **current, "last_msg_id": message_id, "updated_at": datetime.now().isoformat(timespec="seconds"), + "memory_type": CognitiveMemoryType.EPISODIC.value, + "scope": MemoryScope.SESSION.value, + "source": "session-summary", + "version": max(1, int(current.get("version", 0)) + 1), } await self._save_state() diff --git a/app/services/memory/session_store.py b/app/services/memory/session_store.py index 2810bf0..c385c96 100644 --- a/app/services/memory/session_store.py +++ b/app/services/memory/session_store.py @@ -24,6 +24,8 @@ from pathlib import Path from typing import Any, Optional +from app.services.memory.models import CognitiveMemoryType, MemoryScope + from app.utils.async_file import ( aexists, aread_text, @@ -59,6 +61,10 @@ class SessionMemoryInfo: conversation_size_chars: int = 0 turn_count: int = 0 is_active: bool = True + memory_type: str = CognitiveMemoryType.EPISODIC.value + scope: str = MemoryScope.SESSION.value + source: str = "conversation" + version: int = 1 class SessionMemoryStore: @@ -110,6 +116,10 @@ async def ensure_session(self, session_id: str, session_name: str = "") -> None: "updated_at": now, "turn_count": 0, "is_active": True, + "memory_type": CognitiveMemoryType.EPISODIC.value, + "scope": MemoryScope.SESSION.value, + "source": "conversation", + "version": 1, } await awrite_text(meta_path, json.dumps(meta, ensure_ascii=False, indent=2)) @@ -178,6 +188,10 @@ async def append_turn( # Update metadata meta["turn_count"] = turn_num meta["updated_at"] = now_str + meta["memory_type"] = CognitiveMemoryType.EPISODIC.value + meta["scope"] = MemoryScope.SESSION.value + meta["source"] = "conversation" + meta["version"] = max(1, int(meta.get("version", 1))) + 1 await self._save_meta(session_id, meta) # Check if consolidation is needed @@ -251,6 +265,10 @@ async def get_session_info(self, session_id: str) -> Optional[SessionMemoryInfo] conversation_size_chars=conv_size, turn_count=meta.get("turn_count", 0), is_active=meta.get("is_active", True), + memory_type=meta.get("memory_type", CognitiveMemoryType.EPISODIC.value), + scope=meta.get("scope", MemoryScope.SESSION.value), + source=meta.get("source", "legacy-conversation"), + version=max(1, int(meta.get("version", 1))), ) async def list_sessions(self) -> list[SessionMemoryInfo]: @@ -329,6 +347,7 @@ async def _auto_consolidate(self, session_id: str) -> None: if consolidated: conv_path = self._sessions_dir / session_id / "conversation.md" await awrite_text(conv_path, consolidated) + await self._bump_memory_version(session_id) logger.info( "consolidated session=%s: %d → %d chars", session_id, len(content), len(consolidated), @@ -464,6 +483,7 @@ async def trigger_llm_consolidation(self, session_id: str) -> str | None: if result and result.strip(): conv_path = self._sessions_dir / session_id / "conversation.md" await awrite_text(conv_path, result.strip()) + await self._bump_memory_version(session_id) return result.strip() except Exception as exc: @@ -478,7 +498,16 @@ def _build_header(session_id: str, session_name: str = "") -> str: """Build the header for a new conversation.md file.""" name = session_name or session_id now = datetime.now().isoformat(timespec="seconds") - return f"""# 会话记忆: {name} + return f"""--- +memory_type: {CognitiveMemoryType.EPISODIC.value} +scope: {MemoryScope.SESSION.value} +source: conversation +version: 1 +session_id: {session_id} +created_at: {now} +--- + +# 会话记忆: {name} > Session ID: `{session_id}` > 创建时间: {now} @@ -507,6 +536,17 @@ async def _save_meta(self, session_id: str, meta: dict[str, Any]) -> None: meta_path = session_dir / "session.json" await awrite_text(meta_path, json.dumps(meta, ensure_ascii=False, indent=2)) + async def _bump_memory_version(self, session_id: str) -> None: + meta = await self._load_meta(session_id) + if not meta: + return + meta["memory_type"] = CognitiveMemoryType.EPISODIC.value + meta["scope"] = MemoryScope.SESSION.value + meta["source"] = "conversation" + meta["version"] = max(1, int(meta.get("version", 1))) + 1 + meta["updated_at"] = datetime.now().isoformat(timespec="seconds") + await self._save_meta(session_id, meta) + @staticmethod def _truncate_message(msg: str, max_chars: int) -> str: """Truncate a single message to max_chars, preserving structure.""" diff --git a/app/services/memory/storage.py b/app/services/memory/storage.py index f6cdfd8..4271610 100644 --- a/app/services/memory/storage.py +++ b/app/services/memory/storage.py @@ -9,9 +9,11 @@ from app.services.memory.models import ( MEMORY_TYPE_DESCRIPTIONS, MemoryDocument, + CognitiveMemoryType, MemoryHeader, MemoryMeta, MemoryType, + MemoryScope, sanitize_filename, ) from app.utils.async_file import ( @@ -68,6 +70,11 @@ async def save( type_: MemoryType, body: str = "", filename: str | None = None, + *, + memory_type: CognitiveMemoryType | None = None, + scope: MemoryScope | None = None, + source: str | None = None, + version: int | None = None, ) -> MemoryDocument: """Create or overwrite a memory file.""" await self._ensure_dir() @@ -87,6 +94,10 @@ async def save( name=name, description=description, type=type_, + memory_type=memory_type or (existing.meta.memory_type if existing else CognitiveMemoryType.SEMANTIC), + scope=scope or (existing.meta.scope if existing else MemoryScope.USER), + source=source or (existing.meta.source if existing else "manual"), + version=max(1, version if version is not None else ((existing.meta.version + 1) if existing else 1)), created_at=created_at, updated_at=now, ) @@ -310,6 +321,10 @@ async def list_headers(self, max_files: int = 200) -> list[MemoryHeader]: mtime=mtime, description=meta.description, type=meta.type, + memory_type=meta.memory_type, + scope=meta.scope, + source=meta.source, + version=meta.version, name=meta.name, created_at=meta.created_at, updated_at=meta.updated_at, @@ -349,16 +364,19 @@ async def _write_index(self, headers: list[MemoryHeader]) -> None: "## 概述", "", "本目录存储跨会话的持久化记忆。每一条记忆是一个 `.md` 文件,", - "包含 YAML frontmatter(name, description, type)。", + "包含 YAML frontmatter(name, description, type, memory_type, scope, source, version)。", "", - "| 文件名 | 名称 | 类型 | 描述 | 更新于 |", - "|--------|------|------|------|--------|", + "| 文件名 | 名称 | 内容类别 | 认知类型 | 作用域 | 版本 | 描述 | 更新于 |", + "|--------|------|----------|----------|--------|------|------|--------|", ] for h in headers: updated = h.updated_at or datetime.fromtimestamp(h.mtime).isoformat(timespec="seconds") if h.mtime else "-" # Truncate description for table desc = h.description[:60] + "…" if len(h.description) > 60 else h.description - lines.append(f"| {h.filename} | {h.name} | {h.type.value} | {desc} | {updated} |") + lines.append( + f"| {h.filename} | {h.name} | {h.type.value} | {h.memory_type.value} | " + f"{h.scope.value} | {h.version} | {desc} | {updated} |" + ) lines.append("") lines.append("---") diff --git a/app/services/memory/summary_version.py b/app/services/memory/summary_version.py new file mode 100644 index 0000000..3c64f06 --- /dev/null +++ b/app/services/memory/summary_version.py @@ -0,0 +1,25 @@ +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class SummaryVersion: + covered_sequence_start: int = 0 + covered_sequence_end: int = 0 + generated_at: float = 0.0 + source_event_id: str = "" + + +def should_accept_summary(current: SummaryVersion, incoming: SummaryVersion) -> bool: + if incoming.source_event_id and incoming.source_event_id == current.source_event_id: + return False + if ( + incoming.covered_sequence_end > 0 + and current.covered_sequence_end > 0 + and incoming.covered_sequence_end <= current.covered_sequence_end + ): + return False + if incoming.generated_at > 0 and current.generated_at > incoming.generated_at: + return False + return True diff --git a/app/services/memory/test_cognitive_memory_models.py b/app/services/memory/test_cognitive_memory_models.py new file mode 100644 index 0000000..79d1e31 --- /dev/null +++ b/app/services/memory/test_cognitive_memory_models.py @@ -0,0 +1,70 @@ +from __future__ import annotations + +import asyncio + +from app.services.memory.models import ( + CognitiveMemoryType, + MemoryDocument, + MemoryMeta, + MemoryScope, + MemoryType, +) +from app.services.memory.session_store import SessionMemoryStore +from app.services.memory.storage import MemoryStorage + + +def test_legacy_frontmatter_gets_safe_cognitive_defaults() -> None: + document = MemoryDocument.parse( + "---\nname: preference\ndescription: reply style\ntype: user\n---\n\nconcise" + ) + assert document.meta.memory_type == CognitiveMemoryType.SEMANTIC + assert document.meta.scope == MemoryScope.USER + assert document.meta.source == "legacy-file" + assert document.meta.version == 1 + + +def test_cognitive_metadata_round_trip() -> None: + meta = MemoryMeta( + name="deploy-sop", + description="production release procedure", + type=MemoryType.PROJECT, + memory_type=CognitiveMemoryType.PROCEDURAL, + scope=MemoryScope.TEAM, + source="workflow:42", + version=3, + ) + parsed = MemoryDocument.parse(MemoryDocument(meta, "steps").to_markdown()) + assert parsed.meta == meta + + +def test_storage_preserves_metadata_and_increments_version(tmp_path) -> None: + async def run() -> None: + storage = MemoryStorage(tmp_path / "memory") + first = await storage.save( + "preference", "reply style", MemoryType.USER, "concise", + memory_type=CognitiveMemoryType.SEMANTIC, + scope=MemoryScope.USER, + source="manual", + ) + second = await storage.save( + "preference", "reply style", MemoryType.USER, "very concise", + filename=first.file_path.split("\\")[-1], + ) + assert second.meta.memory_type == CognitiveMemoryType.SEMANTIC + assert second.meta.version == 2 + + asyncio.run(run()) + + +def test_session_store_classifies_conversation_as_episodic(tmp_path) -> None: + async def run() -> None: + store = SessionMemoryStore(tmp_path / "user") + await store.append_turn("s1", "question", "answer") + info = await store.get_session_info("s1") + assert info is not None + assert info.memory_type == CognitiveMemoryType.EPISODIC.value + assert info.scope == MemoryScope.SESSION.value + assert info.source == "conversation" + assert info.version == 2 + + asyncio.run(run()) diff --git a/app/services/memory/test_nats_memory_pipeline_integration.py b/app/services/memory/test_nats_memory_pipeline_integration.py new file mode 100644 index 0000000..3b36222 --- /dev/null +++ b/app/services/memory/test_nats_memory_pipeline_integration.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import asyncio +import json +import os +import time +import uuid +from datetime import UTC, datetime + +import pytest + + +def test_nats_rust_summary_online_contract() -> None: + if os.getenv("AGENTHUB_RUN_NATS_INTEGRATION") != "1": + pytest.skip("set AGENTHUB_RUN_NATS_INTEGRATION=1 with NATS/Rust/summarization services running") + + async def run() -> None: + import nats + + nc = await nats.connect(os.getenv("NATS_URL", "nats://127.0.0.1:4222")) + session_id = f"integration-{uuid.uuid4().hex[:10]}" + event_id = f"compact-{uuid.uuid4().hex}" + summary_future: asyncio.Future[dict] = asyncio.get_running_loop().create_future() + + async def on_summary(message) -> None: + envelope = json.loads(message.data) + if envelope.get("session_id") == session_id and not summary_future.done(): + summary_future.set_result(envelope) + + subscription = await nc.subscribe("agenthub.session.summary", cb=on_summary) + messages = [ + { + "sequence": index, + "role": "user" if index % 2 else "assistant", + "content": f"integration memory message {index}", + "token_count": 8, + "timestamp": int(time.time()), + } + for index in range(1, 41) + ] + request = { + "event_id": event_id, + "event_type": "memory.compact.requested", + "event_version": 1, + "occurred_at": datetime.now(UTC).isoformat(), + "trace_id": uuid.uuid4().hex, + "tenant_id": "integration-user", + "session_id": session_id, + "message_id": None, + "actor_id": "integration-test", + "producer": {"service": "integration-test", "instance": "pytest", "region": None}, + "routing": {"channel": "memory", "partition_key": session_id, "priority": "normal"}, + "payload": {"messages": messages}, + } + await nc.publish("agenthub.memory.compact.requested", json.dumps(request).encode("utf-8")) + await nc.flush() + + envelope = await asyncio.wait_for(summary_future, timeout=45) + payload = envelope["payload"] + assert envelope["event_type"] == "session.summary.generated" + assert payload["source"] == "memory-compact" + assert payload["tokens_before"] > 0 + assert payload["tokens_after"] > 0 + assert payload["covered_sequence_start"] == 1 + assert payload["covered_sequence_end"] == 30 + assert payload["summary"] + + await subscription.unsubscribe() + await nc.drain() + + asyncio.run(run()) diff --git a/app/services/memory/test_procedural_memory.py b/app/services/memory/test_procedural_memory.py new file mode 100644 index 0000000..5a3c744 --- /dev/null +++ b/app/services/memory/test_procedural_memory.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from app.services.memory.models import CognitiveMemoryType, MemoryScope +from app.services.memory.procedural_memory import records_from_tools +from app.services.tool_registry import ToolDefinition, ToolExample, ToolParameter + + +def test_tool_definition_projects_to_procedural_memory() -> None: + tool = ToolDefinition( + name="deploy_preview", + description="Create a deployment preview", + category="integration", + parameters=[ToolParameter("environment", "string", True, "Target environment")], + return_type="object", + examples=[ToolExample("preview it", {"environment": "staging"})], + risk_level="L2", + ) + record = records_from_tools([tool])[0] + assert record.memory_type == CognitiveMemoryType.PROCEDURAL.value + assert record.scope == MemoryScope.GLOBAL.value + assert record.kind == "tool" + assert record.source == "tool-registry:deploy_preview" + assert record.risk_level == "L2" + assert len(record.content_hash) == 24 diff --git a/app/services/memory/test_semantic_memory.py b/app/services/memory/test_semantic_memory.py new file mode 100644 index 0000000..1b039d3 --- /dev/null +++ b/app/services/memory/test_semantic_memory.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +import asyncio + +from app.services.memory.semantic_memory import ( + SemanticCandidate, + SemanticMemoryStore, + extract_semantic_candidates, +) + + +def test_extract_semantic_candidates_requires_explicit_signal() -> None: + candidates = extract_semantic_candidates( + "用户询问了部署问题。用户偏好:回答保持简洁。约束:不得自动发布生产环境。" + ) + assert [candidate.category for candidate in candidates] == ["preference", "constraint"] + assert all(candidate.confidence >= 0.7 for candidate in candidates) + + +def test_semantic_store_supersedes_conflicting_value(tmp_path) -> None: + async def run() -> None: + store = SemanticMemoryStore(tmp_path / "user") + await store.upsert_candidates( + [SemanticCandidate("preference:reply-style", "concise", "preference", 0.8)], + source="session-summary", + source_session_id="s1", + source_event_id="e1", + ) + changed = await store.upsert_candidates( + [SemanticCandidate("preference:reply-style", "detailed", "preference", 0.85)], + source="session-summary", + source_session_id="s2", + source_event_id="e2", + ) + all_records = await store.list_records(active_only=False) + active = await store.list_records(active_only=True) + assert len(all_records) == 2 + assert len(active) == 1 + assert active[0].value == "detailed" + assert active[0].version == 2 + assert changed[0].source_session_id == "s2" + old_record = next(record for record in all_records if record.status == "superseded") + assert old_record.superseded_by == active[0].id + + asyncio.run(run()) + + +def test_semantic_search_includes_preferences_and_relevant_facts(tmp_path) -> None: + async def run() -> None: + store = SemanticMemoryStore(tmp_path / "user") + await store.upsert_candidates( + [ + SemanticCandidate("preference:language", "使用中文回复", "preference", 0.9), + SemanticCandidate("fact:deploy", "生产环境运行在 Kubernetes", "fact", 0.85), + ], + source="summary", + source_session_id="s1", + source_event_id="e1", + ) + results = await store.search("如何部署 Kubernetes") + assert any(record.category == "fact" for record in results) + assert any(record.category == "preference" for record in results) + + asyncio.run(run()) diff --git a/app/services/memory/test_summary_version.py b/app/services/memory/test_summary_version.py new file mode 100644 index 0000000..319ca4d --- /dev/null +++ b/app/services/memory/test_summary_version.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +import asyncio + +from app.services.memory.session_memory import SessionMemoryManager +from app.services.memory.storage import MemoryStorage +from app.services.memory.summary_version import SummaryVersion, should_accept_summary + + +def test_summary_version_rejects_duplicate_and_older_coverage() -> None: + current = SummaryVersion(1, 40, 200.0, "event-2") + assert not should_accept_summary(current, SummaryVersion(1, 50, 210.0, "event-2")) + assert not should_accept_summary(current, SummaryVersion(1, 30, 210.0, "event-3")) + assert not should_accept_summary(current, SummaryVersion(1, 50, 190.0, "event-4")) + + +def test_summary_version_accepts_newer_coverage() -> None: + current = SummaryVersion(1, 40, 200.0, "event-2") + assert should_accept_summary(current, SummaryVersion(20, 60, 210.0, "event-3")) + + +def test_session_summary_store_rejects_stale_write(tmp_path) -> None: + async def run() -> None: + manager = SessionMemoryManager(MemoryStorage(tmp_path / "memory")) + assert await manager.write_session_summary( + "s1", "new summary", covered_sequence_end=40, + generated_at=200.0, source_event_id="event-2", force=False, + ) + assert not await manager.write_session_summary( + "s1", "stale summary", covered_sequence_end=30, + generated_at=210.0, source_event_id="event-3", force=False, + ) + assert await manager.get_session_summary("s1") == "new summary" + assert await manager.write_session_summary( + "s1", "newest summary", covered_sequence_start=20, covered_sequence_end=60, + generated_at=220.0, source_event_id="event-4", force=False, + ) + assert await manager.get_session_summary("s1") == "newest summary" + + asyncio.run(run()) diff --git a/app/services/memory_context.py b/app/services/memory_context.py new file mode 100644 index 0000000..9aefb38 --- /dev/null +++ b/app/services/memory_context.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import re +from dataclasses import dataclass + +from app.services.token_budget import count_tokens, truncate_to_tokens + + +@dataclass(frozen=True) +class MemoryContextSection: + name: str + text: str + priority: int + memory_type: str = "episodic" + + +def _normalized(text: str) -> str: + return re.sub(r"[^\w\u3400-\u9fff]+", "", text).lower() + + +def _features(text: str) -> set[str]: + normalized = _normalized(text) + if len(normalized) < 12: + return {normalized} if normalized else set() + return {normalized[index:index + 8] for index in range(0, len(normalized) - 7, 4)} + + +def similarity(left: str, right: str) -> float: + left_norm, right_norm = _normalized(left), _normalized(right) + if not left_norm or not right_norm: + return 0.0 + if left_norm in right_norm or right_norm in left_norm: + return min(len(left_norm), len(right_norm)) / max(len(left_norm), len(right_norm)) + left_set, right_set = _features(left), _features(right) + if not left_set or not right_set: + return 0.0 + return len(left_set & right_set) / max(1, min(len(left_set), len(right_set))) + + +def deduplicate_text(text: str, references: list[str], threshold: float = 0.82) -> str: + blocks = [block.strip() for block in re.split(r"\n{2,}", text) if block.strip()] + kept: list[str] = [] + for block in blocks: + if any(similarity(block, reference) >= threshold for reference in references if reference): + continue + if any(similarity(block, previous) >= threshold for previous in kept): + continue + kept.append(block) + return "\n\n".join(kept) + + +def build_memory_context( + sections: list[MemoryContextSection], + *, + exclude_texts: list[str] | None = None, + max_tokens: int = 3_000, + provider: str = "", + model: str = "", + section_budgets: dict[str, int] | None = None, +) -> tuple[str, dict[str, int | bool]]: + ordered = sorted(sections, key=lambda section: section.priority) + references = [text for text in (exclude_texts or []) if text] + before = sum(count_tokens(section.text, provider, model) for section in ordered) + output: list[str] = [] + truncated = False + used_by_type: dict[str, int] = {} + + for section in ordered: + unique = deduplicate_text(section.text, references) + if not unique: + continue + rendered = f"[{section.name}]\n{unique}" + remaining = max_tokens - count_tokens("\n\n".join(output), provider, model) + if section_budgets is not None: + type_remaining = section_budgets.get(section.memory_type, 0) - used_by_type.get(section.memory_type, 0) + remaining = min(remaining, type_remaining) + if remaining <= 8: + truncated = True + continue + rendered, was_truncated = truncate_to_tokens( + rendered, remaining, provider, model, preserve_tail=0.65, + ) + truncated = truncated or was_truncated + output.append(rendered) + used_by_type[section.memory_type] = used_by_type.get(section.memory_type, 0) + count_tokens( + rendered, provider, model, + ) + references.append(unique) + + result = "\n\n".join(output) + after = count_tokens(result, provider, model) + return result, { + "tokens_before": before, + "tokens_after": after, + "truncated": truncated, + } diff --git a/app/services/memory_summary_consumer.py b/app/services/memory_summary_consumer.py new file mode 100644 index 0000000..f6e42c8 --- /dev/null +++ b/app/services/memory_summary_consumer.py @@ -0,0 +1,208 @@ +from __future__ import annotations + +import asyncio +import json +import logging +import os +import re +import uuid +from collections import deque +from datetime import UTC, datetime +from typing import Any + +from app.config import MEMORY_DIR +from app.services.memory.session_memory import SessionMemoryManager +from app.services.memory.semantic_memory import SemanticMemoryStore, extract_semantic_candidates +from app.services.memory.storage import MemoryStorage +from app.services.performance_monitor import monitor +from app.services.token_budget import count_tokens + + +logger = logging.getLogger("agenthub.memory.summary_consumer") + + +class MemorySummaryConsumer: + """Persist semantic summaries emitted by the offline memory pipeline.""" + + def __init__(self) -> None: + self._nc: Any = None + self._subscription: Any = None + self._seen_order: deque[str] = deque(maxlen=2048) + self._seen: set[str] = set() + + async def start(self) -> bool: + if os.getenv("AGENTHUB_MEMORY_EVENTS_ENABLED", "true").lower() not in {"1", "true", "yes"}: + return False + try: + import nats + + self._nc = await asyncio.wait_for( + nats.connect( + os.getenv("NATS_URL", "nats://127.0.0.1:4222"), + name="agenthub-online-memory-consumer", + connect_timeout=2, + max_reconnect_attempts=-1, + reconnect_time_wait=2, + ), + timeout=3, + ) + js = self._nc.jetstream() + try: + self._subscription = await js.subscribe( + "agenthub.session.summary", + durable="agenthub-online-memory-summary-v1", + cb=self._handle_message, + manual_ack=True, + ) + except Exception: + # Startup ordering must not disable future summary delivery when + # the SESSION JetStream is created by another service later. + self._subscription = await self._nc.subscribe( + "agenthub.session.summary", cb=self._handle_core_message, + ) + logger.warning("SESSION stream unavailable; using live NATS summary subscription") + logger.info("memory summary consumer subscribed") + return True + except Exception as exc: + logger.warning("memory summary consumer disabled: %s", exc) + await self.close() + return False + + async def close(self) -> None: + if self._subscription is not None: + try: + await self._subscription.unsubscribe() + except Exception: + pass + self._subscription = None + if self._nc is not None and not self._nc.is_closed: + try: + await self._nc.drain() + except Exception: + await self._nc.close() + self._nc = None + + async def _handle_message(self, message: Any) -> None: + try: + envelope = json.loads(message.data) + await self.consume_envelope(envelope) + await message.ack() + except Exception: + logger.exception("failed to consume session summary event") + await message.nak() + + async def _handle_core_message(self, message: Any) -> None: + try: + await self.consume_envelope(json.loads(message.data)) + except Exception: + logger.exception("failed to consume live session summary event") + + async def consume_envelope(self, envelope: dict[str, Any]) -> bool: + if envelope.get("event_type") != "session.summary.generated": + return False + event_id = str(envelope.get("event_id", "")) + if event_id and event_id in self._seen: + return False + + payload = envelope.get("payload") or {} + session_id = str(payload.get("session_id") or envelope.get("session_id") or "").strip() + summary = str(payload.get("summary") or "").strip() + tenant_id = str(payload.get("tenant_id") or envelope.get("tenant_id") or "local-admin") + user_id = re.sub(r"[^A-Za-z0-9_.@-]", "_", tenant_id)[:128] or "local-admin" + if not session_id or not summary: + return False + + manager = SessionMemoryManager(MemoryStorage(MEMORY_DIR / "users" / user_id)) + accepted = await manager.write_session_summary( + session_id, + summary, + covered_sequence_start=int(payload.get("covered_sequence_start") or 0), + covered_sequence_end=int(payload.get("covered_sequence_end") or 0), + generated_at=float(payload.get("generated_at") or 0.0), + source_event_id=event_id, + force=False, + ) + if not accepted: + return False + semantic_store = SemanticMemoryStore(MEMORY_DIR / "users" / user_id) + await semantic_store.upsert_candidates( + extract_semantic_candidates(summary), + source=str(payload.get("source") or "session-summary"), + source_session_id=session_id, + source_event_id=event_id, + ) + try: + from app.services.agent_service import _invalidate_memory_cache + + _invalidate_memory_cache() + except Exception: + logger.debug("memory cache invalidation unavailable", exc_info=True) + + monitor.record_summary_usage( + int(payload.get("summary_tokens") or 0), + estimated_cost=float(payload.get("estimated_cost") or 0.0), + quality_score=( + float(payload["quality_score"]) + if payload.get("quality_score") is not None else None + ), + ) + if event_id: + if len(self._seen_order) == self._seen_order.maxlen: + oldest = self._seen_order.popleft() + self._seen.discard(oldest) + self._seen_order.append(event_id) + self._seen.add(event_id) + return True + + async def request_compaction(self, session_id: str, user_id: str) -> bool: + """Publish a bounded conversation window for Rust compaction.""" + if self._nc is None or self._nc.is_closed: + return False + from app.db.session import afetch_all + + rows = await afetch_all( + "SELECT id,sender,content,created_at,sequence FROM (" + "SELECT id,sender,content,created_at," + "ROW_NUMBER() OVER (ORDER BY created_at,id) AS sequence " + "FROM messages WHERE session_id=$1 AND type!='system'" + ") ranked ORDER BY sequence DESC LIMIT 60", + session_id, + ) + rows.reverse() + if len(rows) < 20: + return False + messages = [] + for index, row in enumerate(rows, start=1): + sender = str(row.get("sender") or "user").lower() + role = "user" if sender in {"user", user_id.lower()} else "assistant" + content = str(row.get("content") or "") + messages.append({ + "sequence": int(row.get("sequence") or index), + "role": role, + "content": content, + "token_count": count_tokens(content), + "timestamp": None, + }) + event_id = f"compact-{session_id}-{rows[-1].get('id') or uuid.uuid4().hex[:8]}" + envelope = { + "event_id": event_id, + "event_type": "memory.compact.requested", + "event_version": 1, + "occurred_at": datetime.now(UTC).isoformat(), + "trace_id": uuid.uuid4().hex, + "tenant_id": user_id or "local-admin", + "session_id": session_id, + "message_id": str(rows[-1].get("id") or "") or None, + "actor_id": user_id or None, + "producer": {"service": "agenthub-online", "instance": "local", "region": None}, + "routing": {"channel": "memory", "partition_key": session_id, "priority": "normal"}, + "payload": {"messages": messages}, + } + await self._nc.publish( + "agenthub.memory.compact.requested", + json.dumps(envelope, ensure_ascii=False).encode("utf-8"), + ) + return True + + +memory_summary_consumer = MemorySummaryConsumer() diff --git a/app/services/performance_monitor.py b/app/services/performance_monitor.py index dfaf83a..bde95db 100644 --- a/app/services/performance_monitor.py +++ b/app/services/performance_monitor.py @@ -152,6 +152,20 @@ def __init__(self) -> None: self._total_http_retries = 0 self._total_tool_call_loops = 0 self._total_tool_call_iterations = 0 + self._token_economy: dict[str, dict[str, float]] = defaultdict( + lambda: { + "tokens_before": 0, + "tokens_after": 0, + "truncations": 0, + "operations": 0, + } + ) + self._summary_tokens = 0 + self._summary_cost = 0.0 + self._summary_quality_total = 0.0 + self._summary_quality_samples = 0 + self._answer_quality_total = 0.0 + self._answer_quality_samples = 0 # ── LLM Call tracking (sync — hot-path safe) ─────────────────── @@ -248,6 +262,40 @@ def record_tool_call_loop(self, iterations: int) -> None: self._total_tool_call_loops += 1 self._total_tool_call_iterations += iterations + def record_token_compaction( + self, + section: str, + tokens_before: int, + tokens_after: int, + truncated: bool = False, + ) -> None: + with self._lock: + metrics = self._token_economy[section] + metrics["tokens_before"] += max(0, tokens_before) + metrics["tokens_after"] += max(0, tokens_after) + metrics["operations"] += 1 + if truncated: + metrics["truncations"] += 1 + + def record_summary_usage( + self, + tokens: int, + *, + estimated_cost: float = 0.0, + quality_score: float | None = None, + ) -> None: + with self._lock: + self._summary_tokens += max(0, tokens) + self._summary_cost += max(0.0, estimated_cost) + if quality_score is not None: + self._summary_quality_total += max(0.0, min(1.0, quality_score)) + self._summary_quality_samples += 1 + + def record_answer_quality(self, score: float) -> None: + with self._lock: + self._answer_quality_total += max(0.0, min(1.0, score)) + self._answer_quality_samples += 1 + # ── Snapshot API (can be called from sync or async context) ──── def snapshot(self) -> dict[str, Any]: @@ -309,6 +357,40 @@ def snapshot(self) -> dict[str, Any]: uptime = time.time() - self._started_at + token_sections: dict[str, dict[str, float]] = {} + for name, values in sorted(self._token_economy.items()): + before = int(values["tokens_before"]) + after = int(values["tokens_after"]) + token_sections[name] = { + "tokensBefore": before, + "tokensAfter": after, + "tokensSaved": max(0, before - after), + "reductionRate": round((before - after) / before, 4) if before else 0.0, + "truncations": int(values["truncations"]), + "operations": int(values["operations"]), + } + + try: + from app.services.context_summary_cache import context_summary_cache + summary_cache = context_summary_cache.stats() + except Exception: + summary_cache = {} + + token_economy = { + "sections": token_sections, + "summaryCache": summary_cache, + "summaryTokens": self._summary_tokens, + "summaryEstimatedCost": round(self._summary_cost, 6), + "summaryQualityScore": round( + self._summary_quality_total / self._summary_quality_samples, 4, + ) if self._summary_quality_samples else None, + "summaryQualitySamples": self._summary_quality_samples, + "answerQualityScore": round( + self._answer_quality_total / self._answer_quality_samples, 4, + ) if self._answer_quality_samples else None, + "answerQualitySamples": self._answer_quality_samples, + } + return { "uptimeSeconds": round(uptime, 1), "global": { @@ -323,6 +405,7 @@ def snapshot(self) -> dict[str, Any]: "streaming": stream_snap, "websocket": ws_snap, "degradations": deg_snap, + "tokenEconomy": token_economy, } def model_health(self) -> dict[str, Any]: diff --git a/app/services/prompt_messages.py b/app/services/prompt_messages.py new file mode 100644 index 0000000..fc758d4 --- /dev/null +++ b/app/services/prompt_messages.py @@ -0,0 +1,12 @@ +from __future__ import annotations + + +def split_prompt_for_adapter( + prompt: str, + anchor: str = "符号消息:", +) -> tuple[str, str]: + """Split a composed prompt into non-overlapping system and user messages.""" + split_idx = prompt.rfind(anchor) + if split_idx <= 0: + return "", prompt + return prompt[:split_idx], prompt[split_idx:] diff --git a/app/services/response_quality.py b/app/services/response_quality.py new file mode 100644 index 0000000..99fbaaf --- /dev/null +++ b/app/services/response_quality.py @@ -0,0 +1,28 @@ +from __future__ import annotations + +import re + + +def estimate_response_quality(request: str, response: str) -> float: + """Return a cheap operational quality proxy in the range 0..1. + + This is intentionally not presented as semantic evaluation. It detects + empty/error responses, excessive repetition, and obviously incomplete + answers so token reductions can be correlated with regressions. + """ + clean = response.strip() + if not clean: + return 0.0 + error_markers = ("模型调用异常", "模型调用失败", "traceback", "internal server error") + if any(marker in clean.lower() for marker in error_markers): + return 0.15 + + score = 0.45 + if len(clean) >= min(80, max(20, len(request) // 2)): + score += 0.2 + lines = [re.sub(r"\s+", " ", line.strip()).lower() for line in clean.splitlines() if line.strip()] + unique_ratio = len(set(lines)) / max(1, len(lines)) + score += 0.2 * unique_ratio + if clean[-1] in ".!?。!?`}]))": + score += 0.15 + return round(min(1.0, score), 4) diff --git a/app/services/task_decomposer.py b/app/services/task_decomposer.py index ef5daff..12c2f6f 100644 --- a/app/services/task_decomposer.py +++ b/app/services/task_decomposer.py @@ -10,6 +10,7 @@ from app.db.session import afetch_all from app.schemas.dag import DAGConfig, DAGNode from app.services.context_compaction import build_agent_roster_summary, compact_text +from app.services.context_summary_cache import context_summary_cache from app.services.template_engine import template_engine logger = logging.getLogger("agenthub.task_decomposer") @@ -72,7 +73,6 @@ class ArchitectTaskDecomposer: DECOMPOSE_TIMEOUT = 30.0 def __init__(self) -> None: - self._agent_capability_cache: dict[str, tuple[float, str]] = {} self._historical_context_cache: dict[str, tuple[float, str | None]] = {} self._cache_ttl: float = 300.0 # 5 min @@ -161,14 +161,15 @@ async def _llm_decompose( async def _build_capability_summary(self, agents: list[dict[str, Any]]) -> str: """Build a concise agent capability description for the prompt.""" fingerprint = self._fingerprint_agents(agents) - cached = self._agent_capability_cache.get(fingerprint) - now_ts = time.monotonic() - if cached and (now_ts - cached[0]) < self._cache_ttl: - return cached[1] - - summary = build_agent_roster_summary(agents, max_agents=6, max_tags=3, max_duty_chars=48) - self._agent_capability_cache[fingerprint] = (now_ts, summary) - return summary + owner_id = str(agents[0].get("user_id", "shared")) if agents else "shared" + return context_summary_cache.get_or_build( + "agent", + owner_id, + fingerprint, + lambda: build_agent_roster_summary( + agents, max_agents=6, max_tags=3, max_duty_chars=48, + ), + ) async def _build_historical_context(self, content: str) -> str | None: """Query task_execution_history for relevant past performance.""" diff --git a/app/services/test_context_summary_cache.py b/app/services/test_context_summary_cache.py new file mode 100644 index 0000000..82eac20 --- /dev/null +++ b/app/services/test_context_summary_cache.py @@ -0,0 +1,27 @@ +from __future__ import annotations + +from app.services.context_summary_cache import ContextSummaryCache + + +def test_version_invalidation_forces_summary_rebuild() -> None: + cache = ContextSummaryCache(ttl_seconds=60) + calls = 0 + + def build() -> str: + nonlocal calls + calls += 1 + return f"summary-{calls}" + + assert cache.get_or_build("agent", "u1", "v1", build) == "summary-1" + assert cache.get_or_build("agent", "u1", "v1", build) == "summary-1" + cache.invalidate("agent", "u1") + assert cache.get_or_build("agent", "u1", "v1", build) == "summary-2" + assert cache.stats()["hits"] == 1 + + +def test_versioned_slot_is_removed_by_owner_invalidation() -> None: + cache = ContextSummaryCache(ttl_seconds=60) + cache.set("route", "u1", "active", "route-index") + assert cache.get("route", "u1", "active") == "route-index" + cache.invalidate("route", "u1") + assert cache.get("route", "u1", "active") is None diff --git a/app/services/test_distributed_cache_versions_integration.py b/app/services/test_distributed_cache_versions_integration.py new file mode 100644 index 0000000..59f3594 --- /dev/null +++ b/app/services/test_distributed_cache_versions_integration.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +import asyncio +import os +import uuid + +import pytest + +from app.services.context_summary_cache import context_summary_cache +from app.services.distributed_cache_versions import DistributedCacheVersionBus + + +def test_redis_version_event_invalidates_peer_cache() -> None: + if os.getenv("AGENTHUB_RUN_REDIS_INTEGRATION") != "1": + pytest.skip("set AGENTHUB_RUN_REDIS_INTEGRATION=1 with Redis running") + + async def run() -> None: + listener = DistributedCacheVersionBus() + publisher = DistributedCacheVersionBus() + assert await listener.start() + assert await publisher.start() + owner_id = f"integration-{uuid.uuid4().hex[:12]}" + context_summary_cache.set("route", owner_id, "active-routes", "cached") + assert context_summary_cache.get("route", owner_id, "active-routes") == "cached" + + version = await publisher.publish("route", owner_id) + assert version > 0 + for _ in range(30): + await asyncio.sleep(0.1) + if context_summary_cache.get("route", owner_id, "active-routes") is None: + break + assert context_summary_cache.get("route", owner_id, "active-routes") is None + await publisher.close() + await listener.close() + + asyncio.run(run()) diff --git a/app/services/test_memory_context.py b/app/services/test_memory_context.py new file mode 100644 index 0000000..9b55841 --- /dev/null +++ b/app/services/test_memory_context.py @@ -0,0 +1,45 @@ +from __future__ import annotations + +from app.services.memory_context import ( + MemoryContextSection, + build_memory_context, + deduplicate_text, +) + + +def test_deduplicate_text_removes_history_overlap() -> None: + history = "user: deploy service\nassistant: deployment finished" + candidate = history + "\n\nUnresolved: verify production health" + result = deduplicate_text(candidate, [history]) + assert "deployment finished" not in result + assert "verify production health" in result + + +def test_build_memory_context_prioritizes_session_summary() -> None: + result, stats = build_memory_context( + [ + MemoryContextSection("global-summary", "global preference", 3), + MemoryContextSection("session-summary", "current decision", 1), + ], + max_tokens=20, + provider="unknown", + model="unknown", + ) + assert "session-summary" in result + assert stats["tokens_after"] <= 20 + + +def test_build_memory_context_enforces_per_type_budget() -> None: + result, stats = build_memory_context( + [ + MemoryContextSection("episode", "E" * 2000, 1, "episodic"), + MemoryContextSection("procedure", "P" * 2000, 2, "procedural"), + ], + max_tokens=300, + provider="unknown", + model="unknown", + section_budgets={"episodic": 100, "procedural": 200}, + ) + assert "episode" in result + assert "procedure" in result + assert stats["tokens_after"] <= 300 diff --git a/app/services/test_response_quality.py b/app/services/test_response_quality.py new file mode 100644 index 0000000..bcf2cc1 --- /dev/null +++ b/app/services/test_response_quality.py @@ -0,0 +1,11 @@ +from __future__ import annotations + +from app.services.response_quality import estimate_response_quality + + +def test_response_quality_penalizes_errors_and_repetition() -> None: + good = estimate_response_quality("implement auth", "Implemented authentication and added tests.") + error = estimate_response_quality("implement auth", "模型调用异常:timeout") + repeated = estimate_response_quality("implement auth", "same\nsame\nsame") + assert good > error + assert good > repeated diff --git a/app/services/test_token_budget.py b/app/services/test_token_budget.py new file mode 100644 index 0000000..3bbcd15 --- /dev/null +++ b/app/services/test_token_budget.py @@ -0,0 +1,55 @@ +from __future__ import annotations + +from app.services.token_budget import ( + TokenBudget, + cognitive_memory_budgets, + count_tokens, + fit_prompt, + register_model_tokenizer, + tokenizer_backend, + truncate_to_tokens, + unregister_model_tokenizer, +) + + +def test_count_tokens_handles_cjk_and_latin() -> None: + assert count_tokens("hello world", "unknown", "unknown") > 0 + assert count_tokens("企业多智能体协作平台", "unknown", "unknown") >= 8 + + +def test_truncate_to_tokens_preserves_head_and_tail() -> None: + text = "HEAD-" + ("中" * 5000) + "-TAIL" + result, truncated = truncate_to_tokens(text, 200, "unknown", "unknown") + assert truncated is True + assert result.startswith("HEAD-") + assert result.endswith("-TAIL") + assert count_tokens(result, "unknown", "unknown") <= 200 + + +def test_fit_prompt_applies_model_budget_to_all_prompt_types(monkeypatch) -> None: + monkeypatch.setenv("AGENTHUB_MAX_PROMPT_TOKENS", "2048") + prompt = "system rules\n符号消息: " + ("任务上下文" * 4000) + result, stats = fit_prompt(prompt, "unknown", "custom-model", anchor="符号消息:") + budget = TokenBudget.for_model("unknown", "custom-model") + assert stats["truncated"] is True + assert count_tokens(result, "unknown", "custom-model") <= budget.prompt_limit + assert result.startswith("system rules") + + +def test_cognitive_budgets_follow_task_intent() -> None: + coding = cognitive_memory_budgets(4000, "实现 DAG 部署工具") + research = cognitive_memory_budgets(4000, "调研并比较知识库方案") + chat = cognitive_memory_budgets(4000, "继续刚才的话题") + assert sum(coding.values()) == 4000 + assert coding["procedural"] > chat["procedural"] + assert research["semantic"] > coding["semantic"] + assert chat["working"] > research["working"] + + +def test_provider_native_tokenizer_registration() -> None: + register_model_tokenizer("qwen", lambda text: len(text.split()) * 2) + try: + assert count_tokens("one two three", "qwen", "qwen-max") == 6 + assert tokenizer_backend("qwen", "qwen-max") == "registered-native" + finally: + unregister_model_tokenizer("qwen") diff --git a/app/services/test_tool_prompt_split.py b/app/services/test_tool_prompt_split.py new file mode 100644 index 0000000..b02c191 --- /dev/null +++ b/app/services/test_tool_prompt_split.py @@ -0,0 +1,18 @@ +from __future__ import annotations + +from app.services.prompt_messages import split_prompt_for_adapter + + +def test_static_prefix_is_not_duplicated_in_user_prompt() -> None: + full = "STATIC-SYSTEM\nmemory and tools\n符号消息: dynamic request" + system_prompt, user_prompt = split_prompt_for_adapter(full) + assert system_prompt == "STATIC-SYSTEM\nmemory and tools\n" + assert user_prompt == "符号消息: dynamic request" + assert "STATIC-SYSTEM" not in user_prompt + assert system_prompt + user_prompt == full + + +def test_unsplittable_prompt_remains_user_only() -> None: + system_prompt, user_prompt = split_prompt_for_adapter("plain request") + assert system_prompt == "" + assert user_prompt == "plain request" diff --git a/app/services/test_workflow_contract.py b/app/services/test_workflow_contract.py new file mode 100644 index 0000000..2373257 --- /dev/null +++ b/app/services/test_workflow_contract.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +from app.services.workflow_contract import validate_workflow_contract + + +def _node(node_id: str, dependencies: list[str] | None = None) -> dict: + return { + "id": node_id, + "type": "agent", + "agent": "CodeGen", + "dependencies": dependencies or [], + "x": 10, + "y": 20, + } + + +def test_dependencies_are_normalized_to_editor_edges() -> None: + result = validate_workflow_contract([_node("plan"), _node("code", ["plan"])]) + + assert result.valid + assert result.normalized is not None + assert result.normalized["schemaVersion"] == 1 + assert result.normalized["edges"] == [ + {"id": "plan->code", "from": "plan", "to": "code", "label": "", "condition": ""} + ] + + +def test_explicit_edges_are_authoritative_for_runtime_dependencies() -> None: + result = validate_workflow_contract( + [_node("plan"), _node("code", ["stale"])], + [{"id": "edge-1", "from": "plan", "to": "code"}], + ) + + assert result.valid + assert result.normalized is not None + code = next(node for node in result.normalized["nodes"] if node["id"] == "code") + assert code["dependencies"] == ["plan"] + + +def test_duplicate_nodes_and_edges_return_structured_issues() -> None: + result = validate_workflow_contract( + [_node("same"), _node("same")], + [ + {"id": "duplicate", "from": "same", "to": "same"}, + {"id": "duplicate", "from": "same", "to": "same"}, + ], + ) + + assert not result.valid + assert {issue.code for issue in result.issues} >= { + "duplicate_node_id", "duplicate_edge_id", "duplicate_edge", "self_loop" + } + + +def test_missing_endpoint_and_cycle_are_rejected() -> None: + missing = validate_workflow_contract( + [_node("a")], [{"from": "missing", "to": "a"}], + ) + cyclic = validate_workflow_contract( + [_node("a"), _node("b")], + [{"from": "a", "to": "b"}, {"from": "b", "to": "a"}], + ) + + assert not missing.valid + assert any(issue.code == "missing_edge_source" for issue in missing.issues) + assert not cyclic.valid + assert any(issue.code == "cycle_detected" for issue in cyclic.issues) + + +def test_unassigned_agent_is_a_recoverable_warning() -> None: + result = validate_workflow_contract([{"id": "draft-agent", "type": "agent"}]) + + assert result.valid + assert any(issue.code == "agent_unassigned" and issue.severity == "warning" for issue in result.issues) + + +def test_empty_workflow_and_incomplete_runtime_config_are_rejected() -> None: + empty = validate_workflow_contract([]) + incomplete = validate_workflow_contract([{"id": "request", "type": "http"}]) + + assert not empty.valid + assert any(issue.code == "empty_workflow" for issue in empty.issues) + assert not incomplete.valid + assert any(issue.code == "missing_node_config" for issue in incomplete.issues) diff --git a/app/services/token_budget.py b/app/services/token_budget.py new file mode 100644 index 0000000..1c1cec2 --- /dev/null +++ b/app/services/token_budget.py @@ -0,0 +1,249 @@ +from __future__ import annotations + +import os +import re +from dataclasses import dataclass +from functools import lru_cache +from typing import Iterable +from collections.abc import Callable +from pathlib import Path + + +_MODEL_WINDOWS: tuple[tuple[str, int], ...] = ( + ("gpt-4.1", 1_047_576), + ("gpt-4o", 128_000), + ("gpt-5", 400_000), + ("o1", 200_000), + ("o3", 200_000), + ("claude", 200_000), + ("deepseek", 64_000), + ("qwen", 128_000), + ("doubao", 128_000), + ("glm", 128_000), + ("kimi", 128_000), +) + +_REGISTERED_TOKENIZERS: dict[str, Callable[[str], int]] = {} + + +def register_model_tokenizer( + provider: str, + counter: Callable[[str], int], + model: str = "", +) -> None: + key = f"{provider.lower()}:{model.lower()}" if model else provider.lower() + _REGISTERED_TOKENIZERS[key] = counter + + +def unregister_model_tokenizer(provider: str, model: str = "") -> None: + key = f"{provider.lower()}:{model.lower()}" if model else provider.lower() + _REGISTERED_TOKENIZERS.pop(key, None) + + +@lru_cache(maxsize=64) +def _local_provider_tokenizer(provider: str, model: str): + env_provider = re.sub(r"[^A-Z0-9]+", "_", provider.upper()).strip("_") + configured = os.getenv(f"AGENTHUB_TOKENIZER_{env_provider}_PATH", "").strip() + if not configured: + return None + path = Path(configured) + tokenizer_file = path / "tokenizer.json" if path.is_dir() else path + if not tokenizer_file.is_file(): + return None + try: + from tokenizers import Tokenizer + + return Tokenizer.from_file(str(tokenizer_file)) + except (ImportError, OSError, ValueError): + return None + + +@lru_cache(maxsize=64) +def _tiktoken_encoder(model: str): + try: + import tiktoken + + try: + return tiktoken.encoding_for_model(model) + except KeyError: + return tiktoken.get_encoding("o200k_base") + except (ImportError, ValueError): + return None + + +def count_tokens(text: str, provider: str = "", model: str = "") -> int: + """Count model input tokens, using the provider tokenizer when available. + + OpenAI-compatible models use tiktoken. Other providers fall back to a + multilingual estimator until their native tokenizer is installed. The + fallback counts CJK characters individually and groups latin text, which + is deliberately conservative for budget enforcement. + """ + if not text: + return 0 + provider_key = provider.lower() + model_key = model.lower() + custom = _REGISTERED_TOKENIZERS.get(f"{provider_key}:{model_key}") or _REGISTERED_TOKENIZERS.get(provider_key) + if custom is not None: + return max(1, int(custom(text))) + local_tokenizer = _local_provider_tokenizer(provider_key, model_key) + if local_tokenizer is not None: + return len(local_tokenizer.encode(text).ids) + if provider_key in {"openai", "azure_openai"} or model_key.startswith(("gpt-", "o1", "o3")): + encoder = _tiktoken_encoder(model or "gpt-4o") + if encoder is not None: + return len(encoder.encode(text, disallowed_special=())) + + cjk = len(re.findall(r"[\u3400-\u9fff\uf900-\ufaff]", text)) + non_cjk = len(text) - cjk + return max(1, cjk + (non_cjk + 3) // 4) + + +def tokenizer_backend(provider: str = "", model: str = "") -> str: + provider_key, model_key = provider.lower(), model.lower() + if f"{provider_key}:{model_key}" in _REGISTERED_TOKENIZERS or provider_key in _REGISTERED_TOKENIZERS: + return "registered-native" + if _local_provider_tokenizer(provider_key, model_key) is not None: + return "local-tokenizer-json" + if provider_key in {"openai", "azure_openai"} or model_key.startswith(("gpt-", "o1", "o3")): + return "tiktoken" if _tiktoken_encoder(model or "gpt-4o") is not None else "multilingual-estimator" + return "multilingual-estimator" + + +def model_context_window(provider: str = "", model: str = "") -> int: + override = os.getenv("AGENTHUB_MODEL_CONTEXT_TOKENS", "").strip() + if override.isdigit(): + return max(4_096, int(override)) + key = f"{provider}/{model}".lower() + for fragment, window in _MODEL_WINDOWS: + if fragment in key: + return window + return 32_768 + + +@dataclass(frozen=True) +class TokenBudget: + provider: str + model: str + context_window: int + output_reserve: int + prompt_limit: int + + @classmethod + def for_model( + cls, + provider: str = "", + model: str = "", + *, + output_reserve: int = 4_096, + ) -> "TokenBudget": + context_window = model_context_window(provider, model) + configured_cap = int(os.getenv("AGENTHUB_MAX_PROMPT_TOKENS", "20000")) + prompt_limit = max(2_048, min(configured_cap, context_window - output_reserve)) + return cls(provider, model, context_window, output_reserve, prompt_limit) + + def section_limit(self, section: str) -> int: + shares = { + "history": 0.18, + "memory": 0.14, + "preprocess": 0.08, + "collaboration": 0.10, + "tools": 0.20, + "user": 0.30, + } + return max(256, int(self.prompt_limit * shares.get(section, 0.10))) + + +def cognitive_memory_budgets( + total_tokens: int, + query: str, + domain: str = "", +) -> dict[str, int]: + """Allocate one context pool across the four cognitive memory classes.""" + text = f"{domain} {query}".lower() + if any(term in text for term in ("research", "分析", "调研", "知识", "比较", "search")): + shares = {"working": 0.25, "episodic": 0.15, "semantic": 0.50, "procedural": 0.10} + elif any(term in text for term in ("code", "实现", "修复", "部署", "workflow", "dag", "工具", "sop")): + shares = {"working": 0.30, "episodic": 0.20, "semantic": 0.15, "procedural": 0.35} + elif any(term in text for term in ("plan", "规划", "方案", "架构", "复盘")): + shares = {"working": 0.25, "episodic": 0.30, "semantic": 0.20, "procedural": 0.25} + else: + shares = {"working": 0.40, "episodic": 0.30, "semantic": 0.20, "procedural": 0.10} + + total = max(1024, total_tokens) + budgets = {name: max(128, int(total * share)) for name, share in shares.items()} + difference = total - sum(budgets.values()) + budgets["working"] += difference + return budgets + + +def truncate_to_tokens( + text: str, + max_tokens: int, + provider: str = "", + model: str = "", + *, + preserve_tail: float = 0.75, + marker: str = "\n... [context truncated] ...\n", +) -> tuple[str, bool]: + if max_tokens <= 0: + return "", bool(text) + if count_tokens(text, provider, model) <= max_tokens: + return text, False + + # Binary search character length because all tokenizer implementations are + # monotonic with respect to a prefix/suffix slice. + low, high = 0, len(text) + marker_tokens = count_tokens(marker, provider, model) + target = max(1, max_tokens - marker_tokens) + while low < high: + mid = (low + high + 1) // 2 + head_len = int(mid * (1.0 - preserve_tail)) + candidate = text[:head_len] + text[-(mid - head_len):] + if count_tokens(candidate, provider, model) <= target: + low = mid + else: + high = mid - 1 + head_len = int(low * (1.0 - preserve_tail)) + tail_len = low - head_len + compacted = text[:head_len] + marker + (text[-tail_len:] if tail_len else "") + return compacted, True + + +def fit_prompt( + prompt: str, + provider: str = "", + model: str = "", + *, + output_reserve: int = 4_096, + anchor: str = "", +) -> tuple[str, dict[str, int | bool]]: + budget = TokenBudget.for_model(provider, model, output_reserve=output_reserve) + before = count_tokens(prompt, provider, model) + if before <= budget.prompt_limit: + return prompt, {"tokens_before": before, "tokens_after": before, "truncated": False} + + if anchor and anchor in prompt: + index = prompt.rfind(anchor) + len(anchor) + prefix, dynamic = prompt[:index], prompt[index:] + prefix_tokens = count_tokens(prefix, provider, model) + if prefix_tokens >= budget.prompt_limit - 256: + prefix, _ = truncate_to_tokens( + prefix, + budget.prompt_limit - 256, + provider, + model, + preserve_tail=0.25, + ) + dynamic_budget = max(256, budget.prompt_limit - count_tokens(prefix, provider, model)) + dynamic, _ = truncate_to_tokens(dynamic, dynamic_budget, provider, model) + result = prefix + dynamic + else: + result, _ = truncate_to_tokens(prompt, budget.prompt_limit, provider, model) + + after = count_tokens(result, provider, model) + return result, {"tokens_before": before, "tokens_after": after, "truncated": True} + + +def total_tokens(parts: Iterable[str], provider: str = "", model: str = "") -> int: + return sum(count_tokens(part, provider, model) for part in parts) diff --git a/app/services/workflow_contract.py b/app/services/workflow_contract.py new file mode 100644 index 0000000..59ee155 --- /dev/null +++ b/app/services/workflow_contract.py @@ -0,0 +1,186 @@ +from __future__ import annotations + +from collections import Counter +from typing import Any + +from pydantic import ValidationError + +from app.schemas.dag import DAGConfig, DAGEdge, DAGNode +from app.schemas.workflow import DAGValidationIssue, DAGValidationResult + + +MAX_WORKFLOW_NODES = 200 +MAX_WORKFLOW_EDGES = 400 + + +def validate_workflow_contract( + nodes: list[dict[str, Any]], + edges: list[dict[str, Any]] | None = None, + *, + schema_version: int = 1, +) -> DAGValidationResult: + """Validate and normalize editor data without touching persistence.""" + issues: list[DAGValidationIssue] = [] + raw_edges = edges or [] + if len(nodes) > MAX_WORKFLOW_NODES: + issues.append(_issue("node_limit", f"Workflow exceeds {MAX_WORKFLOW_NODES} nodes")) + if len(raw_edges) > MAX_WORKFLOW_EDGES: + issues.append(_issue("edge_limit", f"Workflow exceeds {MAX_WORKFLOW_EDGES} edges")) + + parsed_nodes: list[DAGNode] = [] + for index, raw in enumerate(nodes[: MAX_WORKFLOW_NODES + 1]): + try: + node = DAGNode.model_validate(raw) + except ValidationError as exc: + issues.append(_issue("invalid_node", f"Node {index}: {exc.errors()[0]['msg']}")) + continue + node.id = node.id.strip() + if not node.id: + issues.append(_issue("empty_node_id", "Node ID cannot be empty", node_id=node.id)) + parsed_nodes.append(node) + + if not parsed_nodes: + issues.append(_issue("empty_workflow", "Workflow must contain at least one node")) + + id_counts = Counter(node.id for node in parsed_nodes if node.id) + for node_id, count in id_counts.items(): + if count > 1: + issues.append(_issue("duplicate_node_id", f"Duplicate node ID: {node_id}", node_id=node_id)) + node_ids = set(id_counts) + + parsed_edges: list[DAGEdge] = [] + if raw_edges: + for index, raw in enumerate(raw_edges[: MAX_WORKFLOW_EDGES + 1]): + try: + edge = DAGEdge.model_validate(raw) + except ValidationError as exc: + issues.append(_issue("invalid_edge", f"Edge {index}: {exc.errors()[0]['msg']}")) + continue + edge.source = edge.source.strip() + edge.target = edge.target.strip() + edge.id = edge.id.strip() or f"{edge.source}->{edge.target}" + parsed_edges.append(edge) + else: + for node in parsed_nodes: + for dependency in node.dependencies: + parsed_edges.append( + DAGEdge(id=f"{dependency}->{node.id}", **{"from": dependency, "to": node.id}) + ) + + edge_id_counts = Counter(edge.id for edge in parsed_edges) + pair_counts = Counter((edge.source, edge.target) for edge in parsed_edges) + for edge in parsed_edges: + if edge_id_counts[edge.id] > 1: + issues.append(_issue("duplicate_edge_id", f"Duplicate edge ID: {edge.id}", edge_id=edge.id)) + if pair_counts[(edge.source, edge.target)] > 1: + issues.append(_issue("duplicate_edge", f"Duplicate edge: {edge.source} -> {edge.target}", edge_id=edge.id)) + if edge.source == edge.target: + issues.append(_issue("self_loop", "An edge cannot connect a node to itself", edge_id=edge.id)) + if edge.source not in node_ids: + issues.append(_issue("missing_edge_source", f"Edge source does not exist: {edge.source}", edge_id=edge.id)) + if edge.target not in node_ids: + issues.append(_issue("missing_edge_target", f"Edge target does not exist: {edge.target}", edge_id=edge.id)) + + incoming: dict[str, list[str]] = {node_id: [] for node_id in node_ids} + for edge in parsed_edges: + if edge.source in node_ids and edge.target in node_ids and edge.source != edge.target: + incoming[edge.target].append(edge.source) + if raw_edges: + for node in parsed_nodes: + node.dependencies = list(dict.fromkeys(incoming.get(node.id, []))) + else: + for node in parsed_nodes: + for dependency in node.dependencies: + if dependency not in node_ids: + issues.append(_issue("missing_dependency", f"Dependency does not exist: {dependency}", node_id=node.id)) + + _validate_cycles(incoming, issues) + start_nodes = [node for node in parsed_nodes if node.type == "start"] + if len(start_nodes) > 1: + issues.append(_issue("multiple_start_nodes", "A workflow can contain at most one start node")) + if parsed_nodes and not any(node.type in {"agent", "tool", "code", "http", "knowledge", "human", "end"} for node in parsed_nodes): + issues.append(_issue("no_executable_node", "Workflow has no executable or end node")) + for node in parsed_nodes: + if node.type == "agent" and not (node.agent or node.domain): + issues.append(_issue("agent_unassigned", "Agent node has no assigned agent", "warning", node_id=node.id)) + config = node.config or (node.model_extra or {}).get(f"{node.type}Config", {}) + required_field = { + "code": "code", + "http": "url", + "knowledge": "collectionId", + "human": "prompt", + }.get(node.type) + if required_field and not str(config.get(required_field, "")).strip(): + issues.append( + _issue( + "missing_node_config", + f"{node.type} node requires {required_field}", + node_id=node.id, + ) + ) + + has_errors = any(issue.severity == "error" for issue in issues) + normalized = None + if not has_errors: + dag = DAGConfig( + schema_version=schema_version, + version=1, + total=len(parsed_nodes), + nodes=parsed_nodes, + edges=parsed_edges, + ) + normalized = dag.model_dump(mode="json", by_alias=True) + return DAGValidationResult(valid=not has_errors, normalized=normalized, issues=_dedupe_issues(issues)) + + +def require_valid_workflow( + nodes: list[dict[str, Any]], edges: list[dict[str, Any]] | None = None, *, schema_version: int = 1, +) -> dict[str, Any]: + result = validate_workflow_contract(nodes, edges, schema_version=schema_version) + if not result.valid or result.normalized is None: + messages = "; ".join(issue.message for issue in result.issues if issue.severity == "error") + raise ValueError(messages or "Invalid workflow") + return result.normalized + + +def _validate_cycles(incoming: dict[str, list[str]], issues: list[DAGValidationIssue]) -> None: + visiting: set[str] = set() + visited: set[str] = set() + + def visit(node_id: str) -> bool: + if node_id in visiting: + return True + if node_id in visited: + return False + visiting.add(node_id) + cyclic = any(visit(dependency) for dependency in incoming.get(node_id, [])) + visiting.remove(node_id) + visited.add(node_id) + return cyclic + + if any(visit(node_id) for node_id in incoming): + issues.append(_issue("cycle_detected", "Workflow contains a dependency cycle")) + + +def _issue( + code: str, + message: str, + severity: str = "error", + *, + node_id: str | None = None, + edge_id: str | None = None, +) -> DAGValidationIssue: + return DAGValidationIssue( + code=code, message=message, severity=severity, nodeId=node_id, edgeId=edge_id, + ) + + +def _dedupe_issues(issues: list[DAGValidationIssue]) -> list[DAGValidationIssue]: + seen: set[tuple[str, str | None, str | None]] = set() + result: list[DAGValidationIssue] = [] + for issue in issues: + key = (issue.code, issue.nodeId, issue.edgeId) + if key not in seen: + seen.add(key) + result.append(issue) + return result diff --git a/app/services/workflow_draft_service.py b/app/services/workflow_draft_service.py new file mode 100644 index 0000000..9cfb646 --- /dev/null +++ b/app/services/workflow_draft_service.py @@ -0,0 +1,112 @@ +from __future__ import annotations + +import json +from typing import Any + +from app.db.init_db import now +from app.db.session import afetch_all, afetch_one, aexecute +from app.schemas.workflow import WorkflowDraftRequest +from app.services.workflow_contract import MAX_WORKFLOW_EDGES, MAX_WORKFLOW_NODES, validate_workflow_contract +from app.services.workflow_errors import WorkflowVersionConflict + + +MAX_DRAFT_BYTES = 1_000_000 + + +class WorkflowDraftService: + async def list_drafts(self, user_id: str) -> list[dict[str, Any]]: + rows = await afetch_all( + "SELECT draft_key,workflow_id,name,base_version,version,created_at,updated_at " + "FROM workflow_drafts WHERE user_id=$1 ORDER BY updated_at DESC", + user_id, + ) + return [self._metadata(row) for row in rows] + + async def get_draft(self, user_id: str, draft_key: str) -> dict[str, Any] | None: + row = await afetch_one( + "SELECT draft_key,workflow_id,name,payload_json,base_version,version,created_at,updated_at " + "FROM workflow_drafts WHERE user_id=$1 AND draft_key=$2", + user_id, + draft_key, + ) + return self._deserialize(row) if row else None + + async def save_draft(self, user_id: str, draft_key: str, data: WorkflowDraftRequest) -> dict[str, Any]: + if len(data.nodes) > MAX_WORKFLOW_NODES or len(data.edges) > MAX_WORKFLOW_EDGES: + raise ValueError("Draft exceeds workflow node or edge limits") + payload = data.model_dump(mode="json", exclude={"draftVersion"}) + payload_json = json.dumps(payload, ensure_ascii=False) + if len(payload_json.encode("utf-8")) > MAX_DRAFT_BYTES: + raise ValueError("Draft payload exceeds 1 MB") + + timestamp = now() + if data.draftVersion == 0: + saved = await afetch_one( + "INSERT INTO workflow_drafts(user_id,workflow_id,draft_key,name,payload_json,base_version," + "version,created_at,updated_at) VALUES($1,$2,$3,$4,$5,$6,1,$7,$8) " + "ON CONFLICT(user_id,draft_key) DO NOTHING RETURNING draft_key,version", + user_id, + data.workflowId, + draft_key, + data.name, + payload_json, + data.baseVersion, + timestamp, + timestamp, + ) + else: + saved = await afetch_one( + "UPDATE workflow_drafts SET workflow_id=$1,name=$2,payload_json=$3,base_version=$4," + "version=version+1,updated_at=$5 WHERE user_id=$6 AND draft_key=$7 AND version=$8 " + "RETURNING draft_key,version", + data.workflowId, + data.name, + payload_json, + data.baseVersion, + timestamp, + user_id, + draft_key, + data.draftVersion, + ) + if not saved: + current = await afetch_one( + "SELECT version FROM workflow_drafts WHERE user_id=$1 AND draft_key=$2", user_id, draft_key, + ) + raise WorkflowVersionConflict(data.draftVersion, int(current["version"]) if current else 0) + result = await self.get_draft(user_id, draft_key) + if result is None: + raise ValueError("Draft save failed") + return result + + async def delete_draft(self, user_id: str, draft_key: str) -> bool: + existing = await self.get_draft(user_id, draft_key) + if not existing: + return False + await aexecute("DELETE FROM workflow_drafts WHERE user_id=$1 AND draft_key=$2", user_id, draft_key) + return True + + @staticmethod + def _metadata(row: dict[str, Any]) -> dict[str, Any]: + return { + "draftKey": row["draft_key"], + "workflowId": row.get("workflow_id"), + "name": row["name"], + "baseVersion": int(row["base_version"]), + "draftVersion": int(row["version"]), + "createdAt": row["created_at"], + "updatedAt": row["updated_at"], + } + + def _deserialize(self, row: dict[str, Any]) -> dict[str, Any]: + payload = json.loads(row["payload_json"]) + validation = validate_workflow_contract( + payload.get("nodes", []), payload.get("edges", []), schema_version=payload.get("schemaVersion", 1), + ) + return { + **self._metadata(row), + "payload": payload, + "validation": validation.model_dump(mode="json"), + } + + +workflow_draft_service = WorkflowDraftService() diff --git a/app/services/workflow_errors.py b/app/services/workflow_errors.py new file mode 100644 index 0000000..243d0fe --- /dev/null +++ b/app/services/workflow_errors.py @@ -0,0 +1,5 @@ +class WorkflowVersionConflict(ValueError): + def __init__(self, expected_version: int, current_version: int) -> None: + super().__init__("Workflow was modified by another editor") + self.expected_version = expected_version + self.current_version = current_version diff --git a/docs/memory-architecture.md b/docs/memory-architecture.md new file mode 100644 index 0000000..0fa8787 --- /dev/null +++ b/docs/memory-architecture.md @@ -0,0 +1,190 @@ +# AgentHub Memory Architecture + +## 1. Current online projection + +Memory now has two orthogonal classifications: + +- Existing `type=user|feedback|project|reference` remains the content category. +- `memory_type=working|episodic|semantic|procedural` is the cognitive type. +- `scope`, `source`, and monotonic `version` define ownership, provenance, and + optimistic evolution without changing the underlying file/database storage. + +AgentHub does not send every stored memory item to every model call. The online +prompt is assembled as a bounded projection: + +| Layer | Current source | Online behavior | Maturity | +| --- | --- | --- | --- | +| L0 working memory | PostgreSQL `messages` | Latest conversation transcript, deduplicated and token-budgeted | Usable | +| L1 session memory | `users/{user}/sessions/{session}/conversation.md` and session summary | Recent durable turns plus semantic session summary | Usable, partly duplicated at rest | +| L2 retrieval memory | File memories and memory search tools | Not injected by default; loaded on demand through `memory_search` | Partial; no unified vector lifecycle | +| L3 global memory | Per-user global summary | Cross-session decisions and preferences, token-budgeted | Partial; freshness and provenance are coarse | +| Segment compaction | Rust memory-segment-core | Triggered from the online session path, emits token reduction and retained messages | Integrated | +| Semantic consolidation | Python summarization-service | Converts Rust structural compaction into a durable semantic summary | Integrated | + +The request path is: + +```text +DB history (L0) + -> token budget + -> session summary + durable recent turns + global summary + -> semantic/exact overlap removal + -> per-model memory budget + -> agent prompt +``` + +Session conversations, session summaries, and `task_execution_history` are +classified as Episodic Memory. Existing Markdown files without the new +frontmatter fields are read as user-scoped Semantic Memory from a legacy file. + +Explicit durable signals in episodic summaries (`preference`, `decision`, +`constraint`, and confirmed `fact`) are projected into structured Semantic +Memory records under each user's existing memory directory. Records retain +source session/event, confidence, status, and version. A conflicting value for +the same semantic key supersedes the old record instead of overwriting its +history. Prompt assembly retrieves query-relevant active records plus durable +preferences; ordinary narrative sentences are not promoted automatically. + +Procedural Memory is exposed as a read-through catalog over existing sources +of truth. Skills, DAG templates, user agent routes, memory files marked as SOP, +registered tools, and tool permission rules keep their original storage and +execution owners. The catalog adds a stable record ID, source/version, +content hash, scope, kind, and risk level without copying procedure bodies. +`GET /api/memory/procedural` provides the combined catalog and query view. + +The prompt budget uses one model-aware cognitive context pool. Allocation is +intent-sensitive: conversational continuity favors Working/Episodic, research +favors Semantic, implementation/deployment/workflow tasks favor Procedural, +and planning balances Episodic with Procedural. Each class has an enforced +sub-budget in addition to the final model context-window guard. + +The asynchronous write-back path is: + +```text +session turn persisted + -> agenthub.memory.compact.requested + -> Rust compact/prune metrics + -> memory.compact.completed + -> Python semantic summary + -> session.summary.generated + -> online summary consumer + -> per-user session summary + prompt cache invalidation +``` + +## 2. Compression policy + +- DB history is limited before prompt construction using the selected model's + tokenizer and context window. +- Session summary has higher prompt priority than raw durable turns. +- Raw durable turns that overlap the DB transcript are removed by normalized + block and shingle similarity. +- Global summary is added last and is removed when it duplicates session + context. +- File-backed knowledge is retrieval-only and does not consume every request's + context window. +- Rust compaction defaults remain 40 messages or 32,000 estimated tokens, with + 10 recent messages retained. +- The online publisher samples every 10 turns; Rust decides whether the actual + threshold has been reached. + +## 3. Token budget + +`app/services/token_budget.py` is the single budget authority. + +- Uses `tiktoken` for supported OpenAI model families. +- Uses a conservative multilingual fallback for providers without a bundled + native tokenizer. +- Resolves model context windows and reserves output capacity. +- Applies section budgets for history, memory, preprocessing, collaboration, + tools, and current user content. +- Applies a final prompt guard to every agent prompt, including specialized + CodeGen, Orchestrator, Architect, and Deploy prompts. + +## 4. Observability + +`GET /api/system/metrics` now exposes `tokenEconomy`: + +- Tokens before and after compaction by section. +- Tokens saved and reduction rate. +- Truncation count. +- Route/agent summary cache entries, hits, misses, and hit ratio. +- Semantic summary tokens and estimated cost. +- Heuristic summary quality score and sample count. +- Operational answer-quality proxy and sample count, for detecting obvious + empty/error/repetition regressions after budget changes. + +The summarization service also exports Prometheus counters/gauges for summary +tokens, estimated cost, compaction tokens before/after, latency, status, and +quality. + +Cross-service verification is available through +`app/services/memory/test_nats_memory_pipeline_integration.py`. With NATS, +memory-segment-core, summarization-service, and model-adapter-service running, +set `AGENTHUB_RUN_NATS_INTEGRATION=1` to verify request, Rust compaction, +semantic summarization, coverage ranges, and final summary publication. + +Summary write-back stores `covered_sequence_start/end`, generation time, and +source event ID. Duplicate events, older coverage, and older generation times +are rejected under a per-session lock. Native Chinese-provider tokenizers can +be registered in code or loaded from local `tokenizer.json` paths through the +`AGENTHUB_TOKENIZER__PATH` variables; the fallback remains explicitly +reported as a multilingual estimator. Route/agent cache invalidation keeps a +local L1 cache while Redis version counters and Pub/Sub invalidate peer workers. +Set `AGENTHUB_RUN_REDIS_INTEGRATION=1` to run the real two-client peer +invalidation test against `REDIS_URL`/`REDIS_ADDR`. + +## 5. Current assessment + +The memory subsystem is now an operational multi-layer pipeline, but it is not +yet a complete ContextOS implementation. + +Strengths: + +- Tenant-scoped durable memory and summaries. +- Prompt projection is bounded and retrieval-first. +- Structural compaction and semantic summarization form an event-driven loop. +- Idempotent online summary consumption and explicit cache invalidation. +- Cost and compression visibility exists at both online and offline layers. + +Remaining gaps: + +1. Native tokenizers are not yet available for every Chinese provider. The + fallback is safe for limits but cannot guarantee exact billing parity. +2. L2 lacks a unified embedding version, retention policy, provenance model, + and deletion propagation across vector and file stores. +3. L3 global summaries do not yet distinguish durable facts, preferences, + hypotheses, and expired information. +4. Summary quality is a heuristic operational signal, not an evaluator-model + or human-rated regression score. +5. In-process caches and the online consumer are worker-local. Multi-replica + deployment needs Redis-backed versions and a single durable consumer group. +6. Summary write-back is last-write-wins. Concurrent summaries need sequence + or covered-range conflict checks. + +## 6. Next improvements + +### Short term + +- Add Qwen, DeepSeek, Doubao, GLM, and Claude native tokenizer adapters. +- Include covered message sequence ranges in summary state and reject stale + write-back events. +- Add integration tests with NATS, Rust core, summarization-service, and the + online consumer in Docker Compose. +- Replace the heuristic history metric with exact before/after counters on all + streaming and non-streaming paths. + +### Medium term + +- Introduce a memory record schema containing tenant, session, source, + provenance, confidence, sensitivity, TTL, embedding version, and tombstone. +- Build L2 vector indexing with deletion propagation and tenant filters. +- Split L3 into durable facts, user/team preferences, decisions, and expired + candidates; refresh incrementally instead of regenerating all summaries. +- Move cache versions and summary checkpoints to Redis/PostgreSQL for replicas. + +### Long term + +- Implement AutoDream consolidation with contradiction detection and approval + rules for high-risk enterprise memory. +- Add knowledge-graph memory with temporal edges and source citations. +- Evaluate summaries against retained facts, unresolved tasks, and answer + quality using offline datasets and canary traffic. diff --git a/docs/optimization-roadmap.md b/docs/optimization-roadmap.md index af1c394..6ef0c09 100644 --- a/docs/optimization-roadmap.md +++ b/docs/optimization-roadmap.md @@ -7,6 +7,13 @@ AgentHub is already past the "chat wrapper" stage. The current codebase has a us What is now true in the repository: - `frontend/app/page.tsx` is already thinner than before and now delegates message recovery, WebSocket URL building, and DAG state to helpers. +- `frontend/hooks/useSessionWebSocket.ts` and `frontend/hooks/useSessionRecovery.ts` now carry the connection lifecycle and reconnect / restore logic out of the page shell. +- `frontend/__tests__/components/chat/taskPreviewReplay.test.tsx` and `frontend/__tests__/components/chat/dagReplay.test.tsx` cover duplicate preview events, reconnect replay, and session switching. +- `frontend/components/admin/PermissionModule.tsx` now replaces the stale `权限` placeholder with a real permission-rule surface backed by `/api/admin/permissions/rules`. +- The admin shell now has explicit mobile, compact-laptop, and wide-desktop behavior; compact layouts collapse the sidebar and remove controls that cannot operate at that width. +- Permission rules now support validated create/edit/toggle/delete flows, request retry feedback, and focused interaction tests. +- Compact admin sidebars now expand to their persisted width without being clipped by the responsive grid. +- Recovery now advances the replay cursor for all identified events, merges final messages over streaming placeholders, and restores persisted DAG progress without overwriting newer live updates. - Prompt/token control has started through `app/services/context_compaction.py`, `app/services/conversation_history.py`, `app/services/orchestrator_preprocessor.py`, `app/services/task_decomposer.py`, and `app/services/result_synthesizer.py`. - `app/api/websocket_message_flow.py` and `app/api/websocket.py` now share compact task preview construction. - The repo already carries the architecture needed for enterprise expansion, but it still has a few oversized hot modules. @@ -147,7 +154,14 @@ Success criteria: - Teams can manage tokens and permissions from the platform. - External developers can integrate without reading internal code first. -### 4.4 Long term +### 4.4 Near-term execution order + +1. Finish permission-rule CRUD and validation in the new admin module. +2. Keep thinning the last frontend coupling in message recovery and DAG replay. +3. Add route / agent pre-summary caches and shrink preview payloads one more layer. +4. Only after those are stable, move to DAG editor, template market, SDKs, and token management. + +### 4.5 Long term Focus: @@ -164,6 +178,8 @@ Success criteria: ### Phase A: Stabilize +Current status: complete. Admin responsiveness, permission rules, message recovery, replay cursors, duplicate-event handling, session isolation, and persisted DAG replay are covered by focused tests. + Priority: 1. Keep transport and state recovery correct. @@ -193,6 +209,29 @@ Deliverables: ### Phase C: Token Economy +Current status: in progress. Route / agent versioned caches, tokenizer-aware +budgets, memory deduplication, prompt prefix de-duplication, and the +Rust-to-Python-to-online summary loop are landed. Native provider tokenizers, +distributed cache versions, and end-to-end quality evaluation remain. + +Cognitive memory migration status: + +- Landed: orthogonal `memory_type/scope/source/version` metadata with legacy + Markdown compatibility. +- Landed: session conversations, summaries, and task execution history are + classified as Episodic Memory. +- Landed: structured Semantic extraction with provenance, confidence, version, + conflict supersession, query retrieval, and prompt projection. +- Landed: Skills, DAG templates/routes, SOP files, tool definitions, and tool + permission rules are exposed through a versioned Procedural Memory catalog. +- Landed: task-intent-aware Working/Episodic/Semantic/Procedural token + allocation with per-class limits and a model-window final guard. +- Landed: local/native provider tokenizer adapters with explicit fallback. +- Landed: sequence-safe and event-idempotent summary write-back. +- Landed: opt-in real NATS Rust/Python/online integration test. +- Landed: Redis version counters and Pub/Sub invalidation for worker-local + route/agent summary caches, with single-node fallback. + Priority: 1. Cache route / agent pre-summaries. @@ -204,6 +243,11 @@ Deliverables: - `app/services/context_compaction.py`. - Shorter history and memory context. - Shorter synthesis inputs. +- `app/services/token_budget.py` as the shared model budget authority. +- `app/services/context_summary_cache.py` with explicit version invalidation. +- `app/services/memory_context.py` for layered projection and overlap removal. +- `app/services/memory_summary_consumer.py` for durable summary write-back. +- Memory architecture and maturity assessment in `docs/memory-architecture.md`. ### Phase D: Platform Hardening @@ -221,6 +265,12 @@ Deliverables: ### Phase E: Productization +Current status: DAG editor authoring loop complete. The backend has a versioned +node/edge contract, validation-only API, workflow optimistic locking, and +user-isolated recoverable drafts. The ReactFlow UI now adds lossless contract +adapters, debounced autosave, refresh recovery, node/edge validation markers, +and explicit reload/overwrite handling for `409` conflicts. + Priority: 1. DAG editor. @@ -230,9 +280,21 @@ Priority: Deliverables: - Workflow authoring UI. +- Versioned workflow schema and explicit edge persistence. +- Structured DAG validation and `409` conflict responses. +- Per-user draft save, recovery, listing, and deletion APIs. - Reusable templates. - Public integration surface. +Next small stage: + +1. Define template metadata, categories, versions, ownership, and publication + states over the stable workflow contract. +2. Add template list/detail/install APIs with tenant isolation and audit events. +3. Build the template market browsing, preview, install, and fork experience. +4. Add curated built-in templates and compatibility checks by schema version. +5. After the template loop is stable, proceed to SDK and API Token management. + ## 6. Already Landed These are the concrete foundation pieces now in the repo: @@ -249,6 +311,11 @@ These are the concrete foundation pieces now in the repo: - `frontend/lib/messageRecovery.ts` - `frontend/lib/websocketUrl.ts` - `frontend/lib/outgoingMessageDraft.ts` +- `frontend/hooks/useSessionWebSocket.ts` +- `frontend/hooks/useSessionRecovery.ts` +- `frontend/components/admin/PermissionModule.tsx` +- `frontend/__tests__/components/chat/taskPreviewReplay.test.tsx` +- `frontend/__tests__/components/chat/dagReplay.test.tsx` ## 7. Rule of Thumb @@ -258,4 +325,3 @@ Do not expand feature surface until the core loop stays: - replayable, - observable, - and cheap enough to run repeatedly. - diff --git a/frontend/__tests__/components/admin/PermissionModule.test.tsx b/frontend/__tests__/components/admin/PermissionModule.test.tsx new file mode 100644 index 0000000..f767d62 --- /dev/null +++ b/frontend/__tests__/components/admin/PermissionModule.test.tsx @@ -0,0 +1,134 @@ +import { describe, it, expect, vi, beforeEach, afterEach } from 'vitest'; +import { render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import PermissionModule from '../../../components/admin/PermissionModule'; + +describe('PermissionModule', () => { + beforeEach(() => { + vi.restoreAllMocks(); + }); + + afterEach(() => { + vi.unstubAllGlobals(); + }); + + it('loads permission rules from the admin API instead of showing the placeholder', async () => { + const fetchMock = vi.fn(async () => ({ + ok: true, + status: 200, + json: async () => ([ + { + id: 1, + agentId: '*', + toolPattern: 'file_*', + pathPattern: '/workspace/**', + behavior: 'ask', + source: 'user', + priority: 10, + enabled: true, + createdAt: '2026-07-19T10:00:00', + }, + { + id: 2, + agentId: 'operator', + toolPattern: 'shell', + pathPattern: '*', + behavior: 'deny', + source: 'system', + priority: 20, + enabled: false, + createdAt: '2026-07-19T10:05:00', + }, + ]), + })); + vi.stubGlobal('fetch', fetchMock); + + render( + ({ Authorization: 'Bearer test' })} + setNotice={vi.fn()} + /> + ); + + expect(await screen.findByText('权限规则中心')).toBeInTheDocument(); + expect(await screen.findByText('file_*')).toBeInTheDocument(); + expect(screen.getByText('shell')).toBeInTheDocument(); + expect(screen.queryByText('该模块已独立,等待配置项接入。')).not.toBeInTheDocument(); + expect(fetchMock).toHaveBeenCalledWith('/api/admin/permissions/rules', expect.any(Object)); + }); + + it('validates required tool pattern and priority range before creating', async () => { + const user = userEvent.setup(); + const fetchMock = vi.fn(async () => ({ ok: true, status: 200, json: async () => [] })); + vi.stubGlobal('fetch', fetchMock); + + render( ({})} setNotice={vi.fn()} />); + await screen.findByText('\u6682\u65e0\u6743\u9650\u89c4\u5219\u3002\u53ef\u4ee5\u5148\u6dfb\u52a0\u4e00\u6761 allow / ask / deny \u89c4\u5219\u3002'); + + const toolInput = screen.getByPlaceholderText('file_*'); + await user.clear(toolInput); + await user.click(screen.getByRole('button', { name: '\u521b\u5efa' })); + expect(await screen.findByRole('alert')).toHaveTextContent('\u5de5\u5177\u6a21\u5f0f\u4e0d\u80fd\u4e3a\u7a7a'); + expect(fetchMock).toHaveBeenCalledTimes(1); + + await user.type(toolInput, 'shell_*'); + const priorityInput = screen.getByRole('spinbutton'); + await user.clear(priorityInput); + await user.type(priorityInput, '10001'); + await user.click(screen.getByRole('button', { name: '\u521b\u5efa' })); + expect(await screen.findByRole('alert')).toHaveTextContent('-10000 \u5230 10000'); + expect(fetchMock).toHaveBeenCalledTimes(1); + }); + + it('edits, toggles and deletes an existing rule through the API', async () => { + const user = userEvent.setup(); + const rule = { + id: 7, + agentId: '*', + toolPattern: 'file_*', + pathPattern: '/workspace/**', + behavior: 'ask', + source: 'user', + priority: 10, + enabled: true, + createdAt: '2026-07-19T10:00:00', + }; + const fetchMock = vi.fn(async (_url: string, init?: RequestInit) => ({ + ok: true, + status: 200, + json: async () => init?.method ? { status: 'ok' } : [rule], + })); + vi.stubGlobal('fetch', fetchMock); + vi.spyOn(window, 'confirm').mockReturnValue(true); + + render( ({ Authorization: 'Bearer test' })} setNotice={vi.fn()} />); + await screen.findByText('file_*'); + + await user.click(screen.getByRole('button', { name: '\u7f16\u8f91\u6743\u9650\u89c4\u5219 file_*' })); + const editTool = screen.getByRole('textbox', { name: '\u7f16\u8f91\u5de5\u5177\u6a21\u5f0f' }); + await user.clear(editTool); + await user.type(editTool, 'shell_*'); + await user.selectOptions(screen.getByRole('combobox', { name: '\u7f16\u8f91\u6743\u9650\u52a8\u4f5c' }), 'deny'); + await user.click(screen.getByRole('button', { name: '\u4fdd\u5b58\u6743\u9650\u89c4\u5219' })); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledWith( + '/api/admin/permissions/rules/7', + expect.objectContaining({ + method: 'PUT', + body: JSON.stringify({ tool_pattern: 'shell_*', path_pattern: '/workspace/**', behavior: 'deny', priority: 10 }), + }), + )); + + await user.click(screen.getByRole('button', { name: '\u542f\u7528' })); + await waitFor(() => expect(fetchMock).toHaveBeenCalledWith( + '/api/admin/permissions/rules/7', + expect.objectContaining({ method: 'PUT', body: JSON.stringify({ enabled: false }) }), + )); + + await user.click(screen.getByRole('button', { name: '\u5220\u9664\u6743\u9650\u89c4\u5219 file_*' })); + await waitFor(() => expect(fetchMock).toHaveBeenCalledWith( + '/api/admin/permissions/rules/7', + expect.objectContaining({ method: 'DELETE' }), + )); + }); +}); diff --git a/frontend/__tests__/components/chat/dagReplay.test.tsx b/frontend/__tests__/components/chat/dagReplay.test.tsx index ed7a7e7..f431f12 100644 --- a/frontend/__tests__/components/chat/dagReplay.test.tsx +++ b/frontend/__tests__/components/chat/dagReplay.test.tsx @@ -1,7 +1,7 @@ import { beforeEach, describe, expect, it, vi } from 'vitest'; import { act, render, screen } from '@testing-library/react'; import DagModal from '../../../components/chat/DagModal'; -import { clearDagSession, useDagState } from '../../../lib/dagStore'; +import { clearDagSession, restoreDagState, setDagState, useDagState } from '../../../lib/dagStore'; import { handleSharedWebSocketEvent } from '../../../lib/websocketSharedEvents'; import type { ChatSession } from '../../../types'; @@ -123,4 +123,30 @@ describe('dag replay UI', () => { expect(screen.getByText('Session B')).toBeInTheDocument(); expect(screen.queryByText('Session A')).not.toBeInTheDocument(); }); + + it('keeps live progress when a slower recovery snapshot arrives', async () => { + setDagState('dag-session-a', { + total: 2, + completed: 1, + nodes: [ + { id: 'node-1', description: 'Recovered plan', status: 'SUCCESS' }, + { id: 'node-2', description: 'Recovered build', status: 'RUNNING' }, + ], + }); + restoreDagState('dag-session-a', { + total: 2, + completed: 0, + nodes: [ + { id: 'node-1', description: 'Recovered plan', status: 'PENDING' }, + { id: 'node-2', description: 'Recovered build', status: 'PENDING' }, + ], + }); + + render(); + + expect(screen.getByText('Recovered plan')).toBeInTheDocument(); + expect(screen.getByText('Recovered build')).toBeInTheDocument(); + expect(screen.getByText('1/2 \u8282\u70b9\u5b8c\u6210')).toBeInTheDocument(); + expect(screen.getByText('\u8fd0\u884c\u4e2d')).toBeInTheDocument(); + }); }); diff --git a/frontend/__tests__/components/flow/WorkflowConflictDialog.test.tsx b/frontend/__tests__/components/flow/WorkflowConflictDialog.test.tsx new file mode 100644 index 0000000..2d9d8a0 --- /dev/null +++ b/frontend/__tests__/components/flow/WorkflowConflictDialog.test.tsx @@ -0,0 +1,33 @@ +import { fireEvent, render, screen } from '@testing-library/react'; +import { describe, expect, it, vi } from 'vitest'; + +import { WorkflowConflictDialog } from '../../../components/flow/WorkflowConflictDialog'; +import { normalizeWorkflowDocument } from '../../../lib/workflowContract'; + +describe('WorkflowConflictDialog', () => { + it('shows graph comparison and exposes reload and overwrite actions', () => { + const onReload = vi.fn(); + const onOverwrite = vi.fn(); + render( + , + ); + + expect(screen.getByText(/当前编辑基于版本/)).toHaveTextContent('2'); + expect(screen.getByText(/名称:本地/)).toHaveTextContent('Local'); + fireEvent.click(screen.getByRole('button', { name: /载入服务器版本/ })); + fireEvent.click(screen.getByRole('button', { name: /以本地版本覆盖/ })); + expect(onReload).toHaveBeenCalledOnce(); + expect(onOverwrite).toHaveBeenCalledOnce(); + }); +}); diff --git a/frontend/__tests__/hooks/useSessionWebSocket.test.ts b/frontend/__tests__/hooks/useSessionWebSocket.test.ts index 499d1cf..800f70a 100644 --- a/frontend/__tests__/hooks/useSessionWebSocket.test.ts +++ b/frontend/__tests__/hooks/useSessionWebSocket.test.ts @@ -1,10 +1,87 @@ -import { describe, expect, it } from 'vitest'; -import { computeReconnectDelay } from '../../hooks/useSessionWebSocket'; +import { act, renderHook } from '@testing-library/react'; +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { buildSocketDedupKey, computeReconnectDelay, useSessionWebSocket } from '../../hooks/useSessionWebSocket'; + +class MockWebSocket { + static readonly CONNECTING = 0; + static readonly OPEN = 1; + static readonly CLOSING = 2; + static readonly CLOSED = 3; + static instances: MockWebSocket[] = []; + + readyState = MockWebSocket.CONNECTING; + sent: string[] = []; + onopen: (() => void) | null = null; + onclose: (() => void) | null = null; + onerror: (() => void) | null = null; + onmessage: ((event: MessageEvent) => void) | null = null; + + constructor(readonly url: string) { + MockWebSocket.instances.push(this); + } + + send(payload: string): void { this.sent.push(payload); } + close(): void { this.readyState = MockWebSocket.CLOSED; } + open(): void { this.readyState = MockWebSocket.OPEN; this.onopen?.(); } + disconnect(): void { this.readyState = MockWebSocket.CLOSED; this.onclose?.(); } + emit(payload: Record): void { + this.onmessage?.({ data: JSON.stringify(payload) } as MessageEvent); + } +} describe('useSessionWebSocket helpers', () => { + afterEach(() => { + vi.useRealTimers(); + vi.unstubAllGlobals(); + vi.restoreAllMocks(); + MockWebSocket.instances = []; + }); + it('computes capped reconnect backoff with deterministic jitter', () => { expect(computeReconnectDelay(0, 0)).toBe(1000); expect(computeReconnectDelay(3, 125)).toBe(8125); expect(computeReconnectDelay(99, 250)).toBe(30250); }); + + it('deduplicates replayable events without collapsing unsequenced stream chunks', () => { + expect(buildSocketDedupKey('task_preview', { messageId: 'preview-1' })).toBe('task_preview:preview-1'); + expect(buildSocketDedupKey('message_chunk', { messageId: 'stream-1' })).toBe(''); + expect(buildSocketDedupKey('message_chunk', { messageId: 'stream-1', sequence: 2 })).toBe('message_chunk:stream-1:2'); + }); + + it('uses the latest replayable event id as the reconnect cursor', async () => { + vi.useFakeTimers(); + vi.spyOn(Math, 'random').mockReturnValue(0); + vi.stubGlobal('WebSocket', MockWebSocket); + const wsRef = { current: new Map() }; + const currentSessionRef = { current: 'session-1' }; + const tokenRef = { current: 'token' }; + + const { result, unmount } = renderHook(() => useSessionWebSocket({ + wsRef, + currentSessionRef, + tokenRef, + setConnected: vi.fn(), + setNotice: vi.fn(), + addToast: vi.fn(), + onSessionClosed: vi.fn(), + onMessage: vi.fn(), + })); + + act(() => result.current.connectSession('session-1')); + const first = MockWebSocket.instances[0]; + act(() => first.open()); + act(() => first.emit({ event: 'task_preview', messageId: 'preview-cursor-1', sessionId: 'session-1' })); + act(() => first.disconnect()); + await act(async () => { vi.advanceTimersByTime(1000); }); + + const reconnected = MockWebSocket.instances[1]; + act(() => reconnected.open()); + expect(reconnected.sent.map((payload) => JSON.parse(payload))).toContainEqual({ + event: 'sync_request', + lastMessageId: 'preview-cursor-1', + }); + + unmount(); + }); }); diff --git a/frontend/__tests__/hooks/useWorkflowEditorSession.test.ts b/frontend/__tests__/hooks/useWorkflowEditorSession.test.ts new file mode 100644 index 0000000..03eba33 --- /dev/null +++ b/frontend/__tests__/hooks/useWorkflowEditorSession.test.ts @@ -0,0 +1,101 @@ +import { act, renderHook, waitFor } from '@testing-library/react'; +import { afterEach, describe, expect, it, vi } from 'vitest'; + +import { useWorkflowEditorSession } from '../../hooks/useWorkflowEditorSession'; + +function response(body: unknown, status = 200): Response { + return new Response(JSON.stringify(body), { status, headers: { 'Content-Type': 'application/json' } }); +} + +const serverWorkflow = { + id: 7, + name: 'Server workflow', + description: '', + triggerKeywords: [], + nodes: [{ id: 'agent', type: 'agent', name: 'Agent', description: '', x: 0, y: 0, agent: 'CodeGen', dependencies: [] }], + edges: [], + isDefault: false, + active: true, + version: 2, + schemaVersion: 1, +}; + +afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllGlobals(); +}); + +describe('useWorkflowEditorSession', () => { + it('recovers a persisted draft on refresh without mixing it with the server graph', async () => { + const draft = { + draftKey: 'workflow-7', workflowId: 7, baseVersion: 2, draftVersion: 3, + createdAt: '', updatedAt: '', + payload: { ...serverWorkflow, name: 'Recovered draft', nodes: [...serverWorkflow.nodes, { id: 'review', type: 'human', name: 'Review', description: '', x: 300, y: 0, dependencies: ['agent'], humanConfig: { prompt: 'Review' } }], edges: [{ id: 'agent->review', from: 'agent', to: 'review' }] }, + validation: { valid: true, normalized: {}, issues: [] }, + }; + vi.stubGlobal('fetch', vi.fn((input: RequestInfo | URL) => { + const url = String(input); + if (url.endsWith('/drafts/workflow-7')) return Promise.resolve(response(draft)); + if (url.endsWith('/workflows/7')) return Promise.resolve(response(serverWorkflow)); + throw new Error(`Unexpected request: ${url}`); + })); + + const { result } = renderHook(() => useWorkflowEditorSession(7)); + + await waitFor(() => expect(result.current.ready).toBe(true)); + expect(result.current.document.name).toBe('Recovered draft'); + expect(result.current.document.nodes.map((node) => node.id)).toEqual(['agent', 'review']); + expect(result.current.message).toBe('已恢复上次草稿'); + }); + + it('debounces graph changes and persists a versioned draft', async () => { + const fetchMock = vi.fn((input: RequestInfo | URL, init?: RequestInit) => { + const url = String(input); + if (url.endsWith('/drafts/new-workflow') && !init?.method) return Promise.resolve(response({}, 404)); + if (url.endsWith('/drafts/new-workflow') && init?.method === 'PUT') { + const payload = JSON.parse(String(init.body)); + return Promise.resolve(response({ + draftKey: 'new-workflow', baseVersion: 0, draftVersion: 1, + createdAt: '', updatedAt: '', payload, + validation: { valid: true, normalized: {}, issues: [] }, + })); + } + throw new Error(`Unexpected request: ${url}`); + }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useWorkflowEditorSession()); + await waitFor(() => expect(result.current.ready).toBe(true)); + + act(() => result.current.setDocument({ ...result.current.document, name: 'Autosaved' })); + + await waitFor(() => expect(result.current.status).toBe('saved'), { timeout: 2000 }); + const saveCall = fetchMock.mock.calls.find(([, init]) => init?.method === 'PUT'); + expect(saveCall).toBeTruthy(); + expect(JSON.parse(String(saveCall?.[1]?.body)).draftVersion).toBe(0); + }); + + it('opens a compare conflict when publish sees a stale workflow version', async () => { + let workflowReads = 0; + vi.stubGlobal('fetch', vi.fn((input: RequestInfo | URL, init?: RequestInit) => { + const url = String(input); + if (url.endsWith('/drafts/workflow-7')) return Promise.resolve(response({}, 404)); + if (url.endsWith('/validate')) return Promise.resolve(response({ valid: true, normalized: {}, issues: [] })); + if (url.endsWith('/workflows/7') && init?.method === 'PUT') { + return Promise.resolve(response({ detail: { code: 'workflow_version_conflict', message: 'conflict', expectedVersion: 2, currentVersion: 3 } }, 409)); + } + if (url.endsWith('/workflows/7')) { + workflowReads += 1; + return Promise.resolve(response({ ...serverWorkflow, name: workflowReads > 1 ? 'Remote update' : serverWorkflow.name, version: workflowReads > 1 ? 3 : 2 })); + } + throw new Error(`Unexpected request: ${url}`); + })); + const { result } = renderHook(() => useWorkflowEditorSession(7)); + await waitFor(() => expect(result.current.ready).toBe(true)); + + await act(async () => { await result.current.publish(); }); + + expect(result.current.conflict?.kind).toBe('workflow'); + expect(result.current.conflict?.currentVersion).toBe(3); + expect(result.current.conflict?.remote.name).toBe('Remote update'); + }); +}); diff --git a/frontend/__tests__/lib/dagStore.test.ts b/frontend/__tests__/lib/dagStore.test.ts index 79b0f07..5b82005 100644 --- a/frontend/__tests__/lib/dagStore.test.ts +++ b/frontend/__tests__/lib/dagStore.test.ts @@ -5,6 +5,8 @@ import { deriveDagStateFromMessages, getDagState, mergeDagTaskUpdate, + mergeRecoveredDagState, + selectLatestPersistedDagState, syncDagFromMessages, } from '../../lib/dagStore'; import type { Message, TaskPreviewEvent } from '../../types'; @@ -144,4 +146,37 @@ describe('dagStore helpers', () => { syncDagFromMessages('session-force-refresh', refreshed, true); expect(getDagState('session-force-refresh').nodes[0].id).toBe('node-new'); }); + + it('merges a persisted snapshot without regressing newer live progress', () => { + const snapshot = { + total: 2, + completed: 0, + nodes: [ + { id: 'node-1', status: 'PENDING', description: 'Plan' }, + { id: 'node-2', status: 'PENDING', description: 'Build' }, + ], + }; + const live = { + total: 2, + completed: 1, + nodes: [ + { id: 'node-1', status: 'SUCCESS', description: 'Plan' }, + { id: 'node-2', status: 'RUNNING', description: 'Build' }, + ], + }; + + expect(mergeRecoveredDagState(snapshot, live)).toMatchObject({ + completed: 1, + nodes: [{ status: 'SUCCESS' }, { status: 'RUNNING' }], + }); + }); + + it('selects the newest persisted task that contains a dag', () => { + const dag = { total: 1, completed: 1, nodes: [{ id: 'done', status: 'SUCCESS' }] }; + expect(selectLatestPersistedDagState([ + {}, + { dagProgress: dag }, + { dagProgress: { total: 1, completed: 0, nodes: [{ id: 'old' }] } }, + ])).toEqual(dag); + }); }); diff --git a/frontend/__tests__/lib/messageRecovery.test.ts b/frontend/__tests__/lib/messageRecovery.test.ts index 57864d8..d665f47 100644 --- a/frontend/__tests__/lib/messageRecovery.test.ts +++ b/frontend/__tests__/lib/messageRecovery.test.ts @@ -120,6 +120,23 @@ describe('messageRecovery helpers', () => { expect(merged[0].isStreaming).toBeFalsy(); }); + it('keeps the final snapshot when it reuses a streaming placeholder id', () => { + const merged = mergeReloadedMessages([ + { + event: 'message', sessionId: 's1', sender: 'agent', content: 'partial', type: 'text', + timestamp: '2026-07-19T00:00:00.000Z', messageId: 'shared-id', isStreaming: true, + }, + ], [ + { + event: 'message', sessionId: 's1', sender: 'agent', content: 'final answer', type: 'text', + timestamp: '2026-07-19T00:00:01.000Z', messageId: 'shared-id', isStreaming: false, + }, + ]); + + expect(merged).toHaveLength(1); + expect(merged[0]).toMatchObject({ content: 'final answer', messageId: 'shared-id', isStreaming: false }); + }); + it('replaces streaming placeholder with the final message payload', () => { const merged = mergeFinalMessage([ { diff --git a/frontend/__tests__/lib/workflowContract.test.ts b/frontend/__tests__/lib/workflowContract.test.ts new file mode 100644 index 0000000..8d834b1 --- /dev/null +++ b/frontend/__tests__/lib/workflowContract.test.ts @@ -0,0 +1,56 @@ +import { describe, expect, it } from 'vitest'; + +import { + fromReactFlowGraph, + normalizeWorkflowDocument, + toReactFlowEdges, + toReactFlowNodes, + workflowDiffSummary, +} from '../../lib/workflowContract'; + +describe('workflow ReactFlow contract adapter', () => { + it('round-trips explicit edge identity, labels, conditions, and dependencies', () => { + const document = normalizeWorkflowDocument({ + name: 'review', + version: 4, + schemaVersion: 1, + nodes: [ + { id: 'plan', type: 'agent', name: 'Plan', description: '', x: 10, y: 20, agent: 'Architect', dependencies: [] }, + { id: 'review', type: 'human', name: 'Review', description: '', x: 320, y: 20, dependencies: ['plan'], humanConfig: { prompt: 'Approve?' } }, + ], + edges: [{ id: 'approval-edge', from: 'plan', to: 'review', label: 'approve', condition: 'score > 0.8' }], + }); + + const nodes = toReactFlowNodes(document); + const edges = toReactFlowEdges(document); + const restored = fromReactFlowGraph(document, nodes, edges); + + expect(restored.edges).toEqual([ + { id: 'approval-edge', from: 'plan', to: 'review', label: 'approve', condition: 'score > 0.8' }, + ]); + expect(restored.nodes.find((node) => node.id === 'review')?.dependencies).toEqual(['plan']); + }); + + it('projects validation issues onto the matching node and edge', () => { + const document = normalizeWorkflowDocument({ + nodes: [{ id: 'a', type: 'agent', name: 'A', description: '', x: 0, y: 0, dependencies: [] }], + edges: [{ id: 'broken', from: 'a', to: 'missing' }], + }); + const issues = [ + { code: 'agent_unassigned', message: 'Agent missing', severity: 'warning' as const, nodeId: 'a' }, + { code: 'missing_edge_target', message: 'Target missing', severity: 'error' as const, edgeId: 'broken' }, + ]; + + expect(toReactFlowNodes(document, issues)[0].data.issues).toHaveLength(1); + const edge = toReactFlowEdges(document, issues)[0]; + expect(edge.data?.issues).toHaveLength(1); + expect(edge.style?.stroke).toBe('#DC2626'); + }); + + it('summarizes graph differences for conflict comparison', () => { + const local = normalizeWorkflowDocument({ name: 'local', nodes: [], edges: [] }); + const remote = normalizeWorkflowDocument({ name: 'remote', nodes: [], edges: [] }); + + expect(workflowDiffSummary(local, remote)[0]).toContain('local'); + }); +}); diff --git a/frontend/__tests__/pages/AdminPage.test.tsx b/frontend/__tests__/pages/AdminPage.test.tsx index 35cf929..17630a6 100644 --- a/frontend/__tests__/pages/AdminPage.test.tsx +++ b/frontend/__tests__/pages/AdminPage.test.tsx @@ -2,9 +2,15 @@ import { describe, it, expect, vi, beforeEach } from 'vitest'; import { render, screen, act } from '@testing-library/react'; import React from 'react'; -// ── All mock setup uses vi.hoisted() for vitest's hoist mechanism ────── - -const { createMockStore, mockAuthState, mockAdminState, mockAgentState, mockMemoryStoreState, mockUserMgmtState, mockWorkflowState } = vi.hoisted(() => { +const { + createMockStore, + mockAuthState, + mockAdminState, + mockAgentState, + mockMemoryStoreState, + mockUserMgmtState, + mockWorkflowState, +} = vi.hoisted(() => { const createMockStore = (defaultState: Record) => { return vi.fn((selector?: (s: unknown) => unknown) => { if (typeof selector === 'function') return selector(defaultState); @@ -19,111 +25,200 @@ const { createMockStore, mockAuthState, mockAdminState, mockAgentState, mockMemo setUser: mockVoid, setToken: mockVoid, authHeaders: vi.fn(() => ({ Authorization: 'Bearer test' })), - fmtErr: vi.fn((d: unknown, f: string) => (typeof d === 'string' ? d : f)), + fmtErr: vi.fn((detail: unknown, fallback: string) => (typeof detail === 'string' ? detail : fallback)), }; const mockAdminState = { - activeMenu: '服务商', + activeMenu: '\u670d\u52a1\u5546', notice: '', setActiveMenu: mockVoid, setNotice: mockVoid, }; - const mockAgentState = { - agents: [], agentTests: {}, adapterOptions: [], - selectedAdapterInfo: null, editSelectedAdapterInfo: null, - defaultChatAgent: 'Orchestrator', isCreatingAgent: false, - showLocalAgentModal: false, editingAgentId: null, - newAgent: { - agentId: '', domain: '', adapterType: 'deepseek', baseModelName: '', - rankLevel: 'L1', dutyNote: '', displayName: '', avatarUrl: '', - capabilityTags: [], baseUrl: '', apiKey: '', - systemPrompt: '', userPrompt: '', assistantPrompt: '', - promptVariables: {}, - publicConfig: { enabled: false, welcomeMessage: '', placeholder: '', themeColor: '#6366f1', logoUrl: '', suggestedQuestions: [] }, + const emptyAgent = { + agentId: '', + domain: '', + adapterType: 'deepseek', + baseModelName: '', + rankLevel: 'L1', + dutyNote: '', + displayName: '', + avatarUrl: '', + capabilityTags: [], + baseUrl: '', + apiKey: '', + systemPrompt: '', + userPrompt: '', + assistantPrompt: '', + promptVariables: {}, + publicConfig: { + enabled: false, + welcomeMessage: '', + placeholder: '', + themeColor: '#6366f1', + logoUrl: '', + suggestedQuestions: [], }, - editAgent: { - agentId: '', domain: '', adapterType: 'deepseek', baseModelName: '', - rankLevel: 'L1', dutyNote: '', displayName: '', avatarUrl: '', - capabilityTags: [], baseUrl: '', apiKey: '', - systemPrompt: '', userPrompt: '', assistantPrompt: '', - promptVariables: {}, - publicConfig: { enabled: false, welcomeMessage: '', placeholder: '', themeColor: '#6366f1', logoUrl: '', suggestedQuestions: [] }, - }, - fetchAdapters: mockVoid, refresh: mockVoid, createAgent: mockVoid, - testAgent: mockVoid, removeAgent: mockVoid, startEditAgent: mockVoid, - cancelEditAgent: mockVoid, saveAgentEdit: mockVoid, - handleSetDefaultChatAgent: mockVoid, handleAdapterChange: mockVoid, - setNewAgent: mockVoid, setEditAgent: mockVoid, - setSelectedAdapterInfo: mockVoid, setEditSelectedAdapterInfo: mockVoid, - setIsCreatingAgent: mockVoid, setShowLocalAgentModal: mockVoid, + }; + + const mockAgentState = { + agents: [], + agentTests: {}, + adapterOptions: [], + selectedAdapterInfo: null, + editSelectedAdapterInfo: null, + defaultChatAgent: 'Orchestrator', + isCreatingAgent: false, + showLocalAgentModal: false, + editingAgentId: null, + newAgent: emptyAgent, + editAgent: emptyAgent, + fetchAdapters: mockVoid, + refresh: mockVoid, + createAgent: mockVoid, + testAgent: mockVoid, + removeAgent: mockVoid, + startEditAgent: mockVoid, + cancelEditAgent: mockVoid, + saveAgentEdit: mockVoid, + handleSetDefaultChatAgent: mockVoid, + handleAdapterChange: mockVoid, + setNewAgent: mockVoid, + setEditAgent: mockVoid, + setSelectedAdapterInfo: mockVoid, + setEditSelectedAdapterInfo: mockVoid, + setIsCreatingAgent: mockVoid, + setShowLocalAgentModal: mockVoid, setEditingAgentId: mockVoid, }; const mockMemoryStoreState = { - init: mockVoid, loadMemoryFiles: mockVoid, loadMemoryDetail: mockVoid, - saveMemoryDetail: mockVoid, setMemoryKeyword: mockVoid, - setActiveMemoryFile: mockVoid, setMemoryBodyDraft: mockVoid, - setMemoryDirty: mockVoid, setMemoryPreview: mockVoid, - setMemorySubTab: mockVoid, setShowTrash: mockVoid, - setShowDeleteConfirm: mockVoid, setPendingDeleteFile: mockVoid, - setConsolidationDryRun: mockVoid, setMemorySearchQuery: mockVoid, - setMemorySearchResults: mockVoid, handleExportMemory: mockVoid, - handleImportMemory: mockVoid, confirmDeleteMemory: mockVoid, - handleDeleteMemory: mockVoid, loadTrash: mockVoid, - handleRecoverFromTrash: mockVoid, handlePurgeFromTrash: mockVoid, - loadSessionSummaries: mockVoid, loadSessionDetail: mockVoid, - loadGlobalSummary: mockVoid, refreshGlobalSummary: mockVoid, - runConsolidation: mockVoid, runMemorySearch: mockVoid, + init: mockVoid, + loadMemoryFiles: mockVoid, + loadMemoryDetail: mockVoid, + saveMemoryDetail: mockVoid, + setMemoryKeyword: mockVoid, + setActiveMemoryFile: mockVoid, + setMemoryBodyDraft: mockVoid, + setMemoryDirty: mockVoid, + setMemoryPreview: mockVoid, + setMemorySubTab: mockVoid, + setShowTrash: mockVoid, + setShowDeleteConfirm: mockVoid, + setPendingDeleteFile: mockVoid, + setConsolidationDryRun: mockVoid, + setMemorySearchQuery: mockVoid, + setMemorySearchResults: mockVoid, + handleExportMemory: mockVoid, + handleImportMemory: mockVoid, + confirmDeleteMemory: mockVoid, + handleDeleteMemory: mockVoid, + loadTrash: mockVoid, + handleRecoverFromTrash: mockVoid, + handlePurgeFromTrash: mockVoid, + loadSessionSummaries: mockVoid, + loadSessionDetail: mockVoid, + loadGlobalSummary: mockVoid, + refreshGlobalSummary: mockVoid, + runConsolidation: mockVoid, + runMemorySearch: mockVoid, getFilteredMemoryFiles: vi.fn(() => []), - loadSessionMemoryList: mockVoid, loadSessionMemoryConversation: mockVoid, - consolidateSessionMemory: mockVoid, createMemorySession: mockVoid, + loadSessionMemoryList: mockVoid, + loadSessionMemoryConversation: mockVoid, + consolidateSessionMemory: mockVoid, + createMemorySession: mockVoid, updateSessionTopic: mockVoid, - memoryLoading: false, memoryError: null, memoryKeyword: '', - memoryFiles: [], activeMemoryFile: null, memoryDetail: null, - memoryBodyDraft: '', memoryDirty: false, memoryPreview: null, - memorySubTab: 'files', sessionList: [], sessionsLoading: false, - activeSessionId: null, activeSessionSummary: null, - globalSummary: null, globalSummaryLoading: false, - consolidationLoading: false, consolidationResult: null, - consolidationError: null, consolidationDryRun: false, - memorySearchQuery: '', memorySearchResults: null, - memorySearchLoading: false, showTrash: false, trashItems: [], - trashLoading: false, showDeleteConfirm: false, pendingDeleteFile: null, - sessionMemoryList: [], sessionMemoryLoading: false, - activeSessionMemoryId: null, sessionMemoryConversation: [], + memoryLoading: false, + memoryError: null, + memoryKeyword: '', + memoryFiles: [], + activeMemoryFile: null, + memoryDetail: null, + memoryBodyDraft: '', + memoryDirty: false, + memoryPreview: null, + memorySubTab: 'files', + sessionList: [], + sessionsLoading: false, + activeSessionId: null, + activeSessionSummary: null, + globalSummary: null, + globalSummaryLoading: false, + consolidationLoading: false, + consolidationResult: null, + consolidationError: null, + consolidationDryRun: false, + memorySearchQuery: '', + memorySearchResults: null, + memorySearchLoading: false, + showTrash: false, + trashItems: [], + trashLoading: false, + showDeleteConfirm: false, + pendingDeleteFile: null, + sessionMemoryList: [], + sessionMemoryLoading: false, + activeSessionMemoryId: null, + sessionMemoryConversation: [], sessionMemoryConversationLoading: false, }; const mockUserMgmtState = { - ...mockMemoryStoreState, - tokenData: null, tokenLoading: false, tokenError: null, - profileBio: '', profileEditingField: null, profileFieldDraft: '', - profileLocation: '', profileEmail: '', profileOrg: '', - profileAvatarUrl: '', profileUploading: false, - userList: [], userListLoading: false, userListError: null, - newUserName: '', newUserPassword: '', newUserRole: 'user', + tokenData: null, + tokenLoading: false, + tokenError: '', + profileBio: '', + profileEditingField: null, + profileFieldDraft: '', + profileLocation: '', + profileEmail: '', + profileOrg: '', + profileAvatarUrl: '', + profileUploading: false, + userList: [], + userListLoading: false, + userListError: '', + newUserName: '', + newUserPassword: '', + newUserRole: 'developer', creatingUser: false, - setProfileEditingField: mockVoid, setProfileFieldDraft: mockVoid, - setNewUserName: mockVoid, setNewUserPassword: mockVoid, + setProfileEditingField: mockVoid, + setProfileFieldDraft: mockVoid, + setNewUserName: mockVoid, + setNewUserPassword: mockVoid, setNewUserRole: mockVoid, - handleStartEditField: mockVoid, handleSaveField: mockVoid, - handleCancelEditField: mockVoid, handleUploadProfileAvatar: mockVoid, - handleCreateUser: mockVoid, handleChangeUserRole: mockVoid, - handleDeleteUser: mockVoid, loadTokenUsage: mockVoid, + handleStartEditField: mockVoid, + handleSaveField: mockVoid, + handleCancelEditField: mockVoid, + handleUploadProfileAvatar: mockVoid, + handleCreateUser: mockVoid, + handleChangeUserRole: mockVoid, + handleDeleteUser: mockVoid, + loadTokenUsage: mockVoid, + init: mockVoid, }; const mockWorkflowState = { - loading: false, error: null, workflows: [], - loadWorkflows: mockVoid, deleteWorkflow: mockVoid, - setDefault: mockVoid, toggleActive: mockVoid, + loading: false, + error: null, + workflows: [], + loadWorkflows: mockVoid, + deleteWorkflow: mockVoid, + setDefault: mockVoid, + toggleActive: mockVoid, }; - return { createMockStore, mockAuthState, mockAdminState, mockAgentState, mockMemoryStoreState, mockUserMgmtState, mockWorkflowState }; + return { + createMockStore, + mockAuthState, + mockAdminState, + mockAgentState, + mockMemoryStoreState, + mockUserMgmtState, + mockWorkflowState, + }; }); -// ── vi.mock calls (hoisted above imports) ────────────────────────────── - vi.mock('next/navigation', () => ({ useRouter: () => ({ push: vi.fn(), replace: vi.fn(), back: vi.fn() }), useSearchParams: () => ({ get: vi.fn(() => null), has: vi.fn(() => false), toString: vi.fn(() => '') }), @@ -132,8 +227,7 @@ vi.mock('next/navigation', () => ({ })); vi.mock('next/dynamic', () => ({ - default: (_importFn: () => Promise, _opts?: unknown) => { - // Always return a simple placeholder — the real modules are too heavy for smoke tests. + default: () => { const Placeholder = () => React.createElement('div', { 'data-testid': 'dynamic-module' }, 'Module loaded'); Placeholder.displayName = 'DynamicMock'; return Placeholder; @@ -146,7 +240,14 @@ vi.mock('../../stores/authStore', () => ({ vi.mock('../../stores/adminStore', () => ({ useAdminStore: createMockStore(mockAdminState), - SETTINGS_MENU: ['服务商', '记忆', '技能', '通用', '审计日志', '用户管理', 'IM 接入', 'MCP', '工作流', '知识库', '模板市场', '工具市场', '工作空间', '上下文引擎', 'AgentNet', 'Agent 身份', 'Docker 沙箱', '多模态工作区', '集中日志', '模块连线', 'RAG 检索', '检索评估', 'A2A 互操作', 'A/B 测试', '成本分析', 'SLO 仪表板', '离线评估'], + SETTINGS_MENU: [ + '\u670d\u52a1\u5546', + '\u5de5\u4f5c\u6d41', + '\u6743\u9650', + '\u901a\u7528', + '\u8bb0\u5fc6', + '\u7528\u6237\u7ba1\u7406', + ], })); vi.mock('../../stores/agentStore', () => ({ @@ -165,13 +266,14 @@ vi.mock('../../stores/workflowStore', () => ({ useWorkflowStore: createMockStore(mockWorkflowState), })); -// Import page AFTER all mocks import AdminPage from '../../app/admin/page'; describe('AdminPage', () => { beforeEach(() => { vi.clearAllMocks(); window.localStorage.clear(); + mockAdminState.activeMenu = '\u670d\u52a1\u5546'; + mockAdminState.notice = ''; }); it('renders without crashing', async () => { @@ -187,4 +289,15 @@ describe('AdminPage', () => { }); expect(screen.queryByText(/warning/i)).toBeNull(); }); + + it('routes the permissions menu to a real module', async () => { + mockAdminState.activeMenu = '\u6743\u9650'; + + await act(async () => { + render(); + }); + + expect(screen.getByTestId('dynamic-module')).toBeInTheDocument(); + expect(screen.queryByText(/\u7b49\u5f85\u914d\u7f6e\u9879\u63a5\u5165/)).toBeNull(); + }); }); diff --git a/frontend/app/admin/layout.tsx b/frontend/app/admin/layout.tsx index ca6b0f6..cc3f677 100644 --- a/frontend/app/admin/layout.tsx +++ b/frontend/app/admin/layout.tsx @@ -57,20 +57,23 @@ export default function AdminLayout({ children }: { children: ReactNode }): JSX. const [sidebarWidthLive, setSidebarWidthLive] = useState(null); const sidebarW = sidebarCollapsed ? '64px' : `${sidebarWidthLive ?? sidebarWidth}px`; const consoleWidth = consoleCollapsed ? '0px' : '380px'; + const gridTemplateColumns = consoleCollapsed + ? `${sidebarW} minmax(0, 1fr)` + : `${sidebarW} minmax(0, 1fr) 4px var(--console-w, ${consoleWidth})`; // Auto-collapse on small screens useEffect(() => { const handleResize = () => { const w = window.innerWidth; - if (w < 1024) { + if (w < 1440) { setSidebarCollapsed(true); - } else if (w >= 1280) { + } else if (w >= 1600) { setSidebarCollapsed(false); } if (w >= 1024) { setMobileDrawerOpen(false); } - if (w < 1280) { + if (w < 1600) { setConsoleCollapsed(true); } }; @@ -196,7 +199,7 @@ export default function AdminLayout({ children }: { children: ReactNode }): JSX.
{/* ═══════════════════════════════════════════════════════ @@ -225,7 +228,7 @@ export default function AdminLayout({ children }: { children: ReactNode }): JSX.