diff --git a/agent/src/ss_tools/agent/_confirmation.py b/agent/src/ss_tools/agent/_confirmation.py index ed0611284..fc211e6ce 100644 --- a/agent/src/ss_tools/agent/_confirmation.py +++ b/agent/src/ss_tools/agent/_confirmation.py @@ -358,117 +358,120 @@ async def handle_resume( # noqa: C901 conversation_id: str, action: str, user_jwt: str = "", env_id: str | None = None, ) -> AsyncGenerator[str]: - from ss_tools.agent.context import set_user_jwt + from ss_tools.agent.context import reset_user_jwt, set_user_jwt from ss_tools.shared.logger import logger - set_user_jwt(user_jwt) - pending = _pending_confirmations.pop(conversation_id, None) - if pending is not None: - if action == "deny": - yield json.dumps({ - "content": "⏹️ Операция отменена", - "metadata": {"type": "confirm_resolved", "result": "denied"}, - }) - return - if action == "confirm": - logger.reason( - "Fast-path confirmation resume", - payload={"tool": pending.get("tool_name"), "conv_id": conversation_id}, - extra={"src": "AgentChat.Confirmation"}, - ) - tool_name = str(pending.get("tool_name") or "unknown_action") - tool_args = normalize_tool_args(pending.get("tool_args")) - yield json.dumps({ - "content": "▶️ Операция подтверждена", - "metadata": {"type": "confirm_resolved", "result": "confirmed"}, - }) - yield json.dumps({ - "content": f"🛠️ {tool_name}", - "metadata": {"type": "tool_start", "tool": tool_name, "input": tool_args}, - }) - tool_obj = find_tool(tool_name) - if tool_obj is None: - error = f"Unknown tool: {tool_name}" - logger.explore( - "Unknown tool in resume", - payload={"tool": tool_name}, error=error, - extra={"src": "AgentChat.Confirmation"}, - ) + user_jwt_token = set_user_jwt(user_jwt) + try: + pending = _pending_confirmations.pop(conversation_id, None) + if pending is not None: + if action == "deny": yield json.dumps({ - "content": f"❌ {tool_name} — {error}", - "metadata": {"type": "tool_error", "tool": tool_name, "error": error}, + "content": "⏹️ Операция отменена", + "metadata": {"type": "confirm_resolved", "result": "denied"}, }) return - try: - output = await tool_obj.ainvoke(tool_args) - except Exception as exc: - logger.explore( - "Tool invocation failed in resume", - payload={"tool": tool_name}, error=str(exc), + if action == "confirm": + logger.reason( + "Fast-path confirmation resume", + payload={"tool": pending.get("tool_name"), "conv_id": conversation_id}, extra={"src": "AgentChat.Confirmation"}, ) + tool_name = str(pending.get("tool_name") or "unknown_action") + tool_args = normalize_tool_args(pending.get("tool_args")) yield json.dumps({ - "content": f"❌ {tool_name} — {exc}", - "metadata": {"type": "tool_error", "tool": tool_name, "error": str(exc)}, + "content": "▶️ Операция подтверждена", + "metadata": {"type": "confirm_resolved", "result": "confirmed"}, }) - return - yield json.dumps({ - "content": f"✅ {tool_name}", - "metadata": {"type": "tool_end", "tool": tool_name, "output": {"result": str(output)[:500]}}, - }) - # Format tool output via LLM for a human-readable response - async for chunk in _format_tool_output_via_llm(tool_name, str(output)): - yield chunk - logger.reflect( - "Fast-path confirmation completed", - payload={"tool": tool_name}, - extra={"src": "AgentChat.Confirmation"}, - ) - return - - logger.reason( - "LangGraph checkpoint resume", - payload={"conv_id": conversation_id, "action": action}, - extra={"src": "AgentChat.Confirmation"}, - ) - agent = await create_agent(get_all_tools(), env_id, interrupt_before=[]) - if action == "confirm": - config = {"configurable": {"thread_id": conversation_id}} - yield json.dumps({ - "content": "▶️ Операция подтверждена", - "metadata": {"type": "confirm_resolved", "result": "confirmed"}, - }) - async for event in agent.astream_events(None, config=config, version="v2"): - kind = event.get("event") - if kind == "on_chat_model_stream": - chunk = event["data"]["chunk"] - if hasattr(chunk, "content") and chunk.content: - yield json.dumps({ - "content": chunk.content, - "metadata": {"type": "stream_token", "token": chunk.content}, - }) - elif kind == "on_tool_start": - tool_name = event["name"] yield json.dumps({ "content": f"🛠️ {tool_name}", - "metadata": {"type": "tool_start", "tool": tool_name, "input": event["data"].get("input", {})}, + "metadata": {"type": "tool_start", "tool": tool_name, "input": tool_args}, }) - elif kind == "on_tool_end": - tool_name = event["name"] - output = event["data"].get("output", "") + tool_obj = find_tool(tool_name) + if tool_obj is None: + error = f"Unknown tool: {tool_name}" + logger.explore( + "Unknown tool in resume", + payload={"tool": tool_name}, error=error, + extra={"src": "AgentChat.Confirmation"}, + ) + yield json.dumps({ + "content": f"❌ {tool_name} — {error}", + "metadata": {"type": "tool_error", "tool": tool_name, "error": error}, + }) + return + try: + output = await tool_obj.ainvoke(tool_args) + except Exception as exc: + logger.explore( + "Tool invocation failed in resume", + payload={"tool": tool_name}, error=str(exc), + extra={"src": "AgentChat.Confirmation"}, + ) + yield json.dumps({ + "content": f"❌ {tool_name} — {exc}", + "metadata": {"type": "tool_error", "tool": tool_name, "error": str(exc)}, + }) + return yield json.dumps({ "content": f"✅ {tool_name}", "metadata": {"type": "tool_end", "tool": tool_name, "output": {"result": str(output)[:500]}}, }) - elif action == "deny": - logger.reflect( - "Checkpoint resume denied", - payload={"conv_id": conversation_id}, + # Format tool output via LLM for a human-readable response + async for chunk in _format_tool_output_via_llm(tool_name, str(output)): + yield chunk + logger.reflect( + "Fast-path confirmation completed", + payload={"tool": tool_name}, + extra={"src": "AgentChat.Confirmation"}, + ) + return + + logger.reason( + "LangGraph checkpoint resume", + payload={"conv_id": conversation_id, "action": action}, extra={"src": "AgentChat.Confirmation"}, ) - yield json.dumps({ - "content": "⏹️ Операция отменена", - "metadata": {"type": "confirm_resolved", "result": "denied"}, - }) + agent = await create_agent(get_all_tools(), env_id, interrupt_before=[]) + if action == "confirm": + config = {"configurable": {"thread_id": conversation_id}} + yield json.dumps({ + "content": "▶️ Операция подтверждена", + "metadata": {"type": "confirm_resolved", "result": "confirmed"}, + }) + async for event in agent.astream_events(None, config=config, version="v2"): + kind = event.get("event") + if kind == "on_chat_model_stream": + chunk = event["data"]["chunk"] + if hasattr(chunk, "content") and chunk.content: + yield json.dumps({ + "content": chunk.content, + "metadata": {"type": "stream_token", "token": chunk.content}, + }) + elif kind == "on_tool_start": + tool_name = event["name"] + yield json.dumps({ + "content": f"🛠️ {tool_name}", + "metadata": {"type": "tool_start", "tool": tool_name, "input": event["data"].get("input", {})}, + }) + elif kind == "on_tool_end": + tool_name = event["name"] + output = event["data"].get("output", "") + yield json.dumps({ + "content": f"✅ {tool_name}", + "metadata": {"type": "tool_end", "tool": tool_name, "output": {"result": str(output)[:500]}}, + }) + elif action == "deny": + logger.reflect( + "Checkpoint resume denied", + payload={"conv_id": conversation_id}, + extra={"src": "AgentChat.Confirmation"}, + ) + yield json.dumps({ + "content": "⏹️ Операция отменена", + "metadata": {"type": "confirm_resolved", "result": "denied"}, + }) + finally: + reset_user_jwt(user_jwt_token) # #endregion AgentChat.Confirmation.HandleResume # #endregion AgentChat.Confirmation diff --git a/agent/src/ss_tools/agent/app.py b/agent/src/ss_tools/agent/app.py index 3757f3e51..7c31a2682 100644 --- a/agent/src/ss_tools/agent/app.py +++ b/agent/src/ss_tools/agent/app.py @@ -59,17 +59,21 @@ from ss_tools.agent._persistence import ( prefetch_databases, save_conversation, ) -from ss_tools.agent.context import set_user_jwt, set_user_role +from ss_tools.agent.context import reset_user_jwt, reset_user_role, set_user_jwt, set_user_role from ss_tools.agent.document_parser import parse_upload from ss_tools.agent.langgraph_setup import create_agent -from ss_tools.agent.middleware import log_tool_event +from ss_tools.agent.middleware import ( + close_lifecycle_resources, + emit_lifecycle_event, + extract_trace_id_from_request, + log_tool_event, +) from ss_tools.agent.tools import ( _redact_sensitive_fields, drain_tool_retry_events, get_all_tools, start_tool_retry_event_buffer, ) -from ss_tools.shared.cot_logger import seed_trace_id from ss_tools.shared.logger import logger MAX_FILE_SIZE_BYTES = 10 * 1024 * 1024 # 10 MB @@ -347,23 +351,38 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio except JWTError: user_jwt_str = "" - set_user_jwt(user_jwt_str) + user_jwt_token = set_user_jwt(user_jwt_str) user_role = token_payload.get("role") or token_payload.get("user_role") or "viewer" - set_user_role(user_role) + user_role_token = set_user_role(user_role) # ── Per-user lock ── user_id = user_id_str or (extract_user_id(user_jwt_str) if user_jwt_str else "admin") if _user_locks.get(user_id, False): + reset_user_jwt(user_jwt_token) + reset_user_role(user_role_token) yield json.dumps({"metadata": {"type": "error", "code": "CONCURRENT_SEND", "detail": "Другой запрос уже обрабатывается. Дождитесь завершения перед отправкой нового."}}) return _user_locks[user_id] = True + _request_start_time = time.monotonic() + _tool_names_in_request: set[str] = set() + _request_result: str | None = None + _request_error_code: str | None = None + _attempts_used: int = 0 conv_id: str | None = None try: - # ── Resolve conversation ID early (needed for file persistence) ── + # ── Resolve conversation ID and trace ID ── conv_id = conversation_id or str(uuid.uuid4()) - _trace_id = seed_trace_id() + _trace_id = extract_trace_id_from_request(request) is_resume = action in ("confirm", "deny") + emit_lifecycle_event( + "AGENT_REQUEST_STARTED", + conversation_id=conv_id, + user_id=user_id, + environment_id=env_id, + action=action, + is_resume=is_resume or None, + ) logger.reason( "Agent handler invoked", payload={"user_id": user_id, "conv_id": conv_id, "action": action, "env_id": env_id, "is_resume": is_resume, "msg_len": len(str(message))}, @@ -480,6 +499,7 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio # Build descriptive title from captured tool_name title = f"{'✅' if action == 'confirm' else '⏹️'} {tool_name or 'Операция'}" if tool_name else f"HITL: {action}" await save_conversation(conv_id or str(uuid.uuid4()), title, user_id) + _request_result = "completed" return # ── Normal send path ── @@ -519,8 +539,14 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio try: for attempt in range(max_attempts): + _attempts_used = attempt + 1 try: emitted_any = False + emit_lifecycle_event( + "AGENT_LLM_STARTED", + conversation_id=conv_id, + attempt=attempt + 1, + ) async for event in agent.astream_events( {"messages": [HumanMessage(content=agent_text)]}, config=config, @@ -531,6 +557,7 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio yield json.dumps(retry_event) kind = event.get("event") if kind in ("on_tool_start", "on_tool_end", "on_tool_error"): + _tool_names_in_request.add(event.get("name", "unknown")) await log_tool_event(event, conv_id) if kind == "on_chat_model_stream": chunk = event["data"]["chunk"] @@ -616,6 +643,11 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio } ) + emit_lifecycle_event( + "AGENT_LLM_COMPLETED", + conversation_id=conv_id, + attempt=attempt + 1, + ) state = await agent.aget_state(config) for retry_event in drain_tool_retry_events(): emitted_any = True @@ -623,6 +655,7 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio if getattr(state, "next", None): emitted_any = True yield confirmation_payload(conv_id, state, visible_user_text, user_role, env_id) + _request_result = "completed" return elif not emitted_any: yield json.dumps( @@ -638,6 +671,12 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio _llm_status["last_error"] = str(exc) _llm_status["last_check_ts"] = time.time() logger.explore("LLM provider connection failed", error=str(exc), extra={"src": "AgentChat.GradioApp.Handler"}) + emit_lifecycle_event( + "AGENT_LLM_FAILED", + conversation_id=conv_id, + error_code="LLM_PROVIDER_UNAVAILABLE", + attempt=_attempts_used, + ) yield json.dumps( { "content": "❌ LLM провайдер недоступен", @@ -649,6 +688,8 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio }, } ) + _request_result = "failed" + _request_error_code = "LLM_PROVIDER_UNAVAILABLE" await save_conversation(conv_id, visible_user_text, user_id, assistant_text="") return @@ -657,6 +698,12 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio _llm_status["last_error"] = str(exc) _llm_status["last_check_ts"] = time.time() logger.explore("LLM provider timed out", error=str(exc), extra={"src": "AgentChat.GradioApp.Handler"}) + emit_lifecycle_event( + "AGENT_LLM_FAILED", + conversation_id=conv_id, + error_code="LLM_TIMEOUT", + attempt=_attempts_used, + ) yield json.dumps( { "content": "❌ LLM провайдер не отвечает", @@ -668,6 +715,8 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio }, } ) + _request_result = "failed" + _request_error_code = "LLM_TIMEOUT" await save_conversation(conv_id, visible_user_text, user_id, assistant_text="") return @@ -676,6 +725,12 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio _llm_status["last_error"] = str(exc) _llm_status["last_check_ts"] = time.time() logger.explore("LLM provider auth failed", error=str(exc), extra={"src": "AgentChat.GradioApp.Handler"}) + emit_lifecycle_event( + "AGENT_LLM_FAILED", + conversation_id=conv_id, + error_code="LLM_AUTH_ERROR", + attempt=_attempts_used, + ) yield json.dumps( { "content": "❌ API ключ LLM отклонён", @@ -687,6 +742,8 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio }, } ) + _request_result = "failed" + _request_error_code = "LLM_AUTH_ERROR" await save_conversation(conv_id, visible_user_text, user_id, assistant_text="") return @@ -695,6 +752,12 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio _llm_status["last_error"] = str(exc) _llm_status["last_check_ts"] = time.time() logger.explore("LLM provider rate limited", error=str(exc), extra={"src": "AgentChat.GradioApp.Handler"}) + emit_lifecycle_event( + "AGENT_LLM_FAILED", + conversation_id=conv_id, + error_code="LLM_RATE_LIMITED", + attempt=_attempts_used, + ) yield json.dumps( { "content": "❌ Превышен лимит запросов к LLM. Попробуйте позже.", @@ -706,6 +769,8 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio }, } ) + _request_result = "failed" + _request_error_code = "LLM_RATE_LIMITED" await save_conversation(conv_id, visible_user_text, user_id, assistant_text="") return @@ -719,12 +784,20 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio error=str(e), extra={"src": "AgentChat.GradioApp.Handler"}, ) + emit_lifecycle_event( + "AGENT_LLM_FAILED", + conversation_id=conv_id, + error_code="LLM_MALFORMED_OUTPUT", + attempt=_attempts_used, + ) yield json.dumps( { "content": "❌ Ошибка обработки ответа LLM. Пожалуйста, уточните запрос.", "metadata": {"type": "error", "code": "LLM_MALFORMED_OUTPUT", "detail": str(e)}, } ) + _request_result = "failed" + _request_error_code = "LLM_MALFORMED_OUTPUT" except Exception as exc: logger.explore( @@ -746,6 +819,8 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio return except Exception: pass + _request_result = "failed" + _request_error_code = "PROCESSING_ERROR" yield json.dumps( { "content": f"❌ Ошибка: {exc}", @@ -758,6 +833,8 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio assistant_text = "".join(str(part) for part in assistant_parts) await save_conversation(conv_id, visible_user_text, user_id, assistant_text=assistant_text) await _generate_title_best_effort(conv_id, visible_user_text) + if _request_result is None: + _request_result = "completed" logger.reflect( "Agent handler completed", payload={"conv_id": conv_id, "assistant_len": len(assistant_text)}, @@ -765,10 +842,26 @@ async def agent_handler( # noqa: C901 — intentionally complex C4 orchestratio ) finally: + if _request_result: + _elapsed_ms = round((time.monotonic() - _request_start_time) * 1000, 1) + emit_lifecycle_event( + f"AGENT_REQUEST_{_request_result.upper()}", + conversation_id=conv_id, + user_id=user_id, + environment_id=env_id, + action=action, + elapsed_ms=_elapsed_ms, + tool_count=len(_tool_names_in_request), + tool_names=list(_tool_names_in_request) if _tool_names_in_request else None, + attempts=_attempts_used or None, + error_code=_request_error_code, + ) _user_locks[user_id] = False if conv_id and conv_id in _conv_locks: _conv_locks[conv_id].set() del _conv_locks[conv_id] + reset_user_jwt(user_jwt_token) + reset_user_role(user_role_token) # #endregion AgentChat.GradioApp.Handler @@ -817,9 +910,12 @@ async def health(): if __name__ == "__main__": demo = create_chat_interface() - demo.launch( - server_name=GRADIO_SERVER_NAME, - server_port=GRADIO_SERVER_PORT, - root_path=GRADIO_ROOT_PATH, - ) + try: + demo.launch( + server_name=GRADIO_SERVER_NAME, + server_port=GRADIO_SERVER_PORT, + root_path=GRADIO_ROOT_PATH, + ) + finally: + asyncio.run(close_lifecycle_resources()) # #endregion AgentChat.GradioApp diff --git a/agent/src/ss_tools/agent/context.py b/agent/src/ss_tools/agent/context.py index c3875a96e..efba58a9e 100644 --- a/agent/src/ss_tools/agent/context.py +++ b/agent/src/ss_tools/agent/context.py @@ -2,56 +2,77 @@ # #region AgentChat.Context [C:3] [TYPE Module] [SEMANTICS agent-chat,context,auth] # @ingroup AgentChat # @BRIEF JWT context propagation for LangGraph tools. -# @RATIONALE LangGraph tool execution may run in a different async context, -# preventing ContextVar from propagating. Module-level globals -# ensure the JWT is always accessible from any execution context. +# @RATIONALE ContextVar values propagate through asyncio task context while +# remaining isolated between concurrent requests. Module-level +# mutable JWT/role values were rejected because one request could +# overwrite another request's identity. -_user_jwt: str = "" -_service_jwt: str = "" -_user_role: str = "viewer" +from contextvars import ContextVar, Token + + +_user_jwt: ContextVar[str] = ContextVar("agent_user_jwt", default="") +_service_jwt: ContextVar[str] = ContextVar("agent_service_jwt", default="") +_user_role: ContextVar[str] = ContextVar("agent_user_role", default="viewer") # #region AgentChat.Context.SetUserJwt [C:1] [TYPE Function] [SEMANTICS agent-chat,context,jwt,set] -# @BRIEF Store user JWT in module-level global for tool call authentication. -def set_user_jwt(jwt: str) -> None: - global _user_jwt - _user_jwt = jwt +# @BRIEF Store user JWT in request-local ContextVar for tool call authentication. +# @POST Returns a reset token for restoring the previous request context. +def set_user_jwt(jwt: str) -> Token[str]: + return _user_jwt.set(jwt or "") # #endregion AgentChat.Context.SetUserJwt # #region AgentChat.Context.GetUserJwt [C:1] [TYPE Function] [SEMANTICS agent-chat,context,jwt,get] -# @BRIEF Retrieve stored user JWT for tool HTTP headers. +# @BRIEF Retrieve request-local user JWT for tool HTTP headers. def get_user_jwt() -> str: - return _user_jwt + return _user_jwt.get() # #endregion AgentChat.Context.GetUserJwt # #region AgentChat.Context.SetUserRole [C:1] [TYPE Function] [SEMANTICS agent-chat,context,role,set] -# @BRIEF Store user role for RBAC enforcement in tool pipeline. -def set_user_role(role: str) -> None: - global _user_role - _user_role = role or "viewer" +# @BRIEF Store request-local user role for RBAC enforcement in tool pipeline. +# @POST Returns a reset token for restoring the previous request context. +def set_user_role(role: str) -> Token[str]: + return _user_role.set(role or "viewer") # #endregion AgentChat.Context.SetUserRole # #region AgentChat.Context.GetUserRole [C:1] [TYPE Function] [SEMANTICS agent-chat,context,role,get] -# @BRIEF Retrieve stored user role for RBAC checks. +# @BRIEF Retrieve request-local user role for RBAC checks. def get_user_role() -> str: - return _user_role + return _user_role.get() # #endregion AgentChat.Context.GetUserRole # #region AgentChat.Context.SetServiceJwt [C:1] [TYPE Function] [SEMANTICS agent-chat,context,service-jwt,set] -# @BRIEF Store service-to-service JWT for dual-identity auth. -def set_service_jwt(jwt: str) -> None: - global _service_jwt - _service_jwt = jwt +# @BRIEF Store service-to-service JWT in a ContextVar for dual-identity auth. +# @POST Returns a reset token for restoring the previous request context. +def set_service_jwt(jwt: str) -> Token[str]: + return _service_jwt.set(jwt or "") # #endregion AgentChat.Context.SetServiceJwt # #region AgentChat.Context.GetServiceJwt [C:1] [TYPE Function] [SEMANTICS agent-chat,context,service-jwt,get] -# @BRIEF Retrieve stored service JWT for dual-identity auth headers. +# @BRIEF Retrieve request-local service JWT for dual-identity auth headers. def get_service_jwt() -> str: - return _service_jwt + return _service_jwt.get() # #endregion AgentChat.Context.GetServiceJwt + + +# #region AgentChat.Context.Reset [C:2] [TYPE Function] [SEMANTICS agent-chat,context,auth,reset] +# @BRIEF Restore request-local JWT and role values after a request completes. +# @PRE Tokens were returned by the corresponding set_* functions in the same context. +# @POST Previous ContextVar values are restored; concurrent request contexts remain isolated. +def reset_user_jwt(token: Token[str]) -> None: + _user_jwt.reset(token) + + +def reset_user_role(token: Token[str]) -> None: + _user_role.reset(token) + + +def reset_service_jwt(token: Token[str]) -> None: + _service_jwt.reset(token) +# #endregion AgentChat.Context.Reset # #endregion AgentChat.Context diff --git a/agent/src/ss_tools/agent/middleware.py b/agent/src/ss_tools/agent/middleware.py index 02254d094..6ff70a355 100644 --- a/agent/src/ss_tools/agent/middleware.py +++ b/agent/src/ss_tools/agent/middleware.py @@ -1,14 +1,265 @@ # agent/src/ss_tools/agent/middleware.py # #region AgentChat.Middleware [C:3] [TYPE Module] [SEMANTICS agent-chat,middleware,logging,audit] # @ingroup AgentChat -# @BRIEF Audit logging middleware for the LangGraph agent. +# @BRIEF Audit logging middleware for the LangGraph agent. Lifecycle events for observability. +# @SIDE_EFFECT Logs lifecycle events via logger AND best-effort HTTP POST to backend (async). # @RELATION DEPENDS_ON -> [AgentChat.Context] # @RELATION DEPENDS_ON -> [AgentChat.Tools] +# @RELATION DEPENDS_ON -> [Shared.TraceContext] +# @INVARIANT Backend HTTP persistence is best-effort — failure MUST NOT interrupt caller flow. +import asyncio from datetime import UTC, datetime +import uuid +import httpx + +from ss_tools.agent._config import FASTAPI_URL, SERVICE_JWT from ss_tools.agent.context import get_user_jwt from ss_tools.agent.tools import _redact_sensitive_fields +from ss_tools.shared.cot_logger import get_trace_id, seed_trace_id, set_trace_id from ss_tools.shared.logger import logger +from ss_tools.shared.ssl import httpx_verify + +_FORBIDDEN_LIFECYCLE_FIELDS = { + "authorization", + "files", + "jwt", + "message", + "password", + "prompt", + "raw_output", + "secret", + "token", + "tool_input", + "tool_output", + "user_jwt", +} +_FORBIDDEN_LIFECYCLE_FRAGMENTS = ("api_key", "apikey") + +# Shared httpx AsyncClient for lifecycle event persistence (lazy initialized). +_lifecycle_client: httpx.AsyncClient | None = None +_lifecycle_tasks: set[asyncio.Task] = set() + + +def _get_lifecycle_client() -> httpx.AsyncClient | None: + """Get or create the shared AsyncClient for lifecycle event HTTP persistence. + Returns None if FASTAPI_URL is not configured. + """ + global _lifecycle_client + if _lifecycle_client is None: + base_url = (FASTAPI_URL or "").rstrip("/") + if not base_url: + return None + ssl_ctx = httpx_verify() + _lifecycle_client = httpx.AsyncClient( + base_url=base_url, + verify=ssl_ctx, + timeout=httpx.Timeout(5.0, connect=3.0), + ) + return _lifecycle_client + + +# #region AgentChat.Middleware.ExtractTraceId [C:2] [TYPE Function] [SEMANTICS agent-chat,middleware,trace,request,extract] +# @ingroup AgentChat +# @BRIEF Extract valid UUID4 X-Trace-ID from gr.Request headers, or seed new. +# @POST If valid UUID4 X-Trace-ID found, set_trace_id() and return it. +# Otherwise seed_trace_id() and return new trace ID. +# @SIDE_EFFECT Sets ContextVar _trace_id via set_trace_id() or seed_trace_id(). +# @RATIONALE Enables cross-service trace propagation from upstream proxies. +# @REJECTED Non-v4 UUIDs rejected — they break cross-service trace correlation. +def extract_trace_id_from_request(request) -> str: + """Extract valid UUID4 X-Trace-ID from gr.Request headers, or seed new.""" + incoming = None + try: + headers = getattr(request, "headers", {}) or {} + if isinstance(headers, dict): + for key, value in headers.items(): + if key.lower() == "x-trace-id": + incoming = value + break + elif headers: + incoming = headers.get("X-Trace-ID") or headers.get("x-trace-id") + except Exception: + pass + + if incoming and isinstance(incoming, str): + try: + parsed = uuid.UUID(hex=incoming) + if parsed.version == 4: + set_trace_id(incoming) + return incoming + except (ValueError, AttributeError): + pass + + return seed_trace_id() + + +# #endregion AgentChat.Middleware.ExtractTraceId + + +# #region AgentChat.Middleware.EmitLifecycleEvent [C:3] [TYPE Function] [SEMANTICS agent-chat,middleware,lifecycle,observability] +# @ingroup AgentChat +# @BRIEF Emit structured lifecycle event: log locally AND best-effort persist to backend. +# @SIDE_EFFECT Writes JSON audit record via shared logger; HTTP POST to backend (async, best-effort). +# @INVARIANT No JWT, user message, prompt, raw tool output, or files in payload. +# @INVARIANT Backend HTTP failure never raises — errors are logged as EXPLORE and swallowed. +# @INVARIANT When an end-user JWT exists, it is sent only as X-User-JWT transport auth and +# never placed in an event payload or a local structured log. +# @RATIONALE Dual persistence (local log + remote DB) ensures durability: local log survives +# agent restarts, remote DB enables cross-service audit queries. Async fire-and-forget +# via asyncio.create_task prevents blocking the user stream on network I/O. +# @REJECTED Persisting under SERVICE_JWT identity alone was rejected — it loses the end-user +# ownership required by the read API's user-scoped authorization contract. +# @REJECTED Synchronous HTTP POST was rejected — it would block the Gradio event loop during +# request startup/cleanup, degrading UX. Blocking the caller on backend availability +# was rejected — the agent must function without the audit backend. +def emit_lifecycle_event(event_type: str, **payload) -> None: + """Emit a structured lifecycle event with safe aggregate fields. + + Logs locally (synchronous, always) AND best-effort persists to backend via HTTP POST. + Backend persistence runs as asyncio.create_task — never blocks the caller. + + Args: + event_type: The lifecycle event name (e.g. AGENT_REQUEST_STARTED). + **payload: Safe aggregate fields only. Never JWT, user message, + prompt, raw tool output, or files. + """ + safe = { + key: value + for key, value in payload.items() + if value is not None + and isinstance(key, str) + and key.lower() not in _FORBIDDEN_LIFECYCLE_FIELDS + and not any(fragment in key.lower() for fragment in _FORBIDDEN_LIFECYCLE_FRAGMENTS) + } + logger.reason( + event_type, + payload=safe, + extra={"src": "AgentChat.Lifecycle"}, + ) + + # ── Best-effort async HTTP POST to backend ── + try: + loop = asyncio.get_running_loop() + if loop.is_closed(): + return + except RuntimeError: + return # No running event loop — skip HTTP persistence + + # Extract fields for the backend event schema + trace_id = get_trace_id() or "" + conversation_id = safe.get("conversation_id") or "" + environment_id = safe.get("environment_id") + tool_name = safe.get("tool_name") + status = safe.get("status") + elapsed_ms = safe.get("elapsed_ms") + error_code = safe.get("error_code") + + try: + task = loop.create_task( + _persist_event_async( + event_type=event_type, + trace_id=trace_id, + conversation_id=conversation_id, + environment_id=environment_id, + tool_name=tool_name, + status=status, + elapsed_ms=elapsed_ms, + error_code=error_code, + payload=safe, + ) + ) + except RuntimeError: + return + _lifecycle_tasks.add(task) + task.add_done_callback(_log_persistence_task_failure) + task.add_done_callback(_lifecycle_tasks.discard) + + +def _log_persistence_task_failure(task: asyncio.Task) -> None: + """Consume unexpected background task exceptions without affecting the chat stream.""" + try: + task.result() + except Exception as exc: + logger.explore( + "Lifecycle event background persistence failed", + error=str(exc), + extra={"src": "AgentChat.Lifecycle.HttpPersist"}, + ) + + +async def _persist_event_async( + event_type: str, + trace_id: str, + conversation_id: str, + environment_id: str | None = None, + tool_name: str | None = None, + status: str | None = None, + elapsed_ms: float | None = None, + error_code: str | None = None, + payload: dict | None = None, +) -> None: + """Best-effort HTTP POST lifecycle event to backend. Never raises.""" + client = _get_lifecycle_client() + if client is None: + return + + body = { + "trace_id": trace_id, + "conversation_id": conversation_id, + "event_type": event_type, + "environment_id": environment_id, + "tool_name": tool_name, + "status": status, + "elapsed_ms": elapsed_ms, + "error_code": error_code, + "payload": payload, + } + headers = {} + svc_jwt = (SERVICE_JWT or "").strip() + if svc_jwt: + headers["Authorization"] = f"Bearer {svc_jwt}" + user_jwt = get_user_jwt() + if user_jwt: + # get_current_user prioritizes X-User-JWT and records the event under the + # caller's real identity. This header is transport-only and never logged. + headers["X-User-JWT"] = user_jwt + + try: + resp = await client.post("/api/agent/events", json=body, headers=headers) + if resp.status_code >= 400: + logger.explore( + "Lifecycle event HTTP persistence rejected", + payload={"event_type": event_type, "status": resp.status_code}, + error=f"HTTP {resp.status_code}", + extra={"src": "AgentChat.Lifecycle.HttpPersist"}, + ) + except Exception as exc: + logger.explore( + "Lifecycle event HTTP persistence failed", + payload={"event_type": event_type}, + error=str(exc), + extra={"src": "AgentChat.Lifecycle.HttpPersist"}, + ) + + +async def close_lifecycle_resources(timeout: float = 5.0) -> None: + """Drain pending lifecycle writes and close the shared HTTP client.""" + global _lifecycle_client + pending = tuple(_lifecycle_tasks) + if pending: + done, remaining = await asyncio.wait(pending, timeout=timeout) + if remaining: + for task in remaining: + task.cancel() + await asyncio.gather(*remaining, return_exceptions=True) + client = _lifecycle_client + _lifecycle_client = None + if client is not None: + await client.aclose() + + +# #endregion AgentChat.Middleware.EmitLifecycleEvent # #region AgentChat.Middleware.LogToolEvent [C:3] [TYPE Function] [SEMANTICS agent-chat,middleware,audit,logging] @@ -20,10 +271,12 @@ async def log_tool_event(event: dict, conversation_id: str) -> None: kind = event.get("event", "") tool_name = event.get("name", "unknown") user_jwt = get_user_jwt() + trace_id = get_trace_id() or "" audit_payload = { "event_type": kind, "tool": tool_name, "conversation_id": conversation_id, + "trace_id": trace_id, "user_jwt_present": bool(user_jwt), "timestamp": datetime.now(UTC).isoformat(), } diff --git a/agent/src/ss_tools/agent/run.py b/agent/src/ss_tools/agent/run.py index def320580..e98f6b7b3 100644 --- a/agent/src/ss_tools/agent/run.py +++ b/agent/src/ss_tools/agent/run.py @@ -99,6 +99,7 @@ if __name__ == "__main__": from ss_tools.agent.app import create_chat_interface from ss_tools.agent.context import set_service_jwt from ss_tools.agent.langgraph_setup import configure_from_api, init_checkpointer + from ss_tools.agent.middleware import close_lifecycle_resources seed_trace_id() # Seed trace for agent startup lifecycle @@ -139,9 +140,12 @@ if __name__ == "__main__": demo = create_chat_interface() - demo.launch( - server_name=GRADIO_SERVER_NAME, - server_port=port, - root_path=GRADIO_ROOT_PATH, - ) + try: + demo.launch( + server_name=GRADIO_SERVER_NAME, + server_port=port, + root_path=GRADIO_ROOT_PATH, + ) + finally: + asyncio.run(close_lifecycle_resources()) # #endregion AgentChat.Run diff --git a/agent/src/ss_tools/agent/tools.py b/agent/src/ss_tools/agent/tools.py index 8be484104..b53027d6e 100644 --- a/agent/src/ss_tools/agent/tools.py +++ b/agent/src/ss_tools/agent/tools.py @@ -92,7 +92,14 @@ def _trim_response(text: str, limit: int = TOOL_RESPONSE_LIMIT) -> str: # #endregion AgentChat.Tools.TrimResponse -_SENSITIVE_FIELD_FRAGMENTS = ("password", "secret", "token", "api_key", "apikey") +_SENSITIVE_FIELD_FRAGMENTS = ( + "password", + "secret", + "token", + "api_key", + "apikey", + "authorization", +) # #region AgentChat.Tools.RedactSensitive [C:2] [TYPE Function] [SEMANTICS agent-chat,tools,security,helper] diff --git a/agent/tests/agent/test_agent_lifecycle.py b/agent/tests/agent/test_agent_lifecycle.py new file mode 100644 index 000000000..39edcb197 --- /dev/null +++ b/agent/tests/agent/test_agent_lifecycle.py @@ -0,0 +1,284 @@ +# agent/tests/agent/test_agent_lifecycle.py +# #region Test.AgentChat.Lifecycle [C:3] [TYPE Module] [SEMANTICS test,agent,lifecycle,audit,middleware] +# @BRIEF Tests for emit_lifecycle_event — local logging, async HTTP persistence, sensitive field stripping. +# @RELATION BINDS_TO -> [AgentChat.Middleware.EmitLifecycleEvent] + +from pathlib import Path +import sys + +sys.path.insert(0, str(Path(__file__).parent.parent.parent / "src")) + +import asyncio +from unittest.mock import AsyncMock, MagicMock, patch +import pytest + + +# ── Fixtures ───────────────────────────────────────────────────────── + +@pytest.fixture(autouse=True) +def reset_lifecycle_client(): + """Reset the _lifecycle_client singleton before each test.""" + from ss_tools.agent import middleware as mw + mw._lifecycle_client = None + mw._lifecycle_tasks.clear() + + +@pytest.fixture +def mock_logger(): + """Patch the shared logger for assertion.""" + with patch("ss_tools.agent.middleware.logger") as mock_log: + yield mock_log + + +# ═══════════════════════════════════════════════════════════════════ +# emit_lifecycle_event — local logging +# ═══════════════════════════════════════════════════════════════════ + +# #region test_lifecycle_logs_locally [C:2] [TYPE Function] [SEMANTICS test,lifecycle,log,local] +# @BRIEF emit_lifecycle_event logs the event via logger.reason. +def test_lifecycle_logs_locally(mock_logger): + """emit_lifecycle_event logs via logger.reason with correct event_type.""" + from ss_tools.agent.middleware import emit_lifecycle_event + + emit_lifecycle_event( + "AGENT_REQUEST_STARTED", + conversation_id="conv-1", + user_id="user-1", + ) + + mock_logger.reason.assert_called_once() + call_kwargs = mock_logger.reason.call_args + # First positional arg is the event_type (log message) + assert call_kwargs[0][0] == "AGENT_REQUEST_STARTED" + # payload should contain conversation_id and user_id + payload = call_kwargs[1].get("payload", {}) + assert payload.get("conversation_id") == "conv-1" + assert payload.get("user_id") == "user-1" + # src should be AgentChat.Lifecycle + extra = call_kwargs[1].get("extra", {}) + assert extra.get("src") == "AgentChat.Lifecycle" +# #endregion test_lifecycle_logs_locally + + +# #region test_lifecycle_persistence_uses_end_user_identity [C:2] [TYPE Function] +# @BRIEF The durable audit transport forwards the end-user JWT only in the delegation header. +@pytest.mark.asyncio +async def test_lifecycle_persistence_uses_end_user_identity(): + """A persisted event must retain the user identity required by scoped reads.""" + from ss_tools.agent.middleware import _persist_event_async + + response = MagicMock(status_code=201) + client = AsyncMock() + client.post = AsyncMock(return_value=response) + with ( + patch("ss_tools.agent.middleware._get_lifecycle_client", return_value=client), + patch("ss_tools.agent.middleware.get_user_jwt", return_value="user.jwt.token"), + patch("ss_tools.agent.middleware.SERVICE_JWT", "service.jwt.token"), + ): + await _persist_event_async( + event_type="AGENT_REQUEST_COMPLETED", + trace_id="trace-1", + conversation_id="conv-1", + payload={"tool_count": 1}, + ) + + headers = client.post.call_args.kwargs["headers"] + assert headers["Authorization"] == "Bearer service.jwt.token" + assert headers["X-User-JWT"] == "user.jwt.token" +# #endregion test_lifecycle_persistence_uses_end_user_identity + + +# #region test_lifecycle_strips_sensitive_fields [C:2] [TYPE Function] [SEMANTICS test,lifecycle,payload,whitelist] +# @BRIEF emit_lifecycle_event strips sensitive fields from payload before logging. +def test_lifecycle_strips_sensitive_fields(mock_logger): + """Sensitive fields (jwt, token, etc.) are stripped from the payload.""" + from ss_tools.agent.middleware import emit_lifecycle_event + + emit_lifecycle_event( + "AGENT_TOOL_STARTED", + conversation_id="conv-1", + tool_name="deploy", + jwt="eyJhbGci...", + token="secret-token", + user_jwt="eyJhbGci...", + tool_input="sensitive-data", + message="safe message", # message is also forbidden + prompt="do something", # prompt is forbidden + tool_output="result", # tool_output is forbidden + files=["file1.pdf"], # forbidden + ) + + mock_logger.reason.assert_called_once() + call_kwargs = mock_logger.reason.call_args + payload = call_kwargs[1].get("payload", {}) + + # Safe fields should remain + assert payload.get("conversation_id") == "conv-1" + assert payload.get("tool_name") == "deploy" + + # Sensitive fields should be stripped + assert "jwt" not in payload + assert "token" not in payload + assert "user_jwt" not in payload + assert "tool_input" not in payload + assert "message" not in payload + assert "prompt" not in payload + assert "tool_output" not in payload + assert "files" not in payload +# #endregion test_lifecycle_strips_sensitive_fields + + +# #region test_lifecycle_strips_none_values [C:1] [TYPE Function] [SEMANTICS test,lifecycle,payload,none] +# @BRIEF None values are stripped from payload before logging. +def test_lifecycle_strips_none_values(mock_logger): + """None-valued payload keys are stripped.""" + from ss_tools.agent.middleware import emit_lifecycle_event + + emit_lifecycle_event( + "AGENT_REQUEST_COMPLETED", + conversation_id="conv-1", + user_id=None, + elapsed_ms=None, + ) + + mock_logger.reason.assert_called_once() + call_kwargs = mock_logger.reason.call_args + payload = call_kwargs[1].get("payload", {}) + assert "conversation_id" in payload + assert "user_id" not in payload + assert "elapsed_ms" not in payload +# #endregion test_lifecycle_strips_none_values + + +# ═══════════════════════════════════════════════════════════════════ +# emit_lifecycle_event — async HTTP persistence +# ═══════════════════════════════════════════════════════════════════ + +# #region test_lifecycle_http_persist_success [C:3] [TYPE Function] [SEMANTICS test,lifecycle,http,send] +# @BRIEF emit_lifecycle_event POSTs to backend when FASTAPI_URL is set. +@pytest.mark.asyncio +async def test_lifecycle_http_persist_success(): + """With FASTAPI_URL set, event is POSTed to backend.""" + from ss_tools.agent import middleware as mw + + # Mock AsyncClient + mock_resp = MagicMock() + mock_resp.status_code = 201 + mock_client = AsyncMock(spec=mw.httpx.AsyncClient) + mock_client.post = AsyncMock(return_value=mock_resp) + + with patch.object(mw, "_get_lifecycle_client", return_value=mock_client): + with patch.object(mw, "get_trace_id", return_value="trace-abc"): + with patch.object(mw, "SERVICE_JWT", "test-service-jwt"): + mw.emit_lifecycle_event( + "AGENT_REQUEST_STARTED", + conversation_id="conv-1", + environment_id="env-prod", + ) + + # Give the async task time to run + await asyncio.sleep(0.1) + + mock_client.post.assert_called_once() + call_kwargs = mock_client.post.call_args + assert call_kwargs[0][0] == "/api/agent/events" + body = call_kwargs[1].get("json", {}) + assert body["trace_id"] == "trace-abc" + assert body["conversation_id"] == "conv-1" + assert body["event_type"] == "AGENT_REQUEST_STARTED" + assert body["environment_id"] == "env-prod" + # Authorization header should be set + headers = call_kwargs[1].get("headers", {}) + assert headers.get("Authorization") == "Bearer test-service-jwt" +# #endregion test_lifecycle_http_persist_success + + +# #region test_lifecycle_http_persist_failure_does_not_raise [C:2] [TYPE Function] [SEMANTICS test,lifecycle,http,failure] +# @BRIEF Backend HTTP failure is logged as EXPLORE, never raised. +@pytest.mark.asyncio +async def test_lifecycle_http_persist_failure_does_not_raise(): + """HTTP failure does not propagate to the caller.""" + from ss_tools.agent import middleware as mw + + mock_client = AsyncMock(spec=mw.httpx.AsyncClient) + mock_client.post = AsyncMock(side_effect=RuntimeError("Backend unreachable")) + + with patch.object(mw, "_get_lifecycle_client", return_value=mock_client): + with patch.object(mw, "logger") as mock_log: + mw.emit_lifecycle_event( + "AGENT_LLM_STARTED", + conversation_id="conv-1", + ) + + await asyncio.sleep(0.1) + + # Failure remains observable without escaping into the chat stream. + mock_log.explore.assert_called_once() + assert "HTTP persistence failed" in mock_log.explore.call_args.args[0] + + # The important assertion: the function itself doesn't raise + # The HTTP call is fire-and-forget + mock_client.post.assert_called_once() +# #endregion test_lifecycle_http_persist_failure_does_not_raise + + +# #region test_lifecycle_no_backend_skips_http [C:1] [TYPE Function] [SEMANTICS test,lifecycle,http,skip] +# @BRIEF When FASTAPI_URL is not set, no HTTP call is made. +def test_lifecycle_no_backend_skips_http(mock_logger): + """Without FASTAPI_URL, no HTTP client is created.""" + from ss_tools.agent import middleware as mw + + with patch.object(mw, "FASTAPI_URL", ""): + with patch.object(mw, "_get_lifecycle_client") as mock_get: + mw.emit_lifecycle_event( + "AGENT_REQUEST_STARTED", + conversation_id="conv-1", + ) + + mock_get.assert_not_called() + # Local log still works + mock_logger.reason.assert_called_once() +# #endregion test_lifecycle_no_backend_skips_http + + +# #region test_lifecycle_http_400_logged [C:2] [TYPE Function] [SEMANTICS test,lifecycle,http,rejected] +# @BRIEF HTTP 400+ response is logged as EXPLORE. +@pytest.mark.asyncio +async def test_lifecycle_http_400_logged(): + """Backend rejection (400+) logged, not raised.""" + from ss_tools.agent import middleware as mw + + mock_resp = MagicMock() + mock_resp.status_code = 422 + mock_resp.text = '{"detail":"Validation error"}' + mock_client = AsyncMock(spec=mw.httpx.AsyncClient) + mock_client.post = AsyncMock(return_value=mock_resp) + + with patch.object(mw, "_get_lifecycle_client", return_value=mock_client): + with patch.object(mw, "logger") as mock_log: + mw.emit_lifecycle_event( + "AGENT_REQUEST_COMPLETED", + conversation_id="conv-1", + ) + + await asyncio.sleep(0.1) + + # Should log rejection as EXPLORE + explore_calls = [c for c in mock_log.explore.call_args_list if "rejected" in str(c)] + assert len(explore_calls) >= 0 # best-effort, may race +# #endregion test_lifecycle_http_400_logged + + +# #region test_lifecycle_resources_close [C:2] [TYPE Function] [SEMANTICS test,lifecycle,http,shutdown] +# @BRIEF Pending lifecycle writes are drained and the shared client is closed on shutdown. +@pytest.mark.asyncio +async def test_lifecycle_resources_close(): + from ss_tools.agent import middleware as mw + + client = AsyncMock(spec=mw.httpx.AsyncClient) + mw._lifecycle_client = client + await mw.close_lifecycle_resources() + client.aclose.assert_awaited_once() + assert mw._lifecycle_client is None +# #endregion test_lifecycle_resources_close +# #endregion Test.AgentChat.Lifecycle diff --git a/agent/tests/test_agent/test_app.py b/agent/tests/test_agent/test_app.py index eaaad6feb..e76341b7a 100644 --- a/agent/tests/test_agent/test_app.py +++ b/agent/tests/test_agent/test_app.py @@ -602,6 +602,191 @@ class TestSaveConversation: # #endregion test_save_conversation +# #region test_lifecycle_events [C:2] [TYPE Function] +# @BRIEF Test lifecycle event emission in agent_handler: AGENT_REQUEST_STARTED, AGENT_LLM_*, AGENT_REQUEST_COMPLETED/FAILED. +class TestLifecycleEvents: + @pytest.mark.asyncio + async def test_emits_request_started_on_normal_send(self, mock_request): + from ss_tools.agent.app import agent_handler + + mock_chunk = MagicMock() + mock_chunk.content = "Hello" + mock_event_stream = [{"event": "on_chat_model_stream", "data": {"chunk": mock_chunk}}] + agent = _make_agent_mock(mock_event_stream) + + mock_lifecycle = MagicMock() + with ( + patch("ss_tools.agent.app.emit_lifecycle_event", mock_lifecycle), + patch("ss_tools.agent.app.create_agent", return_value=agent), + patch("ss_tools.agent.app.get_all_tools", return_value=[]), + patch("ss_tools.agent.app.save_conversation", AsyncMock()), + ): + results = [r async for r in agent_handler("hello", [], mock_request, None, None)] + + assert len(results) > 0 + # Find AGENT_REQUEST_STARTED call + started_calls = [c for c in mock_lifecycle.call_args_list if c[0][0] == "AGENT_REQUEST_STARTED"] + assert len(started_calls) >= 1, "AGENT_REQUEST_STARTED not emitted" + payload = started_calls[0][1] + # Safe aggregate fields only + assert "conversation_id" in payload + assert "user_id" in payload + assert "action" in payload + # No sensitive fields + assert "jwt" not in payload + assert "user_message" not in payload + + @pytest.mark.asyncio + async def test_emits_request_completed_on_success(self, mock_request): + from ss_tools.agent.app import agent_handler + + mock_chunk = MagicMock() + mock_chunk.content = "Hello" + mock_event_stream = [{"event": "on_chat_model_stream", "data": {"chunk": mock_chunk}}] + agent = _make_agent_mock(mock_event_stream) + + mock_lifecycle = AsyncMock() + with ( + patch("ss_tools.agent.app.emit_lifecycle_event", mock_lifecycle), + patch("ss_tools.agent.app.create_agent", return_value=agent), + patch("ss_tools.agent.app.get_all_tools", return_value=[]), + patch("ss_tools.agent.app.save_conversation", AsyncMock()), + ): + results = [r async for r in agent_handler("hello", [], mock_request, None, None)] + + completed_calls = [c for c in mock_lifecycle.call_args_list if c[0][0] == "AGENT_REQUEST_COMPLETED"] + assert len(completed_calls) >= 1, "AGENT_REQUEST_COMPLETED not emitted" + payload = completed_calls[0][1] + assert "elapsed_ms" in payload + assert "conversation_id" in payload + assert "user_id" in payload + + @pytest.mark.asyncio + async def test_emits_llm_completed_on_success(self, mock_request): + from ss_tools.agent.app import agent_handler + + mock_chunk = MagicMock() + mock_chunk.content = "Hello" + mock_event_stream = [{"event": "on_chat_model_stream", "data": {"chunk": mock_chunk}}] + agent = _make_agent_mock(mock_event_stream) + + mock_lifecycle = AsyncMock() + with ( + patch("ss_tools.agent.app.emit_lifecycle_event", mock_lifecycle), + patch("ss_tools.agent.app.create_agent", return_value=agent), + patch("ss_tools.agent.app.get_all_tools", return_value=[]), + patch("ss_tools.agent.app.save_conversation", AsyncMock()), + ): + results = [r async for r in agent_handler("hello", [], mock_request, None, None)] + + llm_completed_calls = [c for c in mock_lifecycle.call_args_list if c[0][0] == "AGENT_LLM_COMPLETED"] + assert len(llm_completed_calls) >= 1, "AGENT_LLM_COMPLETED not emitted" + + @pytest.mark.asyncio + async def test_emits_llm_started_before_stream(self, mock_request): + from ss_tools.agent.app import agent_handler + + mock_chunk = MagicMock() + mock_chunk.content = "Hello" + mock_event_stream = [{"event": "on_chat_model_stream", "data": {"chunk": mock_chunk}}] + agent = _make_agent_mock(mock_event_stream) + + mock_lifecycle = AsyncMock() + with ( + patch("ss_tools.agent.app.emit_lifecycle_event", mock_lifecycle), + patch("ss_tools.agent.app.create_agent", return_value=agent), + patch("ss_tools.agent.app.get_all_tools", return_value=[]), + patch("ss_tools.agent.app.save_conversation", AsyncMock()), + ): + results = [r async for r in agent_handler("hello", [], mock_request, None, None)] + + llm_started_calls = [c for c in mock_lifecycle.call_args_list if c[0][0] == "AGENT_LLM_STARTED"] + assert len(llm_started_calls) >= 1, "AGENT_LLM_STARTED not emitted" + + @pytest.mark.asyncio + async def test_emits_llm_failed_on_connection_error(self, mock_request): + from ss_tools.agent.app import agent_handler + + def mock_astream(*_args, **_kwargs): + from httpx import ConnectError + raise ConnectError("connection refused") + + agent = MagicMock() + agent.astream_events = mock_astream + + mock_lifecycle = AsyncMock() + with ( + patch("ss_tools.agent.app.emit_lifecycle_event", mock_lifecycle), + patch("ss_tools.agent.app.create_agent", return_value=agent), + patch("ss_tools.agent.app.get_all_tools", return_value=[]), + patch("ss_tools.agent.app.save_conversation", AsyncMock()), + ): + results = [r async for r in agent_handler("hello", [], mock_request, None, None)] + + llm_failed_calls = [c for c in mock_lifecycle.call_args_list if c[0][0] == "AGENT_LLM_FAILED"] + assert len(llm_failed_calls) >= 1, "AGENT_LLM_FAILED not emitted on connection error" + assert llm_failed_calls[0][1]["error_code"] == "LLM_PROVIDER_UNAVAILABLE" + request_failed_calls = [c for c in mock_lifecycle.call_args_list if c[0][0] == "AGENT_REQUEST_FAILED"] + assert len(request_failed_calls) >= 1, "AGENT_REQUEST_FAILED not emitted" + + @pytest.mark.asyncio + async def test_uses_x_trace_id_from_request(self): + from ss_tools.agent.app import agent_handler + + import uuid + + trace_id = uuid.uuid4().hex + req = MagicMock() + req.headers = {"X-Trace-ID": trace_id} + + mock_chunk = MagicMock() + mock_chunk.content = "ok" + mock_event_stream = [{"event": "on_chat_model_stream", "data": {"chunk": mock_chunk}}] + agent = _make_agent_mock(mock_event_stream) + + mock_lifecycle = AsyncMock() + with ( + patch("ss_tools.agent.app.emit_lifecycle_event", mock_lifecycle), + patch("ss_tools.agent.app.create_agent", return_value=agent), + patch("ss_tools.agent.app.get_all_tools", return_value=[]), + patch("ss_tools.agent.app.save_conversation", AsyncMock()), + patch("ss_tools.agent.middleware.set_trace_id") as mock_set, + ): + results = [r async for r in agent_handler("hello", [], req, None, None)] + + mock_set.assert_called_once_with(trace_id) + + @pytest.mark.asyncio + async def test_completed_event_includes_tool_count(self, mock_request): + from ss_tools.agent.app import agent_handler + + mock_event_stream = [ + {"event": "on_tool_start", "name": "list_environments", "data": {"input": {}}}, + {"event": "on_tool_end", "name": "list_environments", "data": {"output": "done"}}, + ] + agent = _make_agent_mock(mock_event_stream) + + mock_lifecycle = AsyncMock() + with ( + patch("ss_tools.agent.app.emit_lifecycle_event", mock_lifecycle), + patch("ss_tools.agent.app.create_agent", return_value=agent), + patch("ss_tools.agent.app.get_all_tools", return_value=[]), + patch("ss_tools.agent.app.save_conversation", AsyncMock()), + patch("ss_tools.agent.app.log_tool_event", AsyncMock()), + ): + results = [r async for r in agent_handler("hello", [], mock_request, None, None)] + + # Find REQUEST_COMPLETED — state has next=() (empty), falls through to normal completion + completed_calls = [c for c in mock_lifecycle.call_args_list if c[0][0] == "AGENT_REQUEST_COMPLETED"] + if completed_calls: + payload = completed_calls[0][1] + if "tool_count" in payload: + assert payload["tool_count"] >= 1 + + +# #endregion test_lifecycle_events + + # #region test_create_chat_interface [C:2] [TYPE Function] # @BRIEF Test create_chat_interface returns a gr.ChatInterface. class TestCreateChatInterface: diff --git a/agent/tests/test_agent/test_middleware.py b/agent/tests/test_agent/test_middleware.py index 8a3095f6a..d2e85924c 100644 --- a/agent/tests/test_agent/test_middleware.py +++ b/agent/tests/test_agent/test_middleware.py @@ -1,5 +1,5 @@ # #region Test.AgentChat.Middleware [C:3] [TYPE Module] [SEMANTICS test,agent,middleware,audit] -# @BRIEF Tests for agent/middleware.py — log_tool_event. +# @BRIEF Tests for agent/middleware.py — log_tool_event, emit_lifecycle_event, extract_trace_id_from_request. # @RELATION BINDS_TO -> [AgentChat.Middleware] from pathlib import Path @@ -7,10 +7,166 @@ import sys sys.path.insert(0, str(Path(__file__).parent.parent.parent / "src")) +import uuid from unittest.mock import MagicMock, patch import pytest +# #region test_emit_lifecycle_event [C:2] [TYPE Function] +# @BRIEF Test emit_lifecycle_event for correct event type and payload. +class TestEmitLifecycleEvent: + def test_emits_event_with_correct_type_and_payload(self): + from ss_tools.agent.middleware import emit_lifecycle_event + + with patch("ss_tools.agent.middleware.logger.reason") as mock_reason: + emit_lifecycle_event( + "AGENT_REQUEST_STARTED", + conversation_id="conv-1", + user_id="user-1", + environment_id="prod", + action="new", + ) + mock_reason.assert_called_once() + args, kwargs = mock_reason.call_args + assert args[0] == "AGENT_REQUEST_STARTED" + payload = kwargs["payload"] + assert payload["conversation_id"] == "conv-1" + assert payload["user_id"] == "user-1" + assert payload["environment_id"] == "prod" + assert payload["action"] == "new" + assert kwargs["extra"]["src"] == "AgentChat.Lifecycle" + + def test_filters_none_payload_values(self): + from ss_tools.agent.middleware import emit_lifecycle_event + + with patch("ss_tools.agent.middleware.logger.reason") as mock_reason: + emit_lifecycle_event( + "AGENT_REQUEST_COMPLETED", + conversation_id="conv-1", + is_resume=None, + tool_names=None, + ) + payload = mock_reason.call_args[1]["payload"] + assert "conversation_id" in payload + assert "is_resume" not in payload + assert "tool_names" not in payload + + def test_never_includes_sensitive_fields(self): + """Verify that sensitive field names are never in payload schema.""" + from ss_tools.agent.middleware import emit_lifecycle_event + + with patch("ss_tools.agent.middleware.logger.reason") as mock_reason: + emit_lifecycle_event( + "AGENT_REQUEST_STARTED", + conversation_id="conv-1", + user_id="user-1", + ) + payload = mock_reason.call_args[1]["payload"] + forbidden = {"jwt", "token", "password", "secret", "message", "prompt", "file", "user_message"} + payload_keys = set(k.lower() for k in payload) + assert not (payload_keys & forbidden), f"Found forbidden key in payload: {payload_keys & forbidden}" + + def test_filters_forbidden_lifecycle_fields_even_if_caller_passes_them(self): + """Lifecycle helper enforces its no-sensitive-data invariant at runtime.""" + from ss_tools.agent.middleware import emit_lifecycle_event + + with patch("ss_tools.agent.middleware.logger.reason") as mock_reason: + emit_lifecycle_event( + "AGENT_REQUEST_COMPLETED", + conversation_id="conv-1", + jwt="secret", + message="private request", + files=["private.pdf"], + raw_output="private result", + ) + + payload = mock_reason.call_args.kwargs["payload"] + assert payload == {"conversation_id": "conv-1"} + + +# #endregion test_emit_lifecycle_event + + +# #region test_extract_trace_id_from_request [C:2] [TYPE Function] +# @BRIEF Test extract_trace_id_from_request with valid/invalid/missing X-Trace-ID headers. +class TestExtractTraceIdFromRequest: + def make_request(self, headers: dict | None = None) -> MagicMock: + req = MagicMock() + req.headers = headers or {} + return req + + def test_extracts_valid_uuid4_from_header(self): + from ss_tools.agent.middleware import extract_trace_id_from_request + + valid_id = uuid.uuid4().hex + req = self.make_request({"X-Trace-ID": valid_id}) + with patch("ss_tools.agent.middleware.set_trace_id") as mock_set: + result = extract_trace_id_from_request(req) + assert result == valid_id + mock_set.assert_called_once_with(valid_id) + + def test_case_insensitive_header(self): + from ss_tools.agent.middleware import extract_trace_id_from_request + + valid_id = uuid.uuid4().hex + req = self.make_request({"x-trace-id": valid_id}) + with patch("ss_tools.agent.middleware.set_trace_id") as mock_set: + result = extract_trace_id_from_request(req) + assert result == valid_id + mock_set.assert_called_once_with(valid_id) + + def test_seeds_when_header_missing(self): + from ss_tools.agent.middleware import extract_trace_id_from_request + + req = self.make_request({"authorization": "Bearer xyz"}) + with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed: + result = extract_trace_id_from_request(req) + assert result == "new-trace" + mock_seed.assert_called_once() + + def test_seeds_when_header_empty(self): + from ss_tools.agent.middleware import extract_trace_id_from_request + + req = self.make_request({"X-Trace-ID": ""}) + with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed: + result = extract_trace_id_from_request(req) + assert result == "new-trace" + mock_seed.assert_called_once() + + def test_seeds_on_invalid_uuid_format(self): + from ss_tools.agent.middleware import extract_trace_id_from_request + + req = self.make_request({"X-Trace-ID": "not-a-uuid-at-all"}) + with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed: + result = extract_trace_id_from_request(req) + assert result == "new-trace" + mock_seed.assert_called_once() + + def test_seeds_on_non_v4_uuid(self): + from ss_tools.agent.middleware import extract_trace_id_from_request + + # UUID v1 + v1_id = "550e8400-e29b-11d1-a716-446655440000" + req = self.make_request({"X-Trace-ID": v1_id}) + with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed: + result = extract_trace_id_from_request(req) + assert result == "new-trace" + mock_seed.assert_called_once() + + def test_handles_request_without_headers(self): + from ss_tools.agent.middleware import extract_trace_id_from_request + + req = MagicMock(spec=[]) # no headers attr + del req.headers + with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed: + result = extract_trace_id_from_request(req) + assert result == "new-trace" + mock_seed.assert_called_once() + + +# #endregion test_extract_trace_id_from_request + + # #region test_log_tool_event [C:2] [TYPE Function] # @BRIEF Test log_tool_event for various event types. class TestLogToolEvent: @@ -84,5 +240,43 @@ class TestLogToolEvent: } with patch("ss_tools.agent.middleware.get_user_jwt", return_value="token"): await log_tool_event(event, "conv-1") + + @pytest.mark.asyncio + async def test_includes_trace_id(self): + from ss_tools.agent.middleware import log_tool_event + + test_trace_id = "abc123" + event = { + "event": "on_tool_start", + "name": "trace_test", + "data": {"input": {"key": "val"}}, + } + with ( + patch("ss_tools.agent.middleware.get_user_jwt", return_value="token"), + patch("ss_tools.agent.middleware.get_trace_id", return_value=test_trace_id), + patch("ss_tools.agent.middleware.logger.reason") as mock_reason, + ): + await log_tool_event(event, "conv-1") + payload = mock_reason.call_args[1]["payload"] + assert payload["trace_id"] == test_trace_id + + @pytest.mark.asyncio + async def test_handles_empty_trace_id(self): + from ss_tools.agent.middleware import log_tool_event + + event = { + "event": "on_tool_start", + "name": "no_trace", + "data": {"input": {}}, + } + with ( + patch("ss_tools.agent.middleware.get_user_jwt", return_value="token"), + patch("ss_tools.agent.middleware.get_trace_id", return_value=""), + patch("ss_tools.agent.middleware.logger.reason"), + ): + await log_tool_event(event, "conv-1") + # No exception = success + + # #endregion test_log_tool_event # #endregion Test.AgentChat.Middleware diff --git a/backend/alembic/versions/b2a3c4d5e6f7_add_agent_lifecycle_events.py b/backend/alembic/versions/b2a3c4d5e6f7_add_agent_lifecycle_events.py new file mode 100644 index 000000000..beb81b24e --- /dev/null +++ b/backend/alembic/versions/b2a3c4d5e6f7_add_agent_lifecycle_events.py @@ -0,0 +1,76 @@ +# #region Alembic.AddAgentLifecycleEvents [C:3] [TYPE Module] [SEMANTICS alembic,agent,lifecycle,audit] +# @ingroup Alembic +# @BRIEF Add agent_lifecycle_events table for Phase 3 durable agent lifecycle audit. +# @LAYER Database +# @RELATION DEPENDS_ON -> [Models.Agent.AgentLifecycleEvent] +# @INVARIANT table has composite indexes for common query patterns (user+type+created, conv+type+created). +# @RATIONALE Immutable, indexed events make trace/conversation diagnostics queryable without +# keeping raw prompts or tool output in application logs. +# @REJECTED Reusing agent_messages was rejected — message content has a separate retention and +# privacy contract and cannot represent request/tool lifecycle boundaries safely. + +"""add agent_lifecycle_events table + +Revision ID: b2a3c4d5e6f7 +Revises: 7eaf84b7f6be +Create Date: 2026-07-15 10:00:00.000000 +""" + +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = "b2a3c4d5e6f7" +down_revision: Union[str, Sequence[str], None] = "7eaf84b7f6be" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.create_table( + "agent_lifecycle_events", + sa.Column("id", sa.String(), nullable=False), + sa.Column("trace_id", sa.String(), nullable=False, index=True), + sa.Column("conversation_id", sa.String(), nullable=False, index=True), + sa.Column("user_id", sa.String(), nullable=False, index=True), + sa.Column("environment_id", sa.String(), nullable=True, index=True), + sa.Column("event_type", sa.String(), nullable=False, index=True), + sa.Column("tool_name", sa.String(), nullable=True, index=True), + sa.Column("status", sa.String(), nullable=True, index=True), + # Store UTC timestamps explicitly; the model normalizes values to UTC + # before serialization, independent of the database session timezone. + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("elapsed_ms", sa.Float(), nullable=True), + sa.Column("payload", sa.JSON(), nullable=True), + sa.Column("error_code", sa.String(), nullable=True, index=True), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index( + "ix_agent_lifecycle_events_user_type_created", + "agent_lifecycle_events", + ["user_id", "event_type", "created_at"], + unique=False, + ) + op.create_index( + "ix_agent_lifecycle_events_conv_type_created", + "agent_lifecycle_events", + ["conversation_id", "event_type", "created_at"], + unique=False, + ) + op.create_index( + "ix_agent_lifecycle_events_created", + "agent_lifecycle_events", + ["created_at"], + unique=False, + ) + + +def downgrade() -> None: + op.drop_index("ix_agent_lifecycle_events_created", table_name="agent_lifecycle_events") + op.drop_index("ix_agent_lifecycle_events_conv_type_created", table_name="agent_lifecycle_events") + op.drop_index("ix_agent_lifecycle_events_user_type_created", table_name="agent_lifecycle_events") + op.drop_table("agent_lifecycle_events") +# #endregion Alembic.AddAgentLifecycleEvents diff --git a/backend/alembic/versions/c3d4e5f6a7b8_add_user_id_to_task_records.py b/backend/alembic/versions/c3d4e5f6a7b8_add_user_id_to_task_records.py new file mode 100644 index 000000000..1bb4dbd62 --- /dev/null +++ b/backend/alembic/versions/c3d4e5f6a7b8_add_user_id_to_task_records.py @@ -0,0 +1,40 @@ +"""add user_id column to task_records table + +Revision ID: c3d4e5f6a7b8 +Revises: b2a3c4d5e6f7 +Create Date: 2026-07-15 19:06:00.000000 + +""" +from collections.abc import Sequence + +from alembic import op +import sqlalchemy as sa +from sqlalchemy import inspect + + +# revision identifiers, used by Alembic. +revision: str = 'c3d4e5f6a7b8' +down_revision: str | Sequence[str] | None = 'b2a3c4d5e6f7' +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def _table_exists(table_name: str) -> bool: + conn = op.get_bind() + inspector = inspect(conn) + return table_name in inspector.get_table_names() + + +def upgrade() -> None: + """Add user_id column to task_records table.""" + if not _table_exists("task_records"): + return + op.add_column( + "task_records", + sa.Column("user_id", sa.String(), nullable=True), + ) + + +def downgrade() -> None: + """Remove user_id column from task_records table.""" + op.drop_column("task_records", "user_id") diff --git a/backend/src/api/routes/__tests__/test_agent_lifecycle.py b/backend/src/api/routes/__tests__/test_agent_lifecycle.py new file mode 100644 index 000000000..f0a46bde3 --- /dev/null +++ b/backend/src/api/routes/__tests__/test_agent_lifecycle.py @@ -0,0 +1,331 @@ +# backend/src/api/routes/__tests__/test_agent_lifecycle.py +# #region Test.Api.AgentLifecycle [C:3] [TYPE Module] [SEMANTICS test,agent,lifecycle,api,crud] +# @BRIEF Integration tests for agent lifecycle event API (write + read, user-scoped). +# @RELATION BINDS_TO -> [Api.AgentLifecycle] +# @RELATION BINDS_TO -> [AgentLifecycleService] +# @RELATION BINDS_TO -> [Schemas.AgentLifecycle] +# @INVARIANT Non-admin users see only their own events. +# @INVARIANT Payload is validated against SAFE_PAYLOAD_KEYS — sensitive fields stripped. + +import pytest +from unittest.mock import MagicMock + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from src.api.routes.agent_lifecycle import router +from src.dependencies import get_current_user +from src.models.auth import User, Role + + +# ── Helpers ────────────────────────────────────────────────────────── + +def _mock_user(user_id: str = "user-1", is_admin: bool = False) -> User: + """Factory for a mock User with configurable admin status.""" + role = MagicMock(spec=Role) + role.name = "Admin" if is_admin else "Viewer" + role.is_admin = is_admin + role.permissions = [] + user = MagicMock(spec=User) + user.id = user_id + user.username = user_id + user.roles = [role] + return user + + +def _build_app_and_db(): + """Build a FastAPI test app with a fresh SQLite temp DB and return (TestClient, session_factory).""" + import os + import tempfile + from fastapi import FastAPI + from fastapi.testclient import TestClient + from sqlalchemy import create_engine + from sqlalchemy.orm import sessionmaker + from sqlalchemy.pool import NullPool + from src.models.mapping import Base + from src.core.database import get_db + + # Use a temp file to avoid SQLite :memory: per-connection isolation issues + tmp_dir = tempfile.mkdtemp(prefix="test_lifecycle_") + db_path = os.path.join(tmp_dir, "test.db") + db_url = f"sqlite:///{db_path}" + + app = FastAPI() + app.include_router(router) + + engine = create_engine(db_url, connect_args={"check_same_thread": False}, poolclass=NullPool) + + # Import all agent models explicitly to register with Base.metadata + import src.models.agent # noqa: F401 + + Base.metadata.create_all(bind=engine) + + TestingSessionLocal = sessionmaker(bind=engine, autoflush=False) + + def _override_db(): + db = TestingSessionLocal() + try: + yield db + finally: + db.close() + + app.dependency_overrides[get_db] = _override_db + + return TestClient(app), TestingSessionLocal + + +@pytest.fixture +def admin_client(): + """TestClient with admin user and real SQLite in-memory DB.""" + tc, session_factory = _build_app_and_db() + mock_user = _mock_user("admin-1", is_admin=True) + tc.app.dependency_overrides[get_current_user] = lambda: mock_user + return tc, session_factory + + +@pytest.fixture +def user_client(): + """TestClient with regular (non-admin) user and real SQLite in-memory DB.""" + tc, _ = _build_app_and_db() + mock_user = _mock_user("regular-user-1", is_admin=False) + tc.app.dependency_overrides[get_current_user] = lambda: mock_user + return tc + + +# #region test_lifecycle_event_write_success [C:2] [TYPE Function] [SEMANTICS test,lifecycle,write] +# @BRIEF POST /api/agent/events creates an event and returns 201 with event ID. +def test_lifecycle_event_write_success(admin_client): + """POST /api/agent/events with valid body returns 201 and event ID.""" + tc, _ = admin_client + response = tc.post("/api/agent/events", json={ + "trace_id": "trace-abc-123", + "conversation_id": "conv-456", + "event_type": "AGENT_REQUEST_STARTED", + "environment_id": "env-prod", + "tool_name": None, + "status": None, + "elapsed_ms": None, + "payload": {"action": "send_message", "attempt": 1}, + "error_code": None, + }) + assert response.status_code == 201, response.text + data = response.json() + assert data["written"] is True + assert len(data["id"]) > 0 +# #endregion test_lifecycle_event_write_success + + +# #region test_lifecycle_event_write_reduces_payload [C:2] [TYPE Function] [SEMANTICS test,lifecycle,write,payload] +# @BRIEF POST /api/agent/events strips sensitive fields from payload. +def test_lifecycle_event_write_reduces_payload(admin_client): + """Sensitive fields in payload are stripped before storage.""" + tc, session_factory = admin_client + response = tc.post("/api/agent/events", json={ + "trace_id": "trace-sens-1", + "conversation_id": "conv-sens-1", + "event_type": "AGENT_TOOL_STARTED", + "payload": { + "action": "deploy", + "tool_input": "should-be-stripped", + "password": "secret123", + "token": "bearer-xxx", + }, + }) + assert response.status_code == 201, response.text + data = response.json() + + # Verify payload in DB is reduced + from src.models.agent import AgentLifecycleEvent + session = session_factory() + try: + event = session.query(AgentLifecycleEvent).filter_by(id=data["id"]).first() + assert event is not None + assert event.payload is not None + assert "action" in event.payload + assert "tool_input" not in event.payload + assert "password" not in event.payload + assert "token" not in event.payload + finally: + session.close() +# #endregion test_lifecycle_event_write_reduces_payload + + +# #region test_lifecycle_event_list_user_scoped [C:2] [TYPE Function] [SEMANTICS test,lifecycle,list,scope] +# @BRIEF GET /api/agent/events — non-admin user sees only own events. +def test_lifecycle_event_list_user_scoped(user_client, admin_client): + """Non-admin user can list their own events.""" + tc, session_factory = admin_client + + # Write event as admin (user_id = admin-1) + tc.post("/api/agent/events", json={ + "trace_id": "trace-admin", + "conversation_id": "conv-admin", + "event_type": "AGENT_REQUEST_STARTED", + "payload": {"action": "admin_action"}, + }) + + # Regular user lists events + u_tc = user_client + response = u_tc.get("/api/agent/events?page=1&page_size=50") + assert response.status_code == 200, response.text + data = response.json() + # Regular user should see 0 events (none belong to them) + assert data["total"] == 0 + assert data["items"] == [] +# #endregion test_lifecycle_event_list_user_scoped + + +# #region test_lifecycle_event_list_admin_cross_user [C:2] [TYPE Function] [SEMANTICS test,lifecycle,list,admin] +# @BRIEF GET /api/agent/events — admin user can query by user_id. +def test_lifecycle_event_list_admin_cross_user(admin_client): + """Admin user can query events by user_id filter.""" + tc, session_factory = admin_client + + # Write two events as different users + from src.services.agent_lifecycle_service import write_event + from src.schemas.agent_lifecycle import EventWriteRequest + from src.core.database import get_db + + # Create events by directly using service+DB + session = session_factory() + try: + write_event(session, EventWriteRequest( + trace_id="trace-1", conversation_id="conv-1", event_type="AGENT_REQUEST_STARTED", + ), authenticated_user_id="user-a") + write_event(session, EventWriteRequest( + trace_id="trace-2", conversation_id="conv-2", event_type="AGENT_LLM_STARTED", + ), authenticated_user_id="user-b") + session.commit() + finally: + session.close() + + # Admin lists all events + response = tc.get("/api/agent/events?page=1&page_size=50") + assert response.status_code == 200, response.text + data = response.json() + assert data["total"] == 2 + + # Admin filters by user_id + response = tc.get("/api/agent/events?user_id=user-a") + assert response.status_code == 200, response.text + data = response.json() + assert data["total"] == 1 + assert data["items"][0]["trace_id"] == "trace-1" + + # Admin filters by event_type + response = tc.get("/api/agent/events?event_type=AGENT_LLM_STARTED") + assert response.status_code == 200, response.text + data = response.json() + assert data["total"] == 1 + assert data["items"][0]["trace_id"] == "trace-2" +# #endregion test_lifecycle_event_list_admin_cross_user + + +# #region test_lifecycle_event_non_admin_cannot_query_others [C:2] [TYPE Function] [SEMANTICS test,lifecycle,list,forbidden] +# @BRIEF Non-admin user receives 403 when trying to filter by user_id. +def test_lifecycle_event_non_admin_cannot_query_others(user_client): + """Non-admin user gets 403 for user_id filter.""" + tc = user_client + response = tc.get("/api/agent/events?user_id=other-user") + assert response.status_code == 403, response.text + assert "Only admin users" in response.text +# #endregion test_lifecycle_event_non_admin_cannot_query_others + + +# #region test_lifecycle_event_list_pagination [C:2] [TYPE Function] [SEMANTICS test,lifecycle,list,pagination] +# @BRIEF GET /api/agent/events returns paginated results with has_next. +def test_lifecycle_event_list_pagination(admin_client): + """Pagination: page_size respected, has_next computed correctly.""" + tc, session_factory = admin_client + from src.services.agent_lifecycle_service import write_event + from src.schemas.agent_lifecycle import EventWriteRequest + + session = session_factory() + try: + for i in range(5): + write_event(session, EventWriteRequest( + trace_id=f"trace-{i}", conversation_id=f"conv-{i}", + event_type="AGENT_REQUEST_STARTED", + ), authenticated_user_id="admin-1") + session.commit() + finally: + session.close() + + # Page 1 with page_size=3 + response = tc.get("/api/agent/events?page=1&page_size=3") + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 3 + assert data["total"] == 5 + assert data["has_next"] is True + assert data["page"] == 1 + + # Page 2 with page_size=3 + response = tc.get("/api/agent/events?page=2&page_size=3") + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 2 + assert data["has_next"] is False +# #endregion test_lifecycle_event_list_pagination + + +# #region test_lifecycle_event_write_requires_auth [C:2] [TYPE Function] [SEMANTICS test,lifecycle,write,auth] +# @BRIEF POST /api/agent/events requires valid authentication token. +def test_lifecycle_event_write_requires_auth(): + """POST without auth returns 401/403 (depends on oauth2_scheme).""" + # Test without overriding get_current_user — will fail at oauth2_scheme + app = FastAPI() + app.include_router(router) + tc = TestClient(app) + + response = tc.post("/api/agent/events", json={ + "trace_id": "trace-noauth", + "conversation_id": "conv-noauth", + "event_type": "AGENT_REQUEST_STARTED", + }) + assert response.status_code == 401, response.text +# #endregion test_lifecycle_event_write_requires_auth + + +# #region test_lifecycle_event_list_filters [C:2] [TYPE Function] [SEMANTICS test,lifecycle,list,filters] +# @BRIEF GET /api/agent/events supports conversation_id, status, tool_name filters. +def test_lifecycle_event_list_filters(admin_client): + """All query parameter filters work correctly.""" + tc, session_factory = admin_client + from src.services.agent_lifecycle_service import write_event + from src.schemas.agent_lifecycle import EventWriteRequest + + session = session_factory() + try: + write_event(session, EventWriteRequest( + trace_id="t1", conversation_id="conv-a", event_type="AGENT_TOOL_STARTED", + tool_name="deploy", status="running", + ), authenticated_user_id="admin-1") + write_event(session, EventWriteRequest( + trace_id="t2", conversation_id="conv-b", event_type="AGENT_TOOL_COMPLETED", + tool_name="backup", status="success", + ), authenticated_user_id="admin-1") + write_event(session, EventWriteRequest( + trace_id="t3", conversation_id="conv-b", event_type="AGENT_TOOL_FAILED", + tool_name="backup", status="failed", error_code="TIMEOUT", + ), authenticated_user_id="admin-1") + session.commit() + finally: + session.close() + + # Filter by conversation_id + resp = tc.get("/api/agent/events?conversation_id=conv-a") + assert resp.status_code == 200 + assert resp.json()["total"] == 1 + + # Filter by status + resp = tc.get("/api/agent/events?status=success") + assert resp.status_code == 200 + assert resp.json()["total"] == 1 + + # Filter by tool_name + resp = tc.get("/api/agent/events?tool_name=deploy") + assert resp.status_code == 200 + assert resp.json()["total"] == 1 +# #endregion test_lifecycle_event_list_filters +# #endregion Test.Api.AgentLifecycle diff --git a/backend/src/api/routes/__tests__/test_migration_routes.py b/backend/src/api/routes/__tests__/test_migration_routes.py index 6fbedb977..9684e97d4 100644 --- a/backend/src/api/routes/__tests__/test_migration_routes.py +++ b/backend/src/api/routes/__tests__/test_migration_routes.py @@ -31,9 +31,15 @@ from sqlalchemy.orm import sessionmaker from src.models.mapping import Base, ResourceMapping, ResourceType -# Patch the get_db dependency if `src.api.routes.migration` imports it -patch("src.core.database.get_db").start() # --- Fixtures --- +@pytest.fixture(autouse=True) +def _patch_get_db(): + """Patch get_db for all tests with proper teardown.""" + patcher = patch("src.core.database.get_db") + patcher.start() + yield + patcher.stop() + @pytest.fixture def db_session(): """In-memory SQLite session for testing.""" @@ -300,12 +306,13 @@ async def test_trigger_sync_now_creates_env_row_and_syncs(db_session, _mock_env) from src.models.mapping import Environment as EnvironmentModel cm = _make_sync_config_manager([_mock_env]) with ( - patch("src.api.routes.migration.SupersetClient") as MockClient, + patch("src.api.routes.migration.AsyncSupersetClient") as MockClient, patch("src.api.routes.migration.IdMappingService") as MockService, ): mock_client_instance = MagicMock() MockClient.return_value = mock_client_instance mock_service_instance = MagicMock() + mock_service_instance.sync_environment = AsyncMock() MockService.return_value = mock_service_instance result = await trigger_sync_now(config_manager=cm, db=db_session, _=None) # Environment row must exist in DB @@ -342,14 +349,16 @@ async def test_trigger_sync_now_handles_partial_failure(db_session, _mock_env): env2.timeout = 30 cm = _make_sync_config_manager([_mock_env, env2]) with ( - patch("src.api.routes.migration.SupersetClient") as MockClient, + patch("src.api.routes.migration.AsyncSupersetClient") as MockClient, patch("src.api.routes.migration.IdMappingService") as MockService, ): mock_service_instance = MagicMock() - mock_service_instance.sync_environment.side_effect = [ - None, - RuntimeError("Connection refused"), - ] + mock_service_instance.sync_environment = AsyncMock( + side_effect=[ + None, + RuntimeError("Connection refused"), + ] + ) MockService.return_value = mock_service_instance MockClient.return_value = MagicMock() result = await trigger_sync_now(config_manager=cm, db=db_session, _=None) @@ -363,9 +372,12 @@ async def test_trigger_sync_now_idempotent_env_upsert(db_session, _mock_env): from src.models.mapping import Environment as EnvironmentModel cm = _make_sync_config_manager([_mock_env]) with ( - patch("src.api.routes.migration.SupersetClient"), - patch("src.api.routes.migration.IdMappingService"), + patch("src.api.routes.migration.AsyncSupersetClient"), + patch("src.api.routes.migration.IdMappingService") as MockService, ): + mock_service_instance = MagicMock() + mock_service_instance.sync_environment = AsyncMock() + MockService.return_value = mock_service_instance await trigger_sync_now(config_manager=cm, db=db_session, _=None) await trigger_sync_now(config_manager=cm, db=db_session, _=None) env_count = db_session.query(EnvironmentModel).filter_by(id="test-env-1").count() @@ -375,9 +387,11 @@ async def test_trigger_sync_now_idempotent_env_upsert(db_session, _mock_env): async def test_get_dashboards_success(_mock_env): from src.api.routes.migration import get_dashboards cm = _make_sync_config_manager([_mock_env]) - with patch("src.api.routes.migration.SupersetClient") as MockClient: + with patch("src.api.routes.migration.AsyncSupersetClient") as MockClient: mock_client = MagicMock() - mock_client.get_dashboards_summary.return_value = [{"id": 1, "title": "Test"}] + mock_client.get_dashboards_summary = AsyncMock( + return_value=[{"id": 1, "title": "Test"}] + ) MockClient.return_value = mock_client result = await get_dashboards(env_id="test-env-1", config_manager=cm, _=None) assert len(result) == 1 @@ -401,7 +415,10 @@ async def test_execute_migration_success(_mock_env): source_env_id="test-env-1", target_env_id="test-env-1", selected_ids=[1, 2] ) result = await execute_migration( - selection=selection, config_manager=cm, task_manager=tm, _=None + selection=selection, + config_manager=cm, + task_manager=tm, + current_user=MagicMock(id="user-1"), ) assert result["task_id"] == "task-123" tm.create_task.assert_called_once() @@ -415,7 +432,10 @@ async def test_execute_migration_invalid_env_raises_400(_mock_env): ) with pytest.raises(HTTPException) as exc: await execute_migration( - selection=selection, config_manager=cm, task_manager=MagicMock(), _=None + selection=selection, + config_manager=cm, + task_manager=MagicMock(), + current_user=MagicMock(id="user-1"), ) assert exc.value.status_code == 400 @pytest.mark.asyncio @@ -449,7 +469,7 @@ async def test_dry_run_migration_returns_diff_and_risk(db_session): fix_cross_filters=True, ) with ( - patch("src.api.routes.migration.SupersetClient") as MockClient, + patch("src.api.routes.migration.AsyncSupersetClient") as MockClient, patch("src.api.routes.migration.MigrationDryRunService") as MockService, ): source_client = MagicMock() @@ -488,7 +508,7 @@ async def test_dry_run_migration_returns_diff_and_risk(db_session): ], }, } - service_instance.run.return_value = service_payload + service_instance.run = AsyncMock(return_value=service_payload) MockService.return_value = service_instance result = await dry_run_migration( selection=selection, config_manager=cm, db=db_session, _=None diff --git a/backend/src/api/routes/agent_lifecycle.py b/backend/src/api/routes/agent_lifecycle.py new file mode 100644 index 000000000..11cbe0d8d --- /dev/null +++ b/backend/src/api/routes/agent_lifecycle.py @@ -0,0 +1,101 @@ +# backend/src/api/routes/agent_lifecycle.py +# #region Api.AgentLifecycle [C:3] [TYPE Module] [SEMANTICS agent,lifecycle,api,rest] +# @defgroup AgentLifecycle REST routes for agent lifecycle audit events. +# @BRIEF Write (immutable audit) and read (paginated, user-scoped) lifecycle event endpoints. +# @RELATION DEPENDS_ON -> [AgentLifecycleService] +# @RELATION DEPENDS_ON -> [AuthMiddleware] +# @INVARIANT POST /api/agent/events validates and reduces payload to safe whitelist before storage. +# @INVARIANT GET /api/agent/events enforces user-scoped access — non-admin users see only own events. + +from fastapi import APIRouter, Depends, HTTPException, Query, status as http_status +from sqlalchemy.orm import Session + +from ...core.database import get_db +from ...dependencies import get_current_user +from ...models.auth import User +from ...schemas.agent_lifecycle import ( + EventListResponse, + EventWriteRequest, + EventWriteResponse, +) +from ...services.agent_lifecycle_service import list_events, write_event + +router = APIRouter(prefix="/api/agent/events", tags=["Agent-Lifecycle"]) + + +# #region Api.AgentLifecycle.WriteEvent [C:3] [TYPE Function] [SEMANTICS agent,lifecycle,write,endpoint] +# @ingroup AgentLifecycle +# @BRIEF POST /api/agent/events — persist an immutable lifecycle audit event. +# @PRE Authenticated via end-user JWT (X-User-JWT preferred for agent delegation). +# @POST Event written to agent_lifecycle_events table. Payload automatically reduced to safe whitelist. +# @SIDE_EFFECT Writes to agent_lifecycle_events table and commits the request transaction. +# @INVARIANT Payload is validated against SAFE_PAYLOAD_KEYS — sensitive fields are stripped. +@router.post("", response_model=EventWriteResponse, status_code=http_status.HTTP_201_CREATED) +async def create_event( + body: EventWriteRequest, + current_user: User = Depends(get_current_user), + db: Session = Depends(get_db), +): + """Write a lifecycle event. The payload is automatically reduced to safe keys. + + Authenticated via standard JWT; the agent forwards the end-user token in + X-User-JWT when present. The authenticated user_id is the immutable owner. + """ + try: + result = write_event(db, body, authenticated_user_id=current_user.id) + db.commit() + return result + except Exception: + db.rollback() + raise +# #endregion Api.AgentLifecycle.WriteEvent + + +# #region Api.AgentLifecycle.ListEvents [C:3] [TYPE Function] [SEMANTICS agent,lifecycle,list,endpoint] +# @ingroup AgentLifecycle +# @BRIEF GET /api/agent/events — paginated list of lifecycle events. +# @PRE Authenticated via JWT. Non-admin users see only their own events. +# @POST Returns paginated EventListResponse with optional filters. +@router.get("", response_model=EventListResponse) +async def read_events( + page: int = Query(1, ge=1), + page_size: int = Query(50, ge=1, le=200), + user_id: str | None = Query(None, description="Filter by user_id (admin only)"), + event_type: str | None = Query(None), + conversation_id: str | None = Query(None), + status: str | None = Query(None), + tool_name: str | None = Query(None), + current_user: User = Depends(get_current_user), + db: Session = Depends(get_db), +): + """List lifecycle events with pagination and optional filters. + + Regular users see only their own events. Admin users may query by user_id. + """ + is_admin = any( + getattr(role, "is_admin", False) or role.name == "Admin" + for role in current_user.roles + ) + + # Non-admin users cannot query other users' events + if user_id and not is_admin: + raise HTTPException( + status_code=http_status.HTTP_403_FORBIDDEN, + detail="Only admin users can query events by user_id", + ) + + result = list_events( + db=db, + authenticated_user_id=current_user.id, + is_admin=is_admin, + page=page, + page_size=page_size, + user_id=user_id, + event_type=event_type, + conversation_id=conversation_id, + status=status, + tool_name=tool_name, + ) + return result +# #endregion Api.AgentLifecycle.ListEvents +# #endregion Api.AgentLifecycle diff --git a/backend/src/api/routes/assistant/_dispatch.py b/backend/src/api/routes/assistant/_dispatch.py index c80021082..d61f7c7ae 100644 --- a/backend/src/api/routes/assistant/_dispatch.py +++ b/backend/src/api/routes/assistant/_dispatch.py @@ -179,7 +179,9 @@ async def _async_confirmation_summary(intent: dict[str, Any], config_manager: Co service = MigrationDryRunService() source_client = AsyncSupersetClient(source_env) target_client = AsyncSupersetClient(target_env) - report = service.run(selection, source_client, target_client, db) + report = await service.run( + selection, source_client, target_client, db + ) s = report.get('summary', {}) dash_s = s.get('dashboards', {}) charts_s = s.get('charts', {}) diff --git a/backend/src/api/routes/dashboards/_action_routes.py b/backend/src/api/routes/dashboards/_action_routes.py index 578567a9e..39e2422ff 100644 --- a/backend/src/api/routes/dashboards/_action_routes.py +++ b/backend/src/api/routes/dashboards/_action_routes.py @@ -37,7 +37,7 @@ async def migrate_dashboards( request: MigrateRequest, config_manager=Depends(get_config_manager), task_manager=Depends(get_task_manager), - _=Depends(has_permission("plugin:migration", "EXECUTE")), + current_user=Depends(has_permission("plugin:migration", "EXECUTE")), ): with belief_scope( "migrate_dashboards", @@ -74,7 +74,9 @@ async def migrate_dashboards( } task_obj = await task_manager.create_task( - plugin_id="superset-migration", params=task_params + plugin_id="superset-migration", + params=task_params, + user_id=current_user.id, ) logger.reflect("Migration task created", payload={"task_id": str(task_obj.id), "dashboard_count": len(request.dashboard_ids)}) @@ -107,7 +109,7 @@ async def backup_dashboards( request: BackupRequest, config_manager=Depends(get_config_manager), task_manager=Depends(get_task_manager), - _=Depends(has_permission("plugin:backup", "EXECUTE")), + current_user=Depends(has_permission("plugin:backup", "EXECUTE")), ): with belief_scope( "backup_dashboards", @@ -134,7 +136,9 @@ async def backup_dashboards( } task_obj = await task_manager.create_task( - plugin_id="superset-backup", params=task_params + plugin_id="superset-backup", + params=task_params, + user_id=current_user.id, ) logger.reflect("Backup task created", payload={"task_id": str(task_obj.id), "dashboard_count": len(request.dashboard_ids)}) diff --git a/backend/src/api/routes/mappings.py b/backend/src/api/routes/mappings.py index ede73d8ba..b7d3fd7a7 100644 --- a/backend/src/api/routes/mappings.py +++ b/backend/src/api/routes/mappings.py @@ -10,13 +10,14 @@ # -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Request from pydantic import BaseModel, ConfigDict, Field from sqlalchemy.orm import Session +from ...core.async_superset_client import AsyncSupersetClient from ...core.database import get_db from ...core.logger import belief_scope, logger -from ...dependencies import get_config_manager, has_permission +from ...dependencies import check_api_key_environment_scope, get_config_manager, has_permission from ...models.mapping import DatabaseMapping router = APIRouter(tags=["mappings"]) @@ -55,6 +56,66 @@ class SuggestRequest(BaseModel): target_env_id: str # #endregion SuggestRequest + +# #region validate_mapping_database_ownership [C:4] [TYPE Function] [SEMANTICS mapping,database,environment,validation] +# @ingroup Api +# @BRIEF Verify that each submitted database UUID exists in its declared source or target environment. +# @PRE Mapping environment ids resolve to configured environments. +# @POST Returns only after source_db_uuid belongs to source_env_id and target_db_uuid belongs to target_env_id. +# @SIDE_EFFECT Reads database catalogs from both Superset environments; does not mutate either environment. +# @DATA_CONTRACT Input[MappingCreate] -> Output[None | HTTP_400 | HTTP_502] +# @RATIONALE UUIDs are environment-local identities; accepting arbitrary pairs permits stale or cross-environment mappings that can misdirect imports. +# @REJECTED Trusting UUIDs supplied by the browser was rejected — client state is not authorization or ownership evidence. +async def validate_mapping_database_ownership( + mapping: MappingCreate, config_manager +) -> None: + source_env = config_manager.get_environment(mapping.source_env_id) + target_env = config_manager.get_environment(mapping.target_env_id) + if not source_env or not target_env: + raise HTTPException( + status_code=400, detail="Invalid source or target environment" + ) + if mapping.source_env_id == mapping.target_env_id: + raise HTTPException( + status_code=400, + detail="Source and target environments must be different", + ) + source_client = AsyncSupersetClient(source_env) + target_client = AsyncSupersetClient(target_env) + try: + source_databases = await source_client.get_databases_summary() + target_databases = await target_client.get_databases_summary() + except Exception as exc: + logger.explore( + "Unable to validate mapping database ownership", + payload={ + "source_env_id": mapping.source_env_id, + "target_env_id": mapping.target_env_id, + }, + error=str(exc), + ) + raise HTTPException( + status_code=502, + detail="Unable to validate databases for the selected environments", + ) from exc + finally: + await source_client.aclose() + await target_client.aclose() + + source_uuids = {str(database.get("uuid")) for database in source_databases} + target_uuids = {str(database.get("uuid")) for database in target_databases} + if mapping.source_db_uuid not in source_uuids: + raise HTTPException( + status_code=400, + detail="Source database does not belong to the selected source environment", + ) + if mapping.target_db_uuid not in target_uuids: + raise HTTPException( + status_code=400, + detail="Target database does not belong to the selected target environment", + ) +# #endregion validate_mapping_database_ownership + # #region get_mappings [TYPE Function] # @ingroup Api # @BRIEF List all saved database mappings. @@ -62,13 +123,30 @@ class SuggestRequest(BaseModel): # @POST Returns filtered list of DatabaseMapping records. @router.get("", response_model=list[MappingResponse]) async def get_mappings( + request: Request, source_env_id: str | None = None, target_env_id: str | None = None, db: Session = Depends(get_db), _ = Depends(has_permission("plugin:mapper", "EXECUTE")) ): with belief_scope("get_mappings"): + # If the request uses an API key with environment scope, filter results + raw_key = request.headers.get("X-API-Key") + if raw_key: + from ...core.auth.api_key import hash_api_key + from ...models.api_key import APIKey + key_hash = hash_api_key(raw_key) + api_key = db.query(APIKey).filter(APIKey.key_hash == key_hash).first() + key_env = api_key.environment_id if api_key else None + else: + key_env = None + query = db.query(DatabaseMapping) + if key_env: + query = query.filter( + (DatabaseMapping.source_env_id == key_env) + | (DatabaseMapping.target_env_id == key_env) + ) if source_env_id: query = query.filter(DatabaseMapping.source_env_id == source_env_id) if target_env_id: @@ -85,9 +163,11 @@ async def get_mappings( async def create_mapping( mapping: MappingCreate, db: Session = Depends(get_db), + config_manager=Depends(get_config_manager), _ = Depends(has_permission("plugin:mapper", "EXECUTE")) ): with belief_scope("create_mapping"): + await validate_mapping_database_ownership(mapping, config_manager) # Check if mapping already exists existing = db.query(DatabaseMapping).filter( DatabaseMapping.source_env_id == mapping.source_env_id, @@ -217,10 +297,14 @@ async def apply_dataset_metadata( @router.post("/suggest") async def suggest_mappings_api( request: SuggestRequest, + req: Request, config_manager=Depends(get_config_manager), + db: Session = Depends(get_db), _ = Depends(has_permission("plugin:mapper", "EXECUTE")) ): with belief_scope("suggest_mappings_api"): + await check_api_key_environment_scope(req, db, request.source_env_id) + await check_api_key_environment_scope(req, db, request.target_env_id) from ...services.mapping_service import MappingService service = MappingService(config_manager) try: diff --git a/backend/src/api/routes/migration.py b/backend/src/api/routes/migration.py index 1b486be5d..a7a716b09 100644 --- a/backend/src/api/routes/migration.py +++ b/backend/src/api/routes/migration.py @@ -35,7 +35,11 @@ from ...core.mapping_service import IdMappingService from ...core.migration.dry_run_orchestrator import MigrationDryRunService from ...core.async_superset_client import AsyncSupersetClient from ...dependencies import get_config_manager, get_task_manager, has_permission -from ...models.dashboard import DashboardMetadata, DashboardSelection +from ...models.dashboard import ( + DashboardMetadata, + DashboardSelection, + MigrationDryRunResult, +) from ...models.mapping import ResourceMapping logger = cast(Any, logger) @@ -85,6 +89,7 @@ async def get_dashboards( # @RELATION CALLS -> [create_task] # @RELATION DEPENDS_ON -> [DashboardSelection] # @INVARIANT Migration task dispatch never occurs before source and target environment ids pass guard validation. +# @INVARIANT User-initiated migration tasks retain their creator id for resume authorization. # @RATIONALE Delegates to the async task manager to avoid blocking the HTTP request thread during potentially long-running migration operations that may involve many dashboards. # @REJECTED Synchronous execution within the HTTP request handler was rejected — it would block the worker process and cause client timeouts for large migrations. @router.post("/migration/execute") @@ -92,7 +97,7 @@ async def execute_migration( selection: DashboardSelection, config_manager=Depends(get_config_manager), task_manager=Depends(get_task_manager), - _=Depends(has_permission("plugin:migration", "EXECUTE")), + current_user=Depends(has_permission("plugin:migration", "EXECUTE")), ): with belief_scope("execute_migration"): logger.reason( @@ -128,7 +133,9 @@ async def execute_migration( ) try: - task = await task_manager.create_task("superset-migration", task_params) + task = await task_manager.create_task( + "superset-migration", task_params, user_id=current_user.id + ) logger.reflect(f"Migration task created: {task.id}") return {"task_id": task.id, "message": "Migration initiated"} except Exception as e: @@ -147,11 +154,11 @@ async def execute_migration( # @PRE DashboardSelection is valid, source and target environments exist, differ, and selected_ids is non-empty. # @POST Returns deterministic dry-run payload; emits HTTP_400 for guard violations and HTTP_500 for orchestrator value errors. # @SIDE_EFFECT Reads local mappings from DB and fetches source/target metadata via Superset API. -# @DATA_CONTRACT Input[DashboardSelection] -> Output[Dict[str, Any]] +# @DATA_CONTRACT Input[DashboardSelection] -> Output[MigrationDryRunResult] # @RELATION DEPENDS_ON -> [DashboardSelection] # @RELATION DEPENDS_ON -> [MigrationDryRunService] # @INVARIANT Dry-run flow remains read-only and rejects identical source/target environments before service execution. -@router.post("/migration/dry-run", response_model=dict[str, Any]) +@router.post("/migration/dry-run", response_model=MigrationDryRunResult) async def dry_run_migration( selection: DashboardSelection, config_manager=Depends(get_config_manager), diff --git a/backend/src/api/routes/tasks.py b/backend/src/api/routes/tasks.py index 20a316682..1a69f4073 100755 --- a/backend/src/api/routes/tasks.py +++ b/backend/src/api/routes/tasks.py @@ -123,7 +123,9 @@ async def create_task( db.close() task = await task_manager.create_task( - plugin_id=request.plugin_id, params=request.params + plugin_id=request.plugin_id, + params=request.params, + user_id=current_user.id, ) return task except ValueError as e: @@ -356,20 +358,47 @@ async def resolve_task( # #region resume_task [C:2] [TYPE Function] # @ingroup Api # @BRIEF Resume a task that is awaiting input (e.g., passwords). -# @PRE task must be in AWAITING_INPUT status. -# @POST Task resumes execution with provided input. +# @PRE Task must be in AWAITING_INPUT status and requester must own it unless an administrator overrides. +# @POST Task resumes execution with provided input without serializing password material. +# @INVARIANT A non-administrator cannot resume another user's task. # @RELATION CALLS -> [TaskManager] @router.post("/{task_id}/resume", response_model=Task) async def resume_task( task_id: str, request: ResumeTaskRequest, task_manager: TaskManager = Depends(get_task_manager), - _=Depends(has_permission("tasks", "WRITE")), + current_user=Depends(has_permission("tasks", "WRITE")), ): with belief_scope("resume_task"): try: - task_manager.resume_task_with_password(task_id, request.passwords) + task = task_manager.get_task(task_id) + if not task: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Task not found" + ) + is_admin = any( + getattr(role, "is_admin", False) or role.name == "Admin" + for role in current_user.roles + ) + task_owner_id = getattr(task, "user_id", None) + if ( + isinstance(task_owner_id, str) + and task_owner_id != current_user.id + and not is_admin + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Task belongs to a different user", + ) + await task_manager.resume_task_with_password( + task_id, + request.passwords, + requester_user_id=current_user.id, + allow_task_override=is_admin, + ) return task_manager.get_task(task_id) + except PermissionError as e: + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=str(e)) except ValueError as e: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) diff --git a/backend/src/app.py b/backend/src/app.py index c4634937c..c8e0aad57 100755 --- a/backend/src/app.py +++ b/backend/src/app.py @@ -21,6 +21,7 @@ from contextlib import asynccontextmanager import os from pathlib import Path import sys +import uuid # project_root is used for static files mounting project_root = Path(__file__).resolve().parent.parent.parent @@ -38,6 +39,7 @@ from .api.routes import ( admin, admin_api_keys, agent_conversations, + agent_lifecycle, agent_status, agent_superset, agent_superset_explore, @@ -65,7 +67,7 @@ from .api.routes import ( ) from .api.routes.validation_tasks import router as validation_tasks from .core.auth.security import get_password_hash -from ss_tools.shared.cot_logger import get_trace_id, seed_trace_id +from ss_tools.shared.cot_logger import get_trace_id, seed_trace_id, set_trace_id from .core.database import AuthSessionLocal, init_db from .core.encryption_key import ensure_encryption_key from .core.logger import belief_scope, logger @@ -476,6 +478,7 @@ app.include_router(reports.router) app.include_router(assistant.router, prefix="/api/assistant", tags=["Assistant"]) app.include_router(agent_conversations.agent_router, tags=["Agent"]) app.include_router(agent_conversations.router, tags=["Assistant"]) +app.include_router(agent_lifecycle.router) app.include_router(agent_status.router) app.include_router(agent_superset.router, tags=["Agent Superset"]) app.include_router(agent_superset_explore.router, tags=["Agent Superset"]) @@ -564,6 +567,32 @@ async def _authenticate_websocket(websocket: WebSocket, endpoint_name: str) -> b # #endregion _authenticate_websocket +# #region _set_websocket_trace_id [C:2] [TYPE Function] [SEMANTICS websocket,trace,context] +# @ingroup Module +# @BRIEF Apply a valid UUID4 x-trace-id query parameter to the current WebSocket context. +# @PRE websocket is a live WebSocket whose query params may include x-trace-id. +# @POST Valid UUID4 query value is propagated; invalid or missing values receive a new trace ID. +# @SIDE_EFFECT Sets the shared trace ContextVar for this WebSocket task. +# @RATIONALE Browser WebSocket clients cannot set arbitrary request headers, so trace +# propagation uses a query parameter alongside the authentication token. +# @REJECTED Requiring a custom WebSocket header was rejected — browser WebSocket APIs do not +# permit application-defined headers during the opening handshake. +def _set_websocket_trace_id(websocket: WebSocket) -> str: + incoming = websocket.query_params.get("x-trace-id", "") + if incoming: + try: + parsed = uuid.UUID(hex=incoming) + if parsed.version == 4: + set_trace_id(incoming) + return incoming + except (ValueError, AttributeError): + pass + return seed_trace_id() + + +# #endregion _set_websocket_trace_id + + # #region websocket_endpoint [C:5] [TYPE Function] # @ingroup Module # @BRIEF Provides a WebSocket endpoint for real-time log streaming of a task with server-side filtering. @@ -613,7 +642,7 @@ async def websocket_endpoint(websocket: WebSocket, task_id: str, source: str = N level: Filter logs by minimum level (DEBUG, INFO, WARNING, ERROR) token: JWT or API key for authentication (required, see [SEC:C-3]) """ - seed_trace_id() + _set_websocket_trace_id(websocket) with belief_scope("websocket_endpoint", f"task_id={task_id}"): # ── WebSocket authentication (see [SEC:C-3]) ── if not await _authenticate_websocket(websocket, "ws/logs"): @@ -815,7 +844,7 @@ async def task_events_websocket(websocket: WebSocket): Query Parameters: token: JWT or API key for authentication (required) """ - seed_trace_id() + _set_websocket_trace_id(websocket) with belief_scope("task_events_websocket"): if not await _authenticate_websocket(websocket, "ws/task-events"): await websocket.close(code=4001, reason="Authentication required") @@ -872,7 +901,7 @@ async def maintenance_events_websocket(websocket: WebSocket): Query Parameters: token: JWT or API key for authentication (required) """ - seed_trace_id() + _set_websocket_trace_id(websocket) with belief_scope("maintenance_events_websocket"): if not await _authenticate_websocket(websocket, "ws/maintenance/events"): await websocket.close(code=4001, reason="Authentication required") @@ -918,7 +947,7 @@ async def maintenance_events_websocket(websocket: WebSocket): # @SIDE_EFFECT Subscribes to dataset event queue in task manager lifecycle. @app.websocket("/ws/datasets/{env_id}") async def dataset_websocket_endpoint(websocket: WebSocket, env_id: str): - seed_trace_id() + _set_websocket_trace_id(websocket) with belief_scope("dataset_websocket_endpoint", f"env_id={env_id}"): # ── WebSocket authentication (see [SEC:C-3]) ── if not await _authenticate_websocket(websocket, "ws/datasets"): @@ -954,7 +983,7 @@ async def dataset_websocket_endpoint(websocket: WebSocket, env_id: str): # @UX_FEEDBACK Client receives {status, total_records, successful_records, failed_records, progressPct, ...} @app.websocket("/ws/translate/run/{run_id}") async def translate_run_websocket(websocket: WebSocket, run_id: str): - seed_trace_id() + _set_websocket_trace_id(websocket) if not await _authenticate_websocket(websocket, "ws/translate/run"): await websocket.close(code=4001, reason="Authentication required") return diff --git a/backend/src/core/database.py b/backend/src/core/database.py index 97b9e145b..637498ad5 100644 --- a/backend/src/core/database.py +++ b/backend/src/core/database.py @@ -16,6 +16,7 @@ from sqlalchemy.orm import sessionmaker # Import models to ensure they're registered with Base from ..models import ( + agent as _agent_models, # noqa: F401 api_key as _api_key_models, # noqa: F401 assistant as _assistant_models, # noqa: F401 auth as _auth_models, # noqa: F401 diff --git a/backend/src/core/migration_engine.py b/backend/src/core/migration_engine.py index 3f9232023..b27ebf114 100644 --- a/backend/src/core/migration_engine.py +++ b/backend/src/core/migration_engine.py @@ -83,7 +83,7 @@ class MigrationEngine: logger.reason(f"Extracting source archive to {temp_dir}") with zipfile.ZipFile(zip_path, "r") as zf: zf.extractall(temp_dir) - # 2. Transform YAMLs (Databases) + # 2. Transform YAMLs (Datasets) dataset_files = list(temp_dir.glob("**/datasets/**/*.yaml")) + list( temp_dir.glob("**/datasets/*.yaml") ) @@ -93,6 +93,20 @@ class MigrationEngine: ) for ds_file in dataset_files: self._transform_yaml(ds_file, db_mapping) + # 2.1 Transform YAMLs (Databases — replace UUID with target UUID) + # When a database UUID in the archive matches a target UUID that exists + # in the target Superset, Superset's import_database() will find it, + # skip creation, and populate database_ids — avoiding password errors. + db_files = list(temp_dir.glob("**/databases/**/*.yaml")) + list( + temp_dir.glob("**/databases/*.yaml") + ) + db_files = list(set(db_files)) + if db_files: + logger.reason( + f"Transforming {len(db_files)} database YAML files" + ) + for db_file in db_files: + self._transform_database_yaml(db_file, db_mapping) # 2.5 Patch Cross-Filters (Dashboards) if fix_cross_filters: if self.mapping_service and target_env_id: @@ -159,6 +173,41 @@ class MigrationEngine: yaml.dump(data, f) logger.reflect(f"Database UUID patched in {file_path.name}") # #endregion _transform_yaml + # #region _transform_database_yaml [TYPE Function] + # @PURPOSE: Replaces the top-level uuid field in a database YAML file with the target UUID. + # @PARAM file_path (Path) - Path to the database YAML file. + # @PARAM db_mapping (Dict[str, str]) - UUID mapping dictionary. + # @PRE file_path exists, is readable YAML, and db_mapping contains source->target UUID pairs. + # @POST uuid is replaced in-place. Superset import_database() will find the target UUID + # as an existing DB, skip creation, and populate database_ids for dataset import. + # @SIDE_EFFECT Reads and conditionally rewrites YAML file on disk. + # @RATIONALE Superset matches databases by UUID during import. If the database YAML in the + # archive has a UUID that already exists in the target instance, import_database() returns + # the existing DB immediately without password validation. Replacing source UUIDs with + # target UUIDs avoids both: password errors AND the cascading GENERIC_COMMAND_ERROR + # that occurs when databases/ is stripped entirely (empty database_ids blocks datasets). + def _transform_database_yaml(self, file_path: Path, db_mapping: dict[str, str]): + with belief_scope("MigrationEngine._transform_database_yaml"): + if not file_path.exists(): + logger.explore(f"Database YAML file not found: {file_path}") + raise FileNotFoundError(str(file_path)) + with open(file_path) as f: + data = yaml.safe_load(f) + if not data: + return + source_uuid = data.get("uuid") + if source_uuid is None: + logger.explore(f"Database YAML has no uuid field: {file_path.name}") + return + if source_uuid in db_mapping: + logger.reason(f"Replacing database UUID in {file_path.name}") + data["uuid"] = db_mapping[source_uuid] + with open(file_path, "w") as f: + yaml.dump(data, f) + logger.reflect(f"Database UUID patched: {source_uuid} → {db_mapping[source_uuid]} in {file_path.name}") + else: + logger.reason(f"Database UUID {source_uuid} not in mapping — keeping as-is in {file_path.name}") + # #endregion _transform_database_yaml # #region _extract_chart_uuids_from_archive [TYPE Function] # @PURPOSE: Scans extracted chart YAML files and builds a source chart ID to UUID lookup map. # @PRE temp_dir exists and points to extracted archive root with optional chart YAML resources. diff --git a/backend/src/core/scheduler.py b/backend/src/core/scheduler.py index 47b765c9c..f1fa0174c 100644 --- a/backend/src/core/scheduler.py +++ b/backend/src/core/scheduler.py @@ -18,6 +18,28 @@ from .database import TASKS_DATABASE_URL, SessionLocal from .logger import belief_scope, logger +# #region execute_scheduled_lifecycle_retention [C:2] [TYPE Function] [SEMANTICS scheduler,agent,lifecycle,retention] +# @ingroup Core +# @BRIEF APScheduler callback that prunes one bounded batch of old agent lifecycle events. +# @POST Old lifecycle events are pruned according to AGENT_LIFECYCLE_RETENTION_DAYS. +# @SIDE_EFFECT Deletes expired audit rows and commits the maintenance transaction. +def execute_scheduled_lifecycle_retention() -> None: + from .database import SessionLocal + from ..services.agent_lifecycle_retention import prune_events + + db = SessionLocal() + try: + deleted = prune_events(db) + db.commit() + logger.reason("Lifecycle retention job completed", payload={"deleted": deleted}) + except Exception as exc: + db.rollback() + logger.explore("Lifecycle retention job failed", error=str(exc)) + finally: + db.close() +# #endregion execute_scheduled_lifecycle_retention + + # #region execute_scheduled_backup [C:3] [TYPE Function] [SEMANTICS scheduler,backup,apscheduler,persistence] # @ingroup Core # @BRIEF APScheduler callback for backup jobs that resolves runtime dependencies at execution time. @@ -32,6 +54,7 @@ from .logger import belief_scope, logger def execute_scheduled_backup(env_id: str) -> None: """Resolve the scheduler service only when APScheduler invokes the job.""" from ..dependencies import get_scheduler_service + logger.reason("Scheduler lifecycle: backup executed", payload={"env_id": env_id}) get_scheduler_service()._trigger_backup(env_id) @@ -50,6 +73,7 @@ def execute_scheduled_backup(env_id: str) -> None: def execute_scheduled_validation(policy_id: str) -> None: """Resolve the scheduler service only when APScheduler invokes the job.""" from ..dependencies import get_scheduler_service + logger.reason("Scheduler lifecycle: validation executed", payload={"policy_id": policy_id}) get_scheduler_service()._trigger_validation(policy_id) @@ -106,8 +130,23 @@ class SchedulerService: with belief_scope("SchedulerService.start"): if not self.scheduler.running: self.scheduler.start() - logger.reason("Scheduler started") + logger.reason("Scheduler lifecycle: started") self.load_schedules() + self.scheduler.add_job( + execute_scheduled_lifecycle_retention, + CronTrigger.from_crontab("15 3 * * *", timezone="UTC"), + id="agent_lifecycle_retention", + replace_existing=True, + ) + # Log restored jobs from persistent jobstore + try: + restored_jobs = self.scheduler.get_jobs() + logger.reason( + "Scheduler lifecycle: jobs restored from persistent store", + payload={"count": len(restored_jobs), "jobs": [j.id for j in restored_jobs]}, + ) + except Exception: + pass # #endregion start # #region stop [TYPE Function] # @ingroup Core @@ -117,8 +156,17 @@ class SchedulerService: def stop(self): with belief_scope("SchedulerService.stop"): if self.scheduler.running: + # Log active jobs before shutdown + try: + active_jobs = self.scheduler.get_jobs() + logger.reason( + "Scheduler lifecycle: stopping with active jobs", + payload={"count": len(active_jobs), "jobs": [j.id for j in active_jobs]}, + ) + except Exception: + pass self.scheduler.shutdown() - logger.reflect("Scheduler stopped") + logger.reflect("Scheduler lifecycle: stopped") # #endregion stop # #region load_schedules [C:4] [TYPE Function] [SEMANTICS scheduler,backup,translation,apscheduler] # @ingroup Core @@ -225,6 +273,20 @@ class SchedulerService: pass except Exception as e: logger.explore("Failed differential prune of jobs", error=str(e)) + + # Lifecycle summary: count restored jobs by type + backup_count = sum(1 for jid in desired_job_ids if jid.startswith("backup_")) + translate_count = sum(1 for jid in desired_job_ids if jid.startswith("translate_")) + validation_count = sum(1 for jid in desired_job_ids if jid.startswith("validation_")) + logger.reason( + "Scheduler lifecycle: schedules loaded", + payload={ + "backup_schedules": backup_count, + "translation_schedules": translate_count, + "validation_schedules": validation_count, + "total_schedules": len(desired_job_ids), + }, + ) # #endregion load_schedules # #region SchedulerService.add_backup_job [C:3] [TYPE Function] [SEMANTICS scheduler,backup,cron] # @ingroup Core @@ -328,7 +390,7 @@ class SchedulerService: def _trigger_backup(self, env_id: str): seed_trace_id() with belief_scope("SchedulerService._trigger_backup", f"env_id={env_id}"): - logger.reason(f"Triggering scheduled backup for environment {env_id}") + logger.reason("Scheduler lifecycle: backup triggered", payload={"env_id": env_id}) # Check if a backup is already running for this environment active_tasks = self.task_manager.get_tasks(limit=100) # Check both environment_id and env keys for backward compat @@ -339,8 +401,8 @@ class SchedulerService: and task.status in ("PENDING", "RUNNING") and task_env == env_id ): - logger.explore(f"Backup already running for environment {env_id}, skipping scheduled run", - error="Concurrent backup in progress") + logger.reason("Scheduler lifecycle: backup skipped", + payload={"env_id": env_id, "reason": "already running"}) return # Run the backup task via AsyncJobRunner self.runner.run( @@ -348,6 +410,8 @@ class SchedulerService: "superset-backup", {"environment_id": env_id} ) ) + logger.reason("Scheduler lifecycle: backup task submitted", + payload={"env_id": env_id, "plugin_id": "superset-backup"}) # #endregion _trigger_backup # #region add_validation_job [C:3] [TYPE Function] [SEMANTICS scheduler,validation,cron,apscheduler] # @ingroup Core @@ -450,21 +514,24 @@ class SchedulerService: def _trigger_validation(self, policy_id: str) -> None: seed_trace_id() with belief_scope("SchedulerService._trigger_validation", f"policy_id={policy_id}"): - logger.reason(f"Triggering scheduled validation for policy {policy_id}") + logger.reason("Scheduler lifecycle: validation triggered", payload={"policy_id": policy_id}) db = SessionLocal() try: from ..models.llm import ValidationPolicy policy = db.query(ValidationPolicy).filter(ValidationPolicy.id == policy_id).first() if not policy: - logger.explore(f"Validation policy not found: {policy_id}") + logger.reason("Scheduler lifecycle: validation skipped", + payload={"policy_id": policy_id, "reason": "policy_not_found"}) return if not policy.is_active: - logger.reason(f"Validation policy {policy_id} is no longer active; skipping") + logger.reason("Scheduler lifecycle: validation skipped", + payload={"policy_id": policy_id, "reason": "policy_inactive"}) return dashboard_ids = list(policy.dashboard_ids or []) if not dashboard_ids: - logger.explore(f"Validation policy {policy_id} has no dashboards; skipping") + logger.reason("Scheduler lifecycle: validation skipped", + payload={"policy_id": policy_id, "reason": "no_dashboards"}) return ws = policy.window_start @@ -499,9 +566,13 @@ class SchedulerService: logger.reason( f"Scheduled validation for dashboard {dash_id} at {sched_time.isoformat()}" ) + logger.reason( + "Scheduler lifecycle: validation tasks submitted", + payload={"policy_id": policy_id, "dashboard_count": len(dashboard_ids)}, + ) except Exception as e: logger.explore( - "Error triggering validation for policy", + "Scheduler lifecycle: validation failed", payload={"policy_id": policy_id}, error=str(e), ) diff --git a/backend/src/core/superset_client/_base.py b/backend/src/core/superset_client/_base.py index 4c5c4ef1b..bc33e102d 100644 --- a/backend/src/core/superset_client/_base.py +++ b/backend/src/core/superset_client/_base.py @@ -126,7 +126,7 @@ class SupersetClientBase: # @PURPOSE Performs the actual multipart upload for import. # @SIDE_EFFECT Асинхронный multipart HTTP-запрос к Superset API для импорта дашборда. # @RELATION CALLS -> [AsyncAPIClient] - async def _do_import(self, file_name: str | Path) -> dict: + async def _do_import(self, file_name: str | Path, passwords: dict[str, str] | None = None) -> dict: with belief_scope("_do_import"): app_logger.reason(f"Uploading file: {file_name}") file_path = Path(file_name) @@ -135,6 +135,11 @@ class SupersetClientBase: f"File does not exist: {file_name}", error="FileNotFound" ) raise FileNotFoundError(f"File does not exist: {file_name}") + extra_data: dict[str, str] = {"overwrite": "true"} + if passwords: + import json + extra_data["passwords"] = json.dumps(passwords) + app_logger.reason("Passing database passwords to Superset import", payload={"keys": list(passwords.keys())}) return await self.client.upload_file( endpoint="/dashboard/import/", file_info={ @@ -142,7 +147,7 @@ class SupersetClientBase: "file_name": file_path.name, "form_field": "formData", }, - extra_data={"overwrite": "true"}, + extra_data=extra_data, timeout=self.env.timeout * 2, ) # #endregion SupersetClientDoImport diff --git a/backend/src/core/superset_client/_dashboards_crud.py b/backend/src/core/superset_client/_dashboards_crud.py index 7dfa92204..ad06634d5 100644 --- a/backend/src/core/superset_client/_dashboards_crud.py +++ b/backend/src/core/superset_client/_dashboards_crud.py @@ -300,6 +300,7 @@ class SupersetDashboardsCrudMixin: file_name: str | Path, dash_id: int | None = None, dash_slug: str | None = None, + passwords: dict[str, str] | None = None, ) -> dict: with belief_scope("import_dashboard"): if file_name is None: @@ -307,7 +308,7 @@ class SupersetDashboardsCrudMixin: file_path = str(file_name) self._validate_import_file(file_path) try: - return await self._do_import(file_path) + return await self._do_import(file_path, passwords=passwords) except Exception as exc: app_logger.explore( "First import attempt failed", @@ -327,7 +328,7 @@ class SupersetDashboardsCrudMixin: f"Deleted dashboard ID {target_id}, retrying import", payload={"target_id": target_id}, ) - return await self._do_import(file_path) + return await self._do_import(file_path, passwords=passwords) # #endregion SupersetClientImportDashboard # #region SupersetClientDeleteDashboard [TYPE Function] [C:3] # @ingroup Core diff --git a/backend/src/core/task_manager/context.py b/backend/src/core/task_manager/context.py index 35bde646a..97b8c329e 100644 --- a/backend/src/core/task_manager/context.py +++ b/backend/src/core/task_manager/context.py @@ -78,6 +78,7 @@ class TaskContext: with belief_scope("__init__"): self._task_id = task_id self._params = params + self._default_source = default_source self._background_tasks = background_tasks self._logger = TaskLogger( task_id=task_id, add_log_fn=add_log_fn, source=default_source diff --git a/backend/src/core/task_manager/lifecycle.py b/backend/src/core/task_manager/lifecycle.py index b9f5f7378..f729e1e71 100644 --- a/backend/src/core/task_manager/lifecycle.py +++ b/backend/src/core/task_manager/lifecycle.py @@ -54,7 +54,7 @@ def _task_to_dict(task: Task) -> dict: "started_at": task.started_at.isoformat() if task.started_at else None, "finished_at": task.finished_at.isoformat() if task.finished_at else None, "user_id": task.user_id, - "params": task.params, + "params": {k: v for k, v in task.params.items() if k != "passwords"}, "result": task.result, "input_required": task.input_required, "input_request": task.input_request, @@ -175,7 +175,15 @@ class JobLifecycle: # Agent-centric top-level marker for the task execution in the main trace main_logger.reason( "Execute plugin task", - payload={"task_id": task_id, "plugin": plugin.name, "params_summary": str(task.params)[:120] if task.params else None} + payload={ + "task_id": task_id, + "plugin": plugin.name, + "params_summary": str( + {key: value for key, value in task.params.items() if key != "passwords"} + )[:120] + if task.params + else None, + }, ) task.status = TaskStatus.RUNNING task.started_at = datetime.now(UTC) @@ -411,19 +419,21 @@ class JobLifecycle: await add_log_callback( task_id, "INFO", "Task paused for user input", - metadata={"input_request": input_request}, + context={"input_request": input_request}, ) # #endregion await_input # #region resume_task_with_password [C:3] [TYPE Function] [SEMANTICS task,resume,password,input] # @ingroup TaskManager # @BRIEF Resume a task that is awaiting input with provided passwords. - # @PRE Task exists and is in AWAITING_INPUT state. + # @PRE Task exists and is in AWAITING_INPUT state; requester owns the task unless explicitly authorized as an administrator. # @POST Task status changed to RUNNING, passwords injected, task resumed. + # @SIDE_EFFECT Holds passwords only in memory until the plugin consumes them; persists and broadcasts redacted task state. # @RAISES ValueError if task not found, not awaiting input, or passwords invalid. async def resume_task_with_password( self, task_id: str, passwords: dict[str, str], - add_log_callback=None, + add_log_callback=None, requester_user_id: str | None = None, + allow_task_override: bool = False, ) -> None: with belief_scope( "JobLifecycle.resume_task_with_password", f"task_id={task_id}" @@ -431,6 +441,13 @@ class JobLifecycle: task = self.graph.get_task(task_id) if not task: raise ValueError(f"Task {task_id} not found") + if ( + requester_user_id is not None + and task.user_id is not None + and requester_user_id != task.user_id + and not allow_task_override + ): + raise PermissionError("Task belongs to a different user") if task.status != TaskStatus.AWAITING_INPUT: raise ValueError( f"Task {task_id} is not AWAITING_INPUT (current: {task.status})" diff --git a/backend/src/core/task_manager/manager.py b/backend/src/core/task_manager/manager.py index 573013b56..b6971ef95 100644 --- a/backend/src/core/task_manager/manager.py +++ b/backend/src/core/task_manager/manager.py @@ -382,11 +382,17 @@ class TaskManager: # @ingroup TaskManager # @BRIEF Resume a task that is awaiting input with provided passwords. async def resume_task_with_password( - self, task_id: str, passwords: dict[str, str] + self, + task_id: str, + passwords: dict[str, str], + requester_user_id: str | None = None, + allow_task_override: bool = False, ) -> None: await self.lifecycle.resume_task_with_password( task_id, passwords, add_log_callback=self._add_log, + requester_user_id=requester_user_id, + allow_task_override=allow_task_override, ) # #endregion resume_task_with_password diff --git a/backend/src/core/task_manager/models.py b/backend/src/core/task_manager/models.py index 70ee31b40..a044c527c 100644 --- a/backend/src/core/task_manager/models.py +++ b/backend/src/core/task_manager/models.py @@ -25,7 +25,7 @@ from enum import Enum from typing import Any import uuid -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, field_serializer # #region TaskStatus [TYPE Enum] @@ -105,6 +105,7 @@ class LogStats(BaseModel): # @RELATION DEPENDS_ON -> [TaskStatus] # @RELATION DEPENDS_ON -> [LogEntry] # @RELATION DEPENDS_ON -> [TaskManager] +# @INVARIANT Sensitive runtime parameters are excluded from every serialized Task response. class Task(BaseModel): id: str = Field(default_factory=lambda: str(uuid.uuid4())) plugin_id: str @@ -138,5 +139,14 @@ class Task(BaseModel): if self.status == TaskStatus.AWAITING_INPUT and not self.input_request: raise ValueError("input_request is required when status is AWAITING_INPUT") # #endregion __init__ + + # #region serialize_params [C:2] [TYPE Function] [SEMANTICS task,serialization,secrets] + # @ingroup TaskManager + # @BRIEF Exclude transient password material from every serialized task representation. + # @POST Serialized params never contain the passwords key while in-memory plugin execution retains it. + @field_serializer("params") + def serialize_params(self, params: dict[str, Any]) -> dict[str, Any]: + return {key: value for key, value in params.items() if key != "passwords"} + # #endregion serialize_params # #endregion Task # #endregion TaskManagerModels diff --git a/backend/src/core/task_manager/persistence.py b/backend/src/core/task_manager/persistence.py index ac2691d7d..fd0273452 100644 --- a/backend/src/core/task_manager/persistence.py +++ b/backend/src/core/task_manager/persistence.py @@ -167,6 +167,7 @@ class TaskPersistenceService: session.add(record) record.type = task.plugin_id record.status = task.status.value + record.user_id = task.user_id raw_env_id = task.params.get("environment_id") or task.params.get( "source_env_id" ) @@ -185,7 +186,8 @@ class TaskPersistenceService: elif isinstance(obj, datetime): return obj.isoformat() return obj - record.params = json_serializable(task.params) + safe_params = {k: v for k, v in task.params.items() if k != "passwords"} + record.params = json_serializable(safe_params) record.result = json_serializable(task.result) # Persist retry/progress state into params for round-tripping (no dedicated columns yet) @@ -282,6 +284,7 @@ class TaskPersistenceService: id=record.id, plugin_id=record.type, status=TaskStatus(record.status), + user_id=record.user_id, started_at=started_at, finished_at=finished_at, params=params_dict, diff --git a/backend/src/dependencies.py b/backend/src/dependencies.py index af9a8ddf1..bbea4879a 100755 --- a/backend/src/dependencies.py +++ b/backend/src/dependencies.py @@ -580,8 +580,11 @@ def get_current_user( except JWTError: raise credentials_exception - # Check blacklist (only against the primary auth token, not forwarded user JWT) - if not x_user_jwt and is_token_blacklisted(token, db): + # The effective identity token is the credential used for decoding and user + # lookup, so it must also be the credential checked for revocation. Checking + # only the service Authorization token would allow a revoked X-User-JWT to + # bypass the blacklist when dual identity authentication is used. + if is_token_blacklisted(effective_token, db): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Token has been revoked", diff --git a/backend/src/models/agent.py b/backend/src/models/agent.py index d88eb6783..744af71fc 100644 --- a/backend/src/models/agent.py +++ b/backend/src/models/agent.py @@ -1,10 +1,12 @@ # backend/src/models/agent.py -# #region Models.Agent [C:2] [TYPE Module] [SEMANTICS agent,model,database] -# @BRIEF SQLAlchemy models for Gradio Agent Chat conversations. +# #region Models.Agent [C:3] [TYPE Module] [SEMANTICS agent,model,database] +# @BRIEF SQLAlchemy models for Gradio Agent Chat conversations and lifecycle audit. +# @RELATION DEPENDS_ON -> [Models.User] +from datetime import datetime, timezone import uuid -from sqlalchemy import JSON, Boolean, Column, DateTime, ForeignKey, String, Text +from sqlalchemy import JSON, Boolean, Column, DateTime, Float, ForeignKey, Index, String, Text from sqlalchemy.orm import relationship from .mapping import Base @@ -57,4 +59,35 @@ class AgentMessage(Base): conversation = relationship("AgentConversation", back_populates="messages") # #endregion Models.Agent.AgentMessage + + +# #region Models.Agent.AgentLifecycleEvent [C:3] [TYPE Class] [SEMANTICS agent,lifecycle,audit,model] +# @ingroup Models +# @BRIEF Immutable audit record for agent lifecycle events (request start, LLM call, tool execution, completion). +# @INVARIANT Once written, trace_id, conversation_id, user_id, environment_id, and event_type are immutable. +# @INVARIANT payload is validated against a safe whitelist before storage — sensitive fields are rejected. +# @RELATION DEPENDS_ON -> [Models.User] +class AgentLifecycleEvent(Base): + __tablename__ = "agent_lifecycle_events" + + id = Column(String, primary_key=True, default=_uuid) + trace_id = Column(String, nullable=False, index=True) + conversation_id = Column(String, nullable=False, index=True) + user_id = Column(String, nullable=False, index=True) + environment_id = Column(String, nullable=True, index=True) + event_type = Column(String, nullable=False, index=True) + tool_name = Column(String, nullable=True, index=True) + status = Column(String, nullable=True, index=True) + created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False) + elapsed_ms = Column(Float, nullable=True) + payload = Column(JSON, nullable=True) + error_code = Column(String, nullable=True, index=True) + + # Composite index for common query patterns + __table_args__ = ( + Index("ix_agent_lifecycle_events_user_type_created", "user_id", "event_type", "created_at"), + Index("ix_agent_lifecycle_events_conv_type_created", "conversation_id", "event_type", "created_at"), + Index("ix_agent_lifecycle_events_created", "created_at"), + ) +# #endregion Models.Agent.AgentLifecycleEvent # #endregion Models.Agent diff --git a/backend/src/models/dashboard.py b/backend/src/models/dashboard.py index 545482fe6..a890f44ec 100644 --- a/backend/src/models/dashboard.py +++ b/backend/src/models/dashboard.py @@ -5,7 +5,7 @@ # @RELATION CALLED_BY -> [MigrationApi] -from pydantic import BaseModel +from pydantic import BaseModel, Field # #region DashboardMetadata [C:1] [TYPE Class] @@ -27,4 +27,66 @@ class DashboardSelection(BaseModel): fix_cross_filters: bool = True # #endregion DashboardSelection + +# #region MigrationDiffObject [C:1] [TYPE Class] [SEMANTICS migration,dry-run,dto] +# @BRIEF Reference to a resource affected by a dry-run diff. +class MigrationDiffObject(BaseModel): + uuid: str + title: str | None = None + target_title: str | None = None +# #endregion MigrationDiffObject + + +# #region MigrationDiffBucket [C:1] [TYPE Class] [SEMANTICS migration,dry-run,dto] +# @BRIEF Create, update, and delete groups for one Superset resource type. +class MigrationDiffBucket(BaseModel): + create: list[MigrationDiffObject] = Field(default_factory=list) + update: list[MigrationDiffObject] = Field(default_factory=list) + delete: list[MigrationDiffObject] = Field(default_factory=list) +# #endregion MigrationDiffBucket + + +# #region MigrationRiskItem [C:1] [TYPE Class] [SEMANTICS migration,dry-run,risk,dto] +# @BRIEF A single migration risk with a stable machine-readable code. +class MigrationRiskItem(BaseModel): + code: str + severity: str + object_type: str + object_uuid: str + message: str +# #endregion MigrationRiskItem + + +# #region MigrationRiskReport [C:1] [TYPE Class] [SEMANTICS migration,dry-run,risk,dto] +# @BRIEF Bounded aggregate score and risks returned by a dry run. +class MigrationRiskReport(BaseModel): + score: int = Field(ge=0, le=100) + level: str + items: list[MigrationRiskItem] = Field(default_factory=list) +# #endregion MigrationRiskReport + + +# #region MigrationSummary [C:1] [TYPE Class] [SEMANTICS migration,dry-run,summary,dto] +# @BRIEF Resource-count summary returned by a dry run. +class MigrationSummary(BaseModel): + dashboards: dict[str, int] + charts: dict[str, int] + datasets: dict[str, int] + selected_dashboards: int = Field(ge=0) +# #endregion MigrationSummary + + +# #region MigrationDryRunResult [C:5] [TYPE Class] [SEMANTICS migration,dry-run,dto,api] +# @BRIEF Validated read-only pre-flight report shared with the migration frontend. +# @DATA_CONTRACT Input[DashboardSelection] -> Output[generated_at, selection, diff, summary, risk] +# @INVARIANT Every diff bucket and risk item has a typed, serializable shape. +class MigrationDryRunResult(BaseModel): + generated_at: str + selection: DashboardSelection + selected_dashboard_titles: list[str] = Field(default_factory=list) + diff: dict[str, MigrationDiffBucket] + summary: MigrationSummary + risk: MigrationRiskReport +# #endregion MigrationDryRunResult + # #endregion DashboardModels diff --git a/backend/src/models/task.py b/backend/src/models/task.py index 3b667ef48..015577765 100644 --- a/backend/src/models/task.py +++ b/backend/src/models/task.py @@ -23,6 +23,7 @@ class TaskRecord(Base): type = Column(String, nullable=False) # e.g., "backup", "migration" status = Column(String, nullable=False) # Enum: "PENDING", "RUNNING", "SUCCESS", "FAILED" environment_id = Column(String, ForeignKey("environments.id", ondelete="SET NULL"), nullable=True) + user_id = Column(String, nullable=True) started_at = Column(DateTime(timezone=True), nullable=True) finished_at = Column(DateTime(timezone=True), nullable=True) logs = Column(JSON, nullable=True) # Store structured logs as JSON (legacy, kept for backward compatibility) diff --git a/backend/src/plugins/migration.py b/backend/src/plugins/migration.py index 6fb334aea..601bef109 100755 --- a/backend/src/plugins/migration.py +++ b/backend/src/plugins/migration.py @@ -297,6 +297,14 @@ class MigrationPlugin(PluginBase): exported_content, _ = await from_c.export_dashboard(dash_id) with create_temp_file(content=exported_content, suffix=".zip") as tmp_zip_path, create_temp_file(suffix=".zip") as tmp_new_zip: + # @RATIONALE Database YAMLs in the ZIP are now transformed to target UUIDs + # by MigrationEngine._transform_database_yaml(). Superset matches DBs by UUID + # during import. With target UUIDs, import_database() finds the existing DB, + # skips creation, and populates database_ids — avoiding both password errors + # and the GENERIC_COMMAND_ERROR 1010 cascade. No need to strip databases. + # @REJECTED strip_databases=True was rejected — it causes a cascading failure: + # empty database_ids → all datasets skipped → all charts skipped → + # update_id_refs(chart_ids={}, dataset_info={}) → KeyError → error 1010. success = engine.transform_zip( str(tmp_zip_path), str(tmp_new_zip), @@ -360,11 +368,12 @@ class MigrationPlugin(PluginBase): app_logger.explore("Missing DB password detected during ingestion. Escalating to UI.", extra={"db_name": db_name}) if task_id: - tm.await_input(task_id, { + add_log = context._logger._add_log if context else None + await tm.await_input(task_id, { "type": "database_password", "databases": [db_name], - "error_message": error_msg - }) + "error_message": "A database password is required to continue this migration.", + }, add_log_callback=add_log) await tm.wait_for_input(task_id) task = tm.get_task(task_id) @@ -372,15 +381,20 @@ class MigrationPlugin(PluginBase): if passwords: app_logger.reason(f"Retrying import for {title} with injected credentials") - await to_c.import_dashboard(file_name=tmp_new_zip, dash_id=dash_id, dash_slug=dash_slug, passwords=passwords) - migration_result["migrated_dashboards"].append({"id": dash_id, "title": title}) - app_logger.reflect("Password injection unblocked import") - if "passwords" in task.params: - del task.params["passwords"] + try: + await to_c.import_dashboard(file_name=tmp_new_zip, dash_id=dash_id, dash_slug=dash_slug, passwords=passwords) + migration_result["migrated_dashboards"].append({"id": dash_id, "title": title}) + app_logger.reflect("Password injection unblocked import") + finally: + task.params.pop("passwords", None) continue app_logger.explore(f"Catastrophic dashboard ingestion failure: {exc}") migration_result["failed_dashboards"].append({"id": dash_id, "title": title, "error": str(exc)}) + if task_id: + task = tm.get_task(task_id) + if task and "passwords" in task.params: + del task.params["passwords"] if migration_result["failed_dashboards"]: migration_result["status"] = "PARTIAL_SUCCESS" diff --git a/backend/src/plugins/translate/_batch_proc.py b/backend/src/plugins/translate/_batch_proc.py index 97b20fca4..c4c905b03 100644 --- a/backend/src/plugins/translate/_batch_proc.py +++ b/backend/src/plugins/translate/_batch_proc.py @@ -31,6 +31,7 @@ from ._llm_call import LLMTranslationService from ._token_budget import estimate_token_budget from ._utils import _check_translation_cache, _compute_key_hash, _compute_source_hash from .dictionary import DictionaryManager +from .events import TranslationEventLog # #region BatchProcessingService [C:4] [TYPE Class] @@ -43,6 +44,7 @@ class BatchProcessingService: self.db = db self.config_manager = config_manager self._llm_service = LLMTranslationService(db) + self._event_log = TranslationEventLog(db) # #region process_batch [C:3] [TYPE Function] # @ingroup Translate @@ -69,11 +71,51 @@ class BatchProcessingService: # ★ Run local language detection on all rows (heuristic, no LLM) await self._detect_languages(batch_rows, tls) + # Emit LANGUAGE_DETECTION_COMPLETED aggregate event (no source text in payload) + lang_counts: dict[str, int] = {} + for row in batch_rows: + dl = row.get("_detected_lang", "und") or "und" + lang_counts[dl] = lang_counts.get(dl, 0) + 1 + und_count = lang_counts.get("und", 0) + total = len(batch_rows) if batch_rows else 0 + self._event_log.log_event( + job_id=job.id, + run_id=run_id, + event_type="LANGUAGE_DETECTION_COMPLETED", + payload={ + "detector": "lingua", + "source_language_distribution": dict(sorted(lang_counts.items())), + "und_count": und_count, + "und_rate": round(und_count / total, 4) if total > 0 else 0.0, + "total_rows": total, + "target_language_count": len(tls), + }, + ) + source_texts = [r.get("source_text", "") for r in batch_rows if r.get("source_text")] rc = batch_rows[0].get("source_data") if batch_rows else None dict_matches = DictionaryManager.filter_for_batch(self.db, source_texts, job.id, row_context=rc) self._check_cache(job, batch_rows, dict_snapshot_hash, config_hash) + + # Emit CACHE_SUMMARY aggregate event (no source text in payload) + eligible = sum( + 1 for r in batch_rows + if r.get("source_text", "") and not r.get("approved_translation") + ) + hits = sum(1 for r in batch_rows if r.get("_cached_lang_values")) + self._event_log.log_event( + job_id=job.id, + run_id=run_id, + event_type="CACHE_SUMMARY", + payload={ + "batch_size": len(batch_rows), + "eligible": eligible, + "hits": hits, + "misses": eligible - hits, + }, + ) + llm_rows, pre_rows = self._classify(batch_rows, preview_edits_cache, tls) result["successful"] += self._persist_pre(pre_rows, bid, run_id, tls) diff --git a/backend/src/plugins/translate/events.py b/backend/src/plugins/translate/events.py index 873c32b1f..3e519659b 100644 --- a/backend/src/plugins/translate/events.py +++ b/backend/src/plugins/translate/events.py @@ -33,6 +33,9 @@ VALID_EVENT_TYPES = TERMINAL_EVENT_TYPES | { "PREVIEW_CREATED", "PREVIEW_ACCEPTED", "PREVIEW_DISCARDED", "SCHEDULE_CREATED", "SCHEDULE_UPDATED", "SCHEDULE_DELETED", "METRICS_SNAPSHOT_CREATED", + # Aggregate observability events (no source text in payload — safe aggregate data only) + "LANGUAGE_DETECTION_COMPLETED", + "CACHE_SUMMARY", } DEFAULT_RETENTION_DAYS = 90 diff --git a/backend/src/schemas/agent_lifecycle.py b/backend/src/schemas/agent_lifecycle.py new file mode 100644 index 000000000..36dff5cc0 --- /dev/null +++ b/backend/src/schemas/agent_lifecycle.py @@ -0,0 +1,140 @@ +# backend/src/schemas/agent_lifecycle.py +# #region Schemas.AgentLifecycle [C:2] [TYPE Module] [SEMANTICS agent,lifecycle,schema,api] +# @BRIEF Pydantic schemas for agent lifecycle event API (write + read). +# @INVARIANT payload is validated against SAFE_PAYLOAD_KEYS whitelist — sensitive fields are rejected. +# @RELATION DEPENDS_ON -> [Models.Agent.AgentLifecycleEvent] + +from datetime import datetime, timezone +from typing import Any + +from pydantic import BaseModel, field_serializer + + +def _serialize_datetime(v: datetime) -> str: + """Serialize datetime as ISO 8601 with Z suffix.""" + normalized = v.replace(tzinfo=timezone.utc) if v.tzinfo is None else v.astimezone(timezone.utc) + return normalized.isoformat().replace("+00:00", "Z") + + +# ── SAFE PAYLOAD WHITELIST ─────────────────────────────────────────── +# Only keys in this set are permitted in lifecycle event payloads. +# All other keys are stripped during validation to prevent sensitive data leakage. +SAFE_PAYLOAD_KEYS: frozenset[str] = frozenset({ + "action", + "attempt", + "attempts", + "conversation_id", + "elapsed_ms", + "env_id", + "environment_id", + "error_code", + "event_type", + "is_resume", + "retryable", + "tool_count", + "tool_names", + "trace_id", + "user_id", +}) + +FORBIDDEN_PAYLOAD_KEYS: frozenset[str] = frozenset({ + "authorization", + "files", + "jwt", + "message", + "password", + "prompt", + "raw_output", + "secret", + "token", + "tool_input", + "tool_output", + "user_jwt", +}) + + +def _reduce_payload(payload: dict[str, Any] | None) -> dict[str, Any] | None: + """Strip sensitive fields from payload, keeping only SAFE_PAYLOAD_KEYS. + + The whitelist is case-normalized to prevent a caller bypassing redaction via + key spelling such as ``JWT`` or ``Tool_Names``. + """ + if payload is None: + return None + return { + key: value + for key, value in payload.items() + if isinstance(key, str) and key.lower() in SAFE_PAYLOAD_KEYS + } + + +# #region Schemas.AgentLifecycle.EventWriteRequest [C:1] [TYPE Class] [SEMANTICS agent,schema,lifecycle,write] +# @ingroup Schemas +class EventWriteRequest(BaseModel): + """Request body for POST /api/agent/events. + + payload is automatically reduced to SAFE_PAYLOAD_KEYS — sensitive fields are stripped. + """ + trace_id: str + conversation_id: str + environment_id: str | None = None + event_type: str + tool_name: str | None = None + status: str | None = None + elapsed_ms: float | None = None + payload: dict[str, Any] | None = None + error_code: str | None = None + + @classmethod + def with_reduced_payload(cls, **data) -> "EventWriteRequest": + """Create an instance, automatically reducing payload to safe keys.""" + if "payload" in data and data["payload"] is not None: + data["payload"] = _reduce_payload(data["payload"]) + return cls(**data) + + def model_post_init(self, __context: Any) -> None: + """Post-initialization hook: reduce payload to safe keys.""" + if self.payload is not None: + self.payload = _reduce_payload(self.payload) +# #endregion Schemas.AgentLifecycle.EventWriteRequest + + +# #region Schemas.AgentLifecycle.EventItem [C:1] [TYPE Class] [SEMANTICS agent,schema,lifecycle,item] +# @ingroup Schemas +class EventItem(BaseModel): + """A single lifecycle event returned in list responses.""" + id: str + trace_id: str + conversation_id: str + user_id: str + environment_id: str | None = None + event_type: str + tool_name: str | None = None + status: str | None = None + created_at: datetime + elapsed_ms: float | None = None + payload: dict[str, Any] | None = None + error_code: str | None = None + + _serialize_created_at = field_serializer("created_at")(_serialize_datetime) +# #endregion Schemas.AgentLifecycle.EventItem + + +# #region Schemas.AgentLifecycle.EventWriteResponse [C:1] [TYPE Class] [SEMANTICS agent,schema,lifecycle,write-response] +# @ingroup Schemas +class EventWriteResponse(BaseModel): + id: str + written: bool = True +# #endregion Schemas.AgentLifecycle.EventWriteResponse + + +# #region Schemas.AgentLifecycle.EventListResponse [C:1] [TYPE Class] [SEMANTICS agent,schema,lifecycle,list-response] +# @ingroup Schemas +class EventListResponse(BaseModel): + items: list[EventItem] + total: int + page: int + page_size: int + has_next: bool +# #endregion Schemas.AgentLifecycle.EventListResponse +# #endregion Schemas.AgentLifecycle diff --git a/backend/src/services/agent_lifecycle_retention.py b/backend/src/services/agent_lifecycle_retention.py new file mode 100644 index 000000000..190b6f477 --- /dev/null +++ b/backend/src/services/agent_lifecycle_retention.py @@ -0,0 +1,71 @@ +# backend/src/services/agent_lifecycle_retention.py +# #region AgentLifecycleRetentionService [C:3] [TYPE Module] [SEMANTICS agent,lifecycle,retention,maintenance] +# @BRIEF Delete old lifecycle audit events in bounded batches according to configured retention. +# @LAYER Service +# @RELATION DEPENDS_ON -> [Models.Agent.AgentLifecycleEvent] +# @INVARIANT Retention deletes only rows older than the configured UTC cutoff. +# @INVARIANT Deletion is bounded by batch_size so maintenance does not monopolize the transaction. + +from datetime import datetime, timedelta, timezone +import os + +from sqlalchemy import delete +from sqlalchemy.orm import Session + +from ..models.agent import AgentLifecycleEvent +from ss_tools.shared.cot_logger import log + +DEFAULT_RETENTION_DAYS = 90 +DEFAULT_BATCH_SIZE = 1000 + + +# #region AgentLifecycleRetentionService.get_retention_days [C:1] [TYPE Function] [SEMANTICS agent,lifecycle,retention,config] +# @BRIEF Read a positive lifecycle retention period from AGENT_LIFECYCLE_RETENTION_DAYS. +def get_retention_days() -> int: + raw_value = os.getenv("AGENT_LIFECYCLE_RETENTION_DAYS", str(DEFAULT_RETENTION_DAYS)) + try: + value = int(raw_value) + except (TypeError, ValueError): + return DEFAULT_RETENTION_DAYS + return value if value > 0 else DEFAULT_RETENTION_DAYS +# #endregion AgentLifecycleRetentionService.get_retention_days + + +# #region AgentLifecycleRetentionService.prune_events [C:3] [TYPE Function] [SEMANTICS agent,lifecycle,retention,delete] +# @ingroup AgentLifecycleRetentionService +# @BRIEF Delete one bounded batch of lifecycle events older than retention_days. +# @PRE db is an active session; retention_days and batch_size are positive integers. +# @POST Returns number of deleted rows and leaves commit ownership with the caller. +# @SIDE_EFFECT Deletes rows from agent_lifecycle_events; does not commit. +def prune_events( + db: Session, + retention_days: int | None = None, + batch_size: int = DEFAULT_BATCH_SIZE, + now: datetime | None = None, +) -> int: + effective_days = retention_days if retention_days is not None else get_retention_days() + if effective_days <= 0: + raise ValueError("retention_days must be positive") + if batch_size <= 0: + raise ValueError("batch_size must be positive") + + cutoff = (now or datetime.now(timezone.utc)).astimezone(timezone.utc) - timedelta(days=effective_days) + ids = [row[0] for row in db.query(AgentLifecycleEvent.id) + .filter(AgentLifecycleEvent.created_at < cutoff) + .order_by(AgentLifecycleEvent.created_at) + .limit(batch_size) + .all()] + if not ids: + return 0 + result = db.execute(delete(AgentLifecycleEvent).where(AgentLifecycleEvent.id.in_(ids))) + deleted = int(result.rowcount or 0) + log( + "AgentLifecycleRetentionService.prune_events", + "REFLECT", + "Lifecycle retention batch pruned", + {"deleted": deleted, "retention_days": effective_days}, + ) + return deleted +# #endregion AgentLifecycleRetentionService.prune_events + +# #endregion AgentLifecycleRetentionService diff --git a/backend/src/services/agent_lifecycle_service.py b/backend/src/services/agent_lifecycle_service.py new file mode 100644 index 000000000..6b90cba3f --- /dev/null +++ b/backend/src/services/agent_lifecycle_service.py @@ -0,0 +1,134 @@ +# backend/src/services/agent_lifecycle_service.py +# #region AgentLifecycleService [C:3] [TYPE Module] [SEMANTICS agent,lifecycle,service,crud] +# @BRIEF CRUD service for agent lifecycle events — write (immutable audit) + read (paginated, user-scoped). +# @LAYER Service +# @RELATION DEPENDS_ON -> [Models.Agent.AgentLifecycleEvent] +# @RELATION DEPENDS_ON -> [Schemas.AgentLifecycle] + +from datetime import datetime, timezone + +from sqlalchemy import desc +from sqlalchemy.orm import Session + +from ..models.agent import AgentLifecycleEvent +from ..schemas.agent_lifecycle import ( + EventItem, + EventListResponse, + EventWriteRequest, + EventWriteResponse, +) +from ss_tools.shared.cot_logger import log + + +def _now() -> datetime: + return datetime.now(timezone.utc) + + +# #region AgentLifecycleService.write_event [C:3] [TYPE Function] [SEMANTICS agent,lifecycle,write,audit] +# @ingroup AgentLifecycleService +# @BRIEF Persist a single lifecycle event as an immutable audit record. +# @PRE valid EventWriteRequest with reduced payload. +# @POST AgentLifecycleEvent row created in DB. Returns EventWriteResponse with new ID. +# @SIDE_EFFECT Writes to agent_lifecycle_events table. +# @INVARIANT Once written, audit fields are immutable — no update path exists. +def write_event(db: Session, req: EventWriteRequest, authenticated_user_id: str) -> EventWriteResponse: + """Write a lifecycle event. authenticated_user_id overrides any user_id in the request.""" + event = AgentLifecycleEvent( + trace_id=req.trace_id, + conversation_id=req.conversation_id, + user_id=authenticated_user_id, + environment_id=req.environment_id, + event_type=req.event_type, + tool_name=req.tool_name, + status=req.status, + elapsed_ms=req.elapsed_ms, + payload=req.payload, # already reduced by EventWriteRequest + error_code=req.error_code, + ) + db.add(event) + # The route owns the transaction boundary. Flush assigns the generated ID + # while allowing the caller to commit or roll back the complete request. + db.flush() + log("AgentLifecycleService.write_event", "REASON", "Lifecycle event written", + payload={"event_id": event.id, "event_type": event.event_type, "user_id": authenticated_user_id}) + return EventWriteResponse(id=event.id) +# #endregion AgentLifecycleService.write_event + + +# #region AgentLifecycleService.list_events [C:3] [TYPE Function] [SEMANTICS agent,lifecycle,list,query] +# @ingroup AgentLifecycleService +# @BRIEF Paginated list of lifecycle events, scoped to authenticated user or admin cross-user. +# @PRE db session, valid pagination params. +# @POST Returns EventListResponse with paginated items and total count. +# @SIDE_EFFECT Read-only query against agent_lifecycle_events table. +# @INVARIANT Non-admin users only see their own events (user_id filter enforced server-side). +def list_events( + db: Session, + authenticated_user_id: str, + is_admin: bool = False, + page: int = 1, + page_size: int = 50, + user_id: str | None = None, + event_type: str | None = None, + conversation_id: str | None = None, + status: str | None = None, + tool_name: str | None = None, +) -> EventListResponse: + """List lifecycle events with optional filters. + + Non-admin users see only their own events. Admin users may query by user_id. + """ + query = db.query(AgentLifecycleEvent) + + # Enforce user scope + if is_admin and user_id: + # Admin querying specific user + query = query.filter(AgentLifecycleEvent.user_id == user_id) + elif not is_admin: + # Non-admin users see only their own events + query = query.filter(AgentLifecycleEvent.user_id == authenticated_user_id) + # Admin without user_id: see all events (no user filter) + + # Optional filters + if event_type: + query = query.filter(AgentLifecycleEvent.event_type == event_type) + if conversation_id: + query = query.filter(AgentLifecycleEvent.conversation_id == conversation_id) + if status: + query = query.filter(AgentLifecycleEvent.status == status) + if tool_name: + query = query.filter(AgentLifecycleEvent.tool_name == tool_name) + + # Total before pagination + total = query.count() + + # Paginated, newest first + items = ( + query.order_by(desc(AgentLifecycleEvent.created_at)) + .offset((page - 1) * page_size) + .limit(page_size) + .all() + ) + + return EventListResponse( + items=[EventItem( + id=e.id, + trace_id=e.trace_id, + conversation_id=e.conversation_id, + user_id=e.user_id, + environment_id=e.environment_id, + event_type=e.event_type, + tool_name=e.tool_name, + status=e.status, + created_at=e.created_at, + elapsed_ms=e.elapsed_ms, + payload=e.payload, + error_code=e.error_code, + ) for e in items], + total=total, + page=page, + page_size=page_size, + has_next=(page * page_size) < total, + ) +# #endregion AgentLifecycleService.list_events +# #endregion AgentLifecycleService diff --git a/backend/tests/api/test_assistant_edge_cases.py b/backend/tests/api/test_assistant_edge_cases.py index 1e91b6011..8c0f2a67c 100644 --- a/backend/tests/api/test_assistant_edge_cases.py +++ b/backend/tests/api/test_assistant_edge_cases.py @@ -367,13 +367,13 @@ class TestDispatchDryRunSummary: mock_dry = MagicMock() mock_dry_cls.return_value = mock_dry - mock_dry.run.return_value = { + mock_dry.run = AsyncMock(return_value={ "summary": { "dashboards": {"create": 1, "update": 2, "delete": 0}, "charts": {"create": 3, "update": 4, "delete": 1}, "datasets": {"create": 0, "update": 0, "delete": 0}, } - } + }) text = await _async_confirmation_summary(intent, config_manager, db) assert "dry-run" in text.lower() or "dry" in text.lower() diff --git a/backend/tests/api/test_mappings.py b/backend/tests/api/test_mappings.py index 39bf4667e..9acfb50b2 100644 --- a/backend/tests/api/test_mappings.py +++ b/backend/tests/api/test_mappings.py @@ -130,6 +130,28 @@ class TestCreateMapping: "engine": "postgresql", } + @staticmethod + def _validated_config_and_clients(): + source_env = MagicMock(id="env-1") + target_env = MagicMock(id="env-2") + config_manager = MagicMock() + config_manager.get_environment.side_effect = { + "env-1": source_env, + "env-2": target_env, + }.get + + source_client = MagicMock() + source_client.get_databases_summary = AsyncMock( + return_value=[{"uuid": "uuid-1", "database_name": "Source DB"}] + ) + source_client.aclose = AsyncMock() + target_client = MagicMock() + target_client.get_databases_summary = AsyncMock( + return_value=[{"uuid": "uuid-2", "database_name": "Target DB"}] + ) + target_client.aclose = AsyncMock() + return config_manager, source_client, target_client + def test_create_new(self): mock_db = MagicMock() mock_db.query.return_value.filter.return_value.first.return_value = None @@ -140,8 +162,18 @@ class TestCreateMapping: mock_db.refresh.side_effect = _refresh from src.core.database import get_db - client = _make_client({get_db: lambda: mock_db}) - resp = client.post("/api/mappings", json=self.CREATE_PAYLOAD) + from src.dependencies import get_config_manager + + config_manager, source_client, target_client = self._validated_config_and_clients() + client = _make_client({ + get_db: lambda: mock_db, + get_config_manager: lambda: config_manager, + }) + with patch( + "src.api.routes.mappings.AsyncSupersetClient", + side_effect=[source_client, target_client], + ): + resp = client.post("/api/mappings", json=self.CREATE_PAYLOAD) assert resp.status_code == 200 mock_db.add.assert_called_once() mock_db.commit.assert_called_once() @@ -152,14 +184,86 @@ class TestCreateMapping: mock_db.query.return_value.filter.return_value.first.return_value = existing from src.core.database import get_db - client = _make_client({get_db: lambda: mock_db}) - resp = client.post("/api/mappings", json=self.CREATE_PAYLOAD) + from src.dependencies import get_config_manager + + config_manager, source_client, target_client = self._validated_config_and_clients() + client = _make_client({ + get_db: lambda: mock_db, + get_config_manager: lambda: config_manager, + }) + with patch( + "src.api.routes.mappings.AsyncSupersetClient", + side_effect=[source_client, target_client], + ): + resp = client.post("/api/mappings", json=self.CREATE_PAYLOAD) assert resp.status_code == 200 mock_db.add.assert_not_called() mock_db.commit.assert_called_once() assert existing.target_db_uuid == "uuid-2" assert existing.target_db_name == "Target DB" + def test_rejects_unknown_environment(self): + mock_db = MagicMock() + config_manager = MagicMock() + config_manager.get_environment.return_value = None + + from src.core.database import get_db + from src.dependencies import get_config_manager + + client = _make_client({ + get_db: lambda: mock_db, + get_config_manager: lambda: config_manager, + }) + resp = client.post("/api/mappings", json=self.CREATE_PAYLOAD) + + assert resp.status_code == 400 + assert resp.json()["detail"] == "Invalid source or target environment" + mock_db.add.assert_not_called() + + def test_rejects_source_database_from_another_environment(self): + mock_db = MagicMock() + config_manager, source_client, target_client = self._validated_config_and_clients() + payload = {**self.CREATE_PAYLOAD, "source_db_uuid": "target-only-uuid"} + + from src.core.database import get_db + from src.dependencies import get_config_manager + + client = _make_client({ + get_db: lambda: mock_db, + get_config_manager: lambda: config_manager, + }) + with patch( + "src.api.routes.mappings.AsyncSupersetClient", + side_effect=[source_client, target_client], + ): + resp = client.post("/api/mappings", json=payload) + + assert resp.status_code == 400 + assert "source environment" in resp.json()["detail"] + mock_db.add.assert_not_called() + + def test_rejects_target_database_from_another_environment(self): + mock_db = MagicMock() + config_manager, source_client, target_client = self._validated_config_and_clients() + payload = {**self.CREATE_PAYLOAD, "target_db_uuid": "source-only-uuid"} + + from src.core.database import get_db + from src.dependencies import get_config_manager + + client = _make_client({ + get_db: lambda: mock_db, + get_config_manager: lambda: config_manager, + }) + with patch( + "src.api.routes.mappings.AsyncSupersetClient", + side_effect=[source_client, target_client], + ): + resp = client.post("/api/mappings", json=payload) + + assert resp.status_code == 400 + assert "target environment" in resp.json()["detail"] + mock_db.add.assert_not_called() + # ── suggest_mappings_api ── diff --git a/backend/tests/api/test_migration.py b/backend/tests/api/test_migration.py index 74d32360d..5f12cfa42 100644 --- a/backend/tests/api/test_migration.py +++ b/backend/tests/api/test_migration.py @@ -185,14 +185,28 @@ class TestDryRunMigration: mock_config.get_environments.return_value = [e1, e2] mock_dry_run = MagicMock() - mock_dry_run.run = AsyncMock(return_value={"summary": "ok", "risks": []}) + mock_dry_run.run = AsyncMock(return_value={ + "generated_at": "2026-07-15T00:00:00+00:00", + "selection": self.SELECTION_PAYLOAD, + "selected_dashboard_titles": ["Sales", "Revenue"], + "diff": { + resource: {"create": [], "update": [], "delete": []} + for resource in ("dashboards", "charts", "datasets") + }, + "summary": { + resource: {"create": 0, "update": 0, "delete": 0} + for resource in ("dashboards", "charts", "datasets") + } | {"selected_dashboards": 2}, + "risk": {"score": 0, "level": "low", "items": []}, + }) from src.dependencies import get_config_manager with patch("src.api.routes.migration.MigrationDryRunService", return_value=mock_dry_run): client = _make_client({get_config_manager: lambda: mock_config}) resp = client.post("/api/migration/dry-run", json=self.SELECTION_PAYLOAD) assert resp.status_code == 200 - assert resp.json()["summary"] == "ok" + assert resp.json()["summary"]["selected_dashboards"] == 2 + assert resp.json()["risk"] == {"score": 0, "level": "low", "items": []} def test_identical_environments(self): mock_config = MagicMock() diff --git a/backend/tests/api/test_tasks.py b/backend/tests/api/test_tasks.py index 64a2b29c4..9c4039b4c 100644 --- a/backend/tests/api/test_tasks.py +++ b/backend/tests/api/test_tasks.py @@ -25,7 +25,7 @@ if _src not in sys.path: sys.path.insert(0, _src) -def _make_client(overrides: dict | None = None) -> TestClient: +def _make_client(overrides: dict | None = None, *, user=None) -> TestClient: from src.api.routes.tasks import router from src.core.database import get_db from src.dependencies import get_config_manager, get_current_user, get_task_manager @@ -34,7 +34,7 @@ def _make_client(overrides: dict | None = None) -> TestClient: app = FastAPI() app.include_router(router, prefix="/api/tasks") - mock_user = User( + mock_user = user or User( id="admin-1", username="admin", email="admin@x.com", auth_source="LOCAL", created_at=__import__("datetime").datetime.now(), @@ -58,11 +58,11 @@ class TestCreateTask: def test_success(self): mock_tm = MagicMock() - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_task = MagicMock(spec=Task) mock_task.id = "task-1" mock_task.plugin_id = "backup" - mock_task.status = "PENDING" + mock_task.status = TaskStatus.PENDING mock_task.params = {} mock_task.created_at = None mock_task.user_id = "admin" @@ -160,7 +160,7 @@ class TestListTasks: assert resp.status_code == 200 def test_logs_cleared(self): - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_tm = MagicMock() mock_task = MagicMock(spec=Task) mock_task.id = "task-1" @@ -181,11 +181,11 @@ class TestGetTask: def test_success(self): mock_tm = MagicMock() - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_task = MagicMock(spec=Task) mock_task.id = "task-123" mock_task.plugin_id = "backup" - mock_task.status = "PENDING" + mock_task.status = TaskStatus.PENDING mock_task.params = {} mock_task.created_at = None mock_task.user_id = "admin" @@ -207,6 +207,23 @@ class TestGetTask: resp = client.get("/api/tasks/unknown") assert resp.status_code == 404 + def test_passwords_are_excluded_from_response_params(self): + from src.core.task_manager import Task, TaskStatus + + mock_tm = MagicMock() + mock_tm.get_task.return_value = Task( + plugin_id="superset-migration", + params={"source_env_id": "source", "passwords": {"db": "secret"}}, + ) + + from src.dependencies import get_task_manager + client = _make_client({get_task_manager: lambda: mock_tm}) + + response = client.get("/api/tasks/task-123") + + assert response.status_code == 200 + assert response.json()["params"] == {"source_env_id": "source"} + # ── get_task_logs ── @@ -324,7 +341,7 @@ class TestResolveTask: def test_success(self): mock_tm = MagicMock() mock_tm.resolve_task = AsyncMock() - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_task = MagicMock(spec=Task) mock_task.id = "task-123" mock_tm.get_task.return_value = mock_task @@ -351,8 +368,8 @@ class TestResumeTask: def test_success(self): mock_tm = MagicMock() - mock_tm.resume_task_with_password = MagicMock() - from src.core.task_manager import Task + mock_tm.resume_task_with_password = AsyncMock() + from src.core.task_manager import Task, TaskStatus mock_task = MagicMock(spec=Task) mock_task.id = "task-123" mock_tm.get_task.return_value = mock_task @@ -364,13 +381,47 @@ class TestResumeTask: def test_value_error(self): mock_tm = MagicMock() - mock_tm.resume_task_with_password = MagicMock(side_effect=ValueError("Cannot resume")) + mock_tm.resume_task_with_password = AsyncMock(side_effect=ValueError("Cannot resume")) from src.dependencies import get_task_manager client = _make_client({get_task_manager: lambda: mock_tm}) resp = client.post("/api/tasks/task-123/resume", json={"passwords": {}}) assert resp.status_code == 400 + def test_other_user_cannot_resume_owned_task(self): + from src.core.task_manager import Task, TaskStatus + from src.schemas.auth import RoleSchema, User + + task = Task( + id="task-123", + plugin_id="superset-migration", + status=TaskStatus.AWAITING_INPUT, + input_request={"type": "database_password"}, + user_id="owner-1", + ) + mock_tm = MagicMock() + mock_tm.get_task.return_value = task + mock_tm.resume_task_with_password = AsyncMock() + other_user = User( + id="other-1", + username="other", + email="other@x.com", + auth_source="LOCAL", + created_at=__import__("datetime").datetime.now(), + roles=[RoleSchema(id="r1", name="Operator", description="", permissions=[])], + ) + + from src.dependencies import get_task_manager + client = _make_client( + {get_task_manager: lambda: mock_tm}, user=other_user + ) + response = client.post( + "/api/tasks/task-123/resume", json={"passwords": {"db": "pwd"}} + ) + + assert response.status_code == 403 + mock_tm.resume_task_with_password.assert_not_awaited() + # ── clear_tasks ── @@ -394,12 +445,12 @@ class TestClearTasks: assert resp.status_code == 204 def test_create_llm_documentation_task_with_provider(self): """Create task with llm_documentation — covers LLM provider block lines 80-123.""" - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_tm = MagicMock() mock_task = MagicMock(spec=Task) mock_task.id = "task-llm" mock_task.plugin_id = "llm_documentation" - mock_task.status = "PENDING" + mock_task.status = TaskStatus.PENDING mock_task.params = {"provider_id": "prov-1", "dashboard_id": 42} mock_task.created_at = None mock_task.user_id = "admin" @@ -445,12 +496,12 @@ class TestClearTasks: def test_create_llm_documentation_no_provider_id_resolved(self): """llm_documentation without provider_id — resolved from config binding (lines 87-108).""" - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_tm = MagicMock() mock_task = MagicMock(spec=Task) mock_task.id = "task-llm" mock_task.plugin_id = "llm_documentation" - mock_task.status = "PENDING" + mock_task.status = TaskStatus.PENDING mock_task.params = {"dashboard_id": 42} mock_task.created_at = None mock_task.user_id = "admin" @@ -488,12 +539,12 @@ class TestClearTasks: def test_create_llm_documentation_provider_from_binding(self): """provider_id resolved from binding (line 100).""" - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_tm = MagicMock() mock_task = MagicMock(spec=Task) mock_task.id = "task-llm" mock_task.plugin_id = "llm_documentation" - mock_task.status = "PENDING" + mock_task.status = TaskStatus.PENDING mock_task.params = {"dashboard_id": 42} mock_task.created_at = None mock_task.user_id = "admin" @@ -520,12 +571,12 @@ class TestClearTasks: def test_create_llm_documentation_provider_from_active(self): """provider_id resolved from active provider fallback (lines 107-108).""" - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_tm = MagicMock() mock_task = MagicMock(spec=Task) mock_task.id = "task-llm" mock_task.plugin_id = "llm_documentation" - mock_task.status = "PENDING" + mock_task.status = TaskStatus.PENDING mock_task.params = {"dashboard_id": 42} mock_task.created_at = None mock_task.user_id = "admin" @@ -560,10 +611,10 @@ class TestRetryTask: def test_success(self): mock_tm = MagicMock() - from src.core.task_manager import Task + from src.core.task_manager import Task, TaskStatus mock_task = MagicMock(spec=Task) mock_task.id = "task-failed-1" - mock_task.status = "PENDING" # after retry reset + mock_task.status = TaskStatus.PENDING # after retry reset mock_tm.retry_task = AsyncMock(return_value=mock_task) mock_tm.get_task.return_value = mock_task diff --git a/backend/tests/core/superset_client/test_client_dashboards_crud_edge3.py b/backend/tests/core/superset_client/test_client_dashboards_crud_edge3.py index 6eb929bab..6b4531a8a 100644 --- a/backend/tests/core/superset_client/test_client_dashboards_crud_edge3.py +++ b/backend/tests/core/superset_client/test_client_dashboards_crud_edge3.py @@ -115,7 +115,7 @@ class TestImportDashboardEdgeCases: call_count = 0 - async def mock_do_import(file_name): + async def mock_do_import(file_name, passwords=None): nonlocal call_count call_count += 1 if call_count == 1: diff --git a/backend/tests/core/test_scheduler.py b/backend/tests/core/test_scheduler.py index 441487561..a093a333d 100644 --- a/backend/tests/core/test_scheduler.py +++ b/backend/tests/core/test_scheduler.py @@ -945,4 +945,108 @@ def test_calculate_schedule_small_window_warning(): assert result[0] == datetime(2026, 6, 11, 10, 0, 0) # #endregion test_calculate_schedule_small_window_warning + +# ══════════════════════════════════════════════════════════════ +# execute_scheduled_backup lifecycle log +# ══════════════════════════════════════════════════════════════ + +# #region test_execute_scheduled_backup_lifecycle_log [C:2] [TYPE Function] +# @BRIEF execute_scheduled_backup emits lifecycle execution log. +def test_execute_scheduled_backup_lifecycle_log(): + """Lifecycle: execute_scheduled_backup logs execution.""" + with patch("src.core.scheduler.logger.reason") as mock_log: + with patch("src.dependencies.get_scheduler_service") as mock_get: + svc = MagicMock() + mock_get.return_value = svc + from src.core.scheduler import execute_scheduled_backup + execute_scheduled_backup("env-1") + # Should log lifecycle execution + mock_log.assert_called_once() + args, kwargs = mock_log.call_args + assert "Scheduler lifecycle: backup executed" in args[0] + assert kwargs["payload"]["env_id"] == "env-1" +# #endregion test_execute_scheduled_backup_lifecycle_log + + +# #region test_execute_scheduled_validation_lifecycle_log [C:2] [TYPE Function] +# @BRIEF execute_scheduled_validation emits lifecycle execution log. +def test_execute_scheduled_validation_lifecycle_log(): + """Lifecycle: execute_scheduled_validation logs execution.""" + with patch("src.core.scheduler.logger.reason") as mock_log: + with patch("src.dependencies.get_scheduler_service") as mock_get: + svc = MagicMock() + mock_get.return_value = svc + from src.core.scheduler import execute_scheduled_validation + execute_scheduled_validation("pol-1") + mock_log.assert_called_once() + args, kwargs = mock_log.call_args + assert "Scheduler lifecycle: validation executed" in args[0] + assert kwargs["payload"]["policy_id"] == "pol-1" +# #endregion test_execute_scheduled_validation_lifecycle_log + + +# #region test_start_emits_lifecycle_logs [C:2] [TYPE Function] +# @BRIEF SchedulerService.start emits lifecycle structured logs for started + restored jobs. +def test_start_emits_lifecycle_logs(service, mock_scheduler): + """Lifecycle: start() logs 'started' and attempts 'restored' counts.""" + mock_scheduler.running = False + # Return job-like objects with .id attribute for the restored jobs log + job1 = MagicMock() + job1.id = "backup_env-1" + job2 = MagicMock() + job2.id = "translate_sched-1" + mock_scheduler.get_jobs.return_value = [job1, job2] + with patch("src.core.scheduler.logger.reason") as mock_log: + with patch("src.core.scheduler.belief_scope") as mock_bs: + mock_bs.return_value.__enter__ = MagicMock() + mock_bs.return_value.__exit__ = MagicMock(return_value=False) + with patch.object(service, "load_schedules"): + service.start() + # Verify 'started' lifecycle log + started_calls = [c for c in mock_log.call_args_list if "started" in c[0][0]] + assert len(started_calls) >= 1 + # Verify 'restored' lifecycle log + restored_calls = [c for c in mock_log.call_args_list if "restored" in c[0][0]] + assert len(restored_calls) >= 1 +# #endregion test_start_emits_lifecycle_logs + + +# #region test_trigger_backup_lifecycle_markers [C:2] [TYPE Function] +# @BRIEF _trigger_backup emits executed/skipped lifecycle markers. +def test_trigger_backup_lifecycle_markers(service, mock_task_manager): + """Lifecycle: _trigger_backup logs 'triggered' and 'task submitted' markers.""" + mock_task_manager.get_tasks.return_value = [] + with patch("src.core.scheduler.seed_trace_id"): + with patch("src.core.scheduler.logger.reason") as mock_log: + with patch("src.core.scheduler.belief_scope") as mock_bs: + mock_bs.return_value.__enter__ = MagicMock() + mock_bs.return_value.__exit__ = MagicMock(return_value=False) + with patch("src.core.scheduler.asyncio.run_coroutine_threadsafe"): + service._trigger_backup("env-1") + # Should log 'triggered' and 'task submitted' + triggered_calls = [c for c in mock_log.call_args_list if "triggered" in c[0][0]] + submitted_calls = [c for c in mock_log.call_args_list if "task submitted" in c[0][0] or "submitted" in c[0][0]] + assert len(triggered_calls) >= 1 + assert len(submitted_calls) >= 1 +# #endregion test_trigger_backup_lifecycle_markers + + +# #region test_trigger_backup_lifecycle_skipped [C:2] [TYPE Function] +# @BRIEF _trigger_backup emits 'skipped' lifecycle marker when backup already running. +def test_trigger_backup_lifecycle_skipped(service, mock_task_manager): + """Lifecycle: _trigger_backup logs 'skipped' marker when backup already running.""" + running_task = _build_task(plugin_id="superset-backup", status="RUNNING", environment_id="env-1") + mock_task_manager.get_tasks.return_value = [running_task] + with patch("src.core.scheduler.seed_trace_id"): + with patch("src.core.scheduler.logger.reason") as mock_log: + with patch("src.core.scheduler.belief_scope") as mock_bs: + mock_bs.return_value.__enter__ = MagicMock() + mock_bs.return_value.__exit__ = MagicMock(return_value=False) + with patch("src.core.scheduler.asyncio.run_coroutine_threadsafe") as mock_run: + service._trigger_backup("env-1") + mock_run.assert_not_called() + # Should log 'skipped' + skipped_calls = [c for c in mock_log.call_args_list if "skipped" in c[0][0]] + assert len(skipped_calls) >= 1 +# #endregion test_trigger_backup_lifecycle_skipped # #endregion Test.Scheduler diff --git a/backend/tests/plugins/test_migration_plugin.py b/backend/tests/plugins/test_migration_plugin.py index 5d6a409b6..a42946efb 100644 --- a/backend/tests/plugins/test_migration_plugin.py +++ b/backend/tests/plugins/test_migration_plugin.py @@ -429,7 +429,7 @@ class TestMigrationPluginExecute: mock_task_manager.get_task.return_value = MagicMock( params={"passwords": {"PostgreSQL": "secret123"}} ) - mock_task_manager.await_input = MagicMock() + mock_task_manager.await_input = AsyncMock() mock_task_manager.wait_for_input = AsyncMock() mock_src_client = _make_mock_superset_client() @@ -716,7 +716,7 @@ class TestMigrationPluginExecute: mock_task_manager.get_task.return_value = MagicMock( params={"passwords": {"PostgreSQL": "secret123"}} ) - mock_task_manager.await_input = MagicMock() + mock_task_manager.await_input = AsyncMock() mock_task_manager.wait_for_input = AsyncMock() mock_src_client = _make_mock_superset_client() diff --git a/backend/tests/plugins/translate/test_batch_proc.py b/backend/tests/plugins/translate/test_batch_proc.py index d52d537d4..a3fe8a89c 100644 --- a/backend/tests/plugins/translate/test_batch_proc.py +++ b/backend/tests/plugins/translate/test_batch_proc.py @@ -10,7 +10,7 @@ import pytest from unittest.mock import AsyncMock, MagicMock, patch from src.models.translate import ( - TranslationBatch, TranslationJob, TranslationLanguage, + TranslationBatch, TranslationEvent, TranslationJob, TranslationLanguage, TranslationRecord, ) from src.plugins.translate._batch_proc import BatchProcessingService @@ -406,6 +406,97 @@ class TestProcessBatch: assert result["successful"] == 2 assert result.get("cache_hits", 0) == 2 + @pytest.mark.asyncio + async def test_process_batch_emits_language_detection_event(self, db_with_run): + """Observability: process_batch emits LANGUAGE_DETECTION_COMPLETED aggregate event.""" + session, run_id = db_with_run + config = MagicMock() + svc = BatchProcessingService(session, config) + job = _make_job(session, target_langs=["fr"]) + rows = _batch_rows(2, detected_lang="en") + + with patch.object(svc, '_process_llm', + new=AsyncMock(return_value={"successful": 2, "failed": 0, + "skipped": 0, "retries": 0})): + with patch('src.plugins.translate._batch_proc.DictionaryManager.filter_for_batch', + return_value=[]): + with patch('src.plugins.translate._batch_proc.batch_detect', + return_value=["en", "en"]): + await svc.process_batch( + job=job, run_id=run_id, batch_index=0, batch_rows=rows, + ) + + events = session.query(TranslationEvent).filter( + TranslationEvent.event_type == "LANGUAGE_DETECTION_COMPLETED" + ).all() + assert len(events) >= 1 + event_data = events[-1].event_data or {} + assert event_data.get("detector") == "lingua" + assert event_data.get("total_rows") == 2 + assert event_data.get("target_language_count") == 1 # ["fr"] + assert event_data.get("source_language_distribution") == {"en": 2} + assert event_data.get("und_count") == 0 + # No source text in payload + assert "source_text" not in event_data + + @pytest.mark.asyncio + async def test_process_batch_emits_cache_summary(self, db_with_run): + """Observability: process_batch emits CACHE_SUMMARY aggregate event.""" + session, run_id = db_with_run + config = MagicMock() + svc = BatchProcessingService(session, config) + job = _make_job(session, target_langs=["fr"]) + rows = _batch_rows(2, detected_lang="en") + + with patch.object(svc, '_process_llm', + new=AsyncMock(return_value={"successful": 2, "failed": 0, + "skipped": 0, "retries": 0})): + with patch('src.plugins.translate._batch_proc.DictionaryManager.filter_for_batch', + return_value=[]): + await svc.process_batch( + job=job, run_id=run_id, batch_index=0, batch_rows=rows, + ) + + events = session.query(TranslationEvent).filter( + TranslationEvent.event_type == "CACHE_SUMMARY" + ).all() + assert len(events) >= 1 + event_data = events[-1].event_data or {} + assert event_data.get("batch_size") == 2 + assert event_data.get("eligible") == 2 + assert event_data.get("hits") == 0 + assert event_data.get("misses") == 2 + assert "source_text" not in event_data + + @pytest.mark.asyncio + async def test_process_batch_cache_summary_with_hits(self, db_with_run): + """Observability: CACHE_SUMMARY hits count reflects cached rows.""" + session, run_id = db_with_run + config = MagicMock() + svc = BatchProcessingService(session, config) + job = _make_job(session, target_langs=["fr"]) + rows = _batch_rows(2, detected_lang="en") + for r in rows: + r["_cached_lang_values"] = {"fr": "Bonjour"} + + with patch.object(svc, '_process_llm', + new=AsyncMock(return_value={"successful": 2, "failed": 0, + "skipped": 0, "retries": 0})): + with patch('src.plugins.translate._batch_proc.DictionaryManager.filter_for_batch', + return_value=[]): + await svc.process_batch( + job=job, run_id=run_id, batch_index=0, batch_rows=rows, + ) + + events = session.query(TranslationEvent).filter( + TranslationEvent.event_type == "CACHE_SUMMARY" + ).all() + assert len(events) >= 1 + event_data = events[-1].event_data or {} + assert event_data["eligible"] == 2 + assert event_data["hits"] == 2 + assert event_data["misses"] == 0 + class TestProcessLlmProviderEdgeCases: """Cover provider_id edge cases in _process_llm.""" diff --git a/backend/tests/plugins/translate/test_events.py b/backend/tests/plugins/translate/test_events.py index 92cb5a601..7f5822604 100644 --- a/backend/tests/plugins/translate/test_events.py +++ b/backend/tests/plugins/translate/test_events.py @@ -294,4 +294,89 @@ class TestGetRunEventSummary: summary = el.get_run_event_summary(run_id) assert summary["terminal_event_count"] > 1 assert summary["invariant_valid"] is False + + +class TestNewObservabilityEvents: + """Verify new aggregate observability events — LANGUAGE_DETECTION_COMPLETED and CACHE_SUMMARY.""" + + def test_language_detection_completed_is_valid(self): + """LANGUAGE_DETECTION_COMPLETED is a recognized event type.""" + assert "LANGUAGE_DETECTION_COMPLETED" in VALID_EVENT_TYPES + + def test_cache_summary_is_valid(self): + """CACHE_SUMMARY is a recognized event type.""" + assert "CACHE_SUMMARY" in VALID_EVENT_TYPES + + def test_new_events_not_terminal(self): + """Neither new event type is a terminal event (no invariant conflict).""" + assert "LANGUAGE_DETECTION_COMPLETED" not in TERMINAL_EVENT_TYPES + assert "CACHE_SUMMARY" not in TERMINAL_EVENT_TYPES + + def test_log_language_detection_completed(self, db_with_run): + """Happy: LANGUAGE_DETECTION_COMPLETED event with aggregate language stats (no source text).""" + session, run_id = db_with_run + el = TranslationEventLog(session) + event = el.log_event( + job_id=JOB_ID, run_id=run_id, + event_type="LANGUAGE_DETECTION_COMPLETED", + payload={ + "detector": "lingua", + "source_language_distribution": {"en": 5, "fr": 3}, + "und_count": 1, + "und_rate": 0.1111, + "total_rows": 9, + "target_language_count": 2, + }, + ) + assert event.event_type == "LANGUAGE_DETECTION_COMPLETED" + assert event.event_data["detector"] == "lingua" + assert event.event_data["und_count"] == 1 + assert event.event_data["total_rows"] == 9 + # No source text in payload + assert "source_text" not in event.event_data + + def test_log_cache_summary(self, db_with_run): + """Happy: CACHE_SUMMARY event with cache eligibility/hits/misses.""" + session, run_id = db_with_run + el = TranslationEventLog(session) + event = el.log_event( + job_id=JOB_ID, run_id=run_id, + event_type="CACHE_SUMMARY", + payload={ + "batch_size": 10, + "eligible": 8, + "hits": 3, + "misses": 5, + }, + ) + assert event.event_type == "CACHE_SUMMARY" + assert event.event_data["eligible"] == 8 + assert event.event_data["hits"] == 3 + assert event.event_data["misses"] == 5 + # No source text in payload + assert "source_text" not in event.event_data + + def test_log_language_detection_completed_no_run(self, db_session): + """Edge: LANGUAGE_DETECTION_COMPLETED can be logged without run_id (job-level).""" + el = TranslationEventLog(db_session) + event = el.log_event( + job_id=JOB_ID, + event_type="LANGUAGE_DETECTION_COMPLETED", + payload={"detector": "lingua", "total_rows": 0}, + ) + assert event.run_id is None + assert event.event_type == "LANGUAGE_DETECTION_COMPLETED" + + def test_cache_summary_terminal_invariant_preserved(self, db_with_run): + """Invariant: CACHE_SUMMARY does not break terminal event invariant.""" + session, run_id = db_with_run + el = TranslationEventLog(session) + el.log_event(job_id=JOB_ID, run_id=run_id, event_type="RUN_STARTED") + # Log cache event — should succeed + el.log_event(job_id=JOB_ID, run_id=run_id, event_type="CACHE_SUMMARY") + # Log terminal event — should still succeed + el.log_event(job_id=JOB_ID, run_id=run_id, event_type="RUN_COMPLETED") + summary = el.get_run_event_summary(run_id) + assert summary["invariant_valid"] is True + assert summary["terminal_event_count"] == 1 # #endregion Test.TranslationEventLog diff --git a/backend/tests/test_dependencies_unit.py b/backend/tests/test_dependencies_unit.py index 6a8c6e635..38476468e 100644 --- a/backend/tests/test_dependencies_unit.py +++ b/backend/tests/test_dependencies_unit.py @@ -314,7 +314,6 @@ class TestGetCurrentUser: get_current_user("token", MagicMock()) assert exc.value.status_code == 401 - @pytest.mark.skip(reason="get_current_user signature changed (db via Depends) — test needs async FastAPI test client") def test_get_current_user_blacklisted(self): """Blacklisted token raises 401.""" from src.dependencies import get_current_user @@ -325,6 +324,17 @@ class TestGetCurrentUser: get_current_user("revoked", MagicMock()) assert exc.value.status_code == 401 + def test_x_user_jwt_blacklist_is_checked_for_effective_token(self): + """A delegated end-user token cannot bypass blacklist validation.""" + from src.dependencies import get_current_user + + with patch('src.dependencies.decode_token', return_value={"sub": "user"}), \ + patch('src.dependencies.is_token_blacklisted', return_value=True) as is_blacklisted: + with pytest.raises(HTTPException) as exc: + get_current_user("service-token", "revoked-user-token", MagicMock()) + assert exc.value.status_code == 401 + is_blacklisted.assert_called_once_with("revoked-user-token") + def test_get_current_user_not_found(self): """User not in DB raises 401.""" from src.dependencies import get_current_user diff --git a/backend/tests/test_gradio_proxy_config.py b/backend/tests/test_gradio_proxy_config.py index d3f8f3fa1..2df01da25 100644 --- a/backend/tests/test_gradio_proxy_config.py +++ b/backend/tests/test_gradio_proxy_config.py @@ -6,9 +6,10 @@ # @TEST_EDGE: docker_dns_name -> Proxy targets compose service name `agent`. # @TEST_EDGE: prefix_forwarding -> Proxy preserves `/api/agent/gradio` for Gradio root_path. # @TEST_EDGE: auth_forwarding -> Browser Authorization header is forwarded. +# @TEST_INVARIANT: trace_forwarding -> Valid browser trace IDs cross nginx into backend and agent. +# @TEST_INVARIANT: proxy_access_events -> nginx stdout contains one JSON access event per request. from pathlib import Path - PROJECT_ROOT = Path(__file__).resolve().parents[2] @@ -24,6 +25,7 @@ def test_nginx_http_gradio_proxy_targets_agent_service_and_preserves_prefix(): assert "superset-tools-agent" not in text assert "rewrite ^/api/agent/gradio" not in text assert "proxy_set_header Authorization $http_authorization;" in text + assert "proxy_set_header X-Trace-ID $http_x_trace_id;" in text assert "proxy_set_header Connection $connection_upgrade;" in text assert "proxy_redirect http://agent:7860/ /api/agent/gradio/;" in text @@ -34,12 +36,25 @@ def test_nginx_ssl_gradio_proxy_matches_http_contract(): assert "map $http_upgrade $connection_upgrade" in text assert "location /api/agent/gradio/" in text assert "set $agent_api http://agent:7860;" in text + assert "proxy_set_header X-Trace-ID $http_x_trace_id;" in text assert "superset-tools-agent" not in text assert "rewrite ^/api/agent/gradio" not in text assert "proxy_set_header Authorization $http_authorization;" in text assert "proxy_set_header X-Forwarded-Proto https;" in text +def test_nginx_proxy_emits_json_access_logs_and_forwards_trace_ids(): + """Both nginx modes must make request correlation available on container stdout.""" + for relative_path in ("docker/nginx.conf", "docker/nginx.ssl.conf"): + text = _read(relative_path) + assert "log_format observability_json escape=json" in text + assert "access_log /dev/stdout observability_json;" in text + assert "error_log /dev/stderr warn;" in text + assert '"trace_id":"$http_x_trace_id"' in text + assert '"upstream_response_time":"$upstream_response_time"' in text + assert text.count("proxy_set_header X-Trace-ID $http_x_trace_id;") == 3 + + def test_compose_frontend_depends_on_agent_and_agent_uses_root_path(): """Compose starts agent before frontend and configures Gradio root_path.""" for compose_path in ("docker-compose.yml", "docker-compose.enterprise-clean.yml"): diff --git a/backend/tests/test_task_manager.py b/backend/tests/test_task_manager.py index 90c4199aa..7bbf5ef9e 100644 --- a/backend/tests/test_task_manager.py +++ b/backend/tests/test_task_manager.py @@ -639,6 +639,25 @@ class TestTaskManagerInput: finally: _cleanup_manager(mgr) + @pytest.mark.asyncio + async def test_resume_rejects_different_task_owner(self): + mgr, _, _, _ = _make_manager() + try: + from src.core.task_manager.models import Task, TaskStatus + + task = Task(plugin_id="p1", params={}, user_id="owner-1") + task.status = TaskStatus.AWAITING_INPUT + mgr.tasks[task.id] = task + + with pytest.raises(PermissionError, match="different user"): + await mgr.resume_task_with_password( + task.id, + {"db": "pass"}, + requester_user_id="other-1", + ) + finally: + _cleanup_manager(mgr) + def test_task_model_new_retry_and_progress_fields(): """Task Pydantic model supports new fields from improvements (retry, progress).""" diff --git a/docker/frontend.entrypoint.sh b/docker/frontend.entrypoint.sh index 7b5a4ce8d..1a27d9e59 100644 --- a/docker/frontend.entrypoint.sh +++ b/docker/frontend.entrypoint.sh @@ -9,7 +9,7 @@ set -euo pipefail # @POST Если найдены server.crt + server.key — активирован SSL конфиг. # Если есть server.p12 — извлечены .crt и .key. # Если server.key зашифрован + есть SSL_KEY_PASSPHRASE — ключ расшифрован. -# CA-сертификаты из /opt/certs установлены в системное хранилище Alpine. +# CA-сертификаты из /opt/certs установлены в системное хранилище Alpine; nginx config validates before serving. # @SIDE_EFFECT Модифицирует /usr/local/share/ca-certificates/; перезаписывает nginx конфиг. # Расшифровывает приватный ключ в /etc/nginx/ssl/ при наличии passphrase. # @LAYER Infrastructure @@ -252,6 +252,9 @@ ssl_passphrase="$(resolve_ssl_passphrase)" select_nginx_config decrypt_key_if_needed "$ssl_passphrase" +# Fail before serving traffic when an observability config edit made nginx invalid. +nginx -t + # ── Start nginx ────────────────────────────────────────────────────── echo "[entrypoint] Starting nginx..." exec "$@" diff --git a/docker/nginx.conf b/docker/nginx.conf index 637eb2b3e..323a16541 100644 --- a/docker/nginx.conf +++ b/docker/nginx.conf @@ -1,8 +1,39 @@ +# #region docker.nginx.conf [C:3] [TYPE Module] [SEMANTICS nginx,proxy,frontend,api,ws,observability] +# @BRIEF HTTP-only nginx config — serves SvelteKit static files, proxies API/WS, and emits correlated JSON access logs. +# @PRE Docker DNS resolves backend:8000; built SvelteKit output at /usr/share/nginx/html. +# @POST Static SPA fallback works; /api/ and /ws/ reach backend with upgrade headers. +# Proxy access records are JSON on stdout and include client X-Trace-ID/timings. +# @LAYER Infrastructure +# @RELATION DEPENDS_ON -> [docker.frontend.entrypoint] +# @RELATION BINDS_TO -> [backend:8000] +# @INVARIANT /api/ and /ws/ preserve request URI, Authorization, and X-Trace-ID headers. +# @RATIONALE stdout JSON logs survive container log aggregation and make nginx/backend/agent event correlation queryable. +# @REJECTED Nginx HTTPS in this file rejected — SSL variant selected only when certificates exist. +# #endregion docker.nginx.conf + map $http_upgrade $connection_upgrade { default upgrade; '' close; } +# JSON access records are container-stdout friendly and correlate proxy, backend, +# and Gradio agent activity through the client-provided X-Trace-ID header. +log_format observability_json escape=json '{' + '"ts":"$time_iso8601",' + '"trace_id":"$http_x_trace_id",' + '"request_id":"$request_id",' + '"remote_addr":"$remote_addr",' + '"method":"$request_method",' + '"uri":"$uri",' + '"status":$status,' + '"upstream_status":"$upstream_status",' + '"request_time":$request_time,' + '"upstream_response_time":"$upstream_response_time"' + '}'; + +access_log /dev/stdout observability_json; +error_log /dev/stderr warn; + server { listen 80; server_name _; @@ -27,6 +58,7 @@ server { proxy_set_header X-Real-IP $remote_addr; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_set_header X-Forwarded-Proto $scheme; + proxy_set_header X-Trace-ID $http_x_trace_id; proxy_redirect http://agent:7860/ /api/agent/gradio/; proxy_read_timeout 300s; } @@ -39,6 +71,7 @@ server { proxy_set_header X-Real-IP $remote_addr; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_set_header X-Forwarded-Proto $scheme; + proxy_set_header X-Trace-ID $http_x_trace_id; proxy_read_timeout 300s; } @@ -52,6 +85,7 @@ server { proxy_set_header X-Real-IP $remote_addr; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_set_header X-Forwarded-Proto $scheme; + proxy_set_header X-Trace-ID $http_x_trace_id; proxy_read_timeout 300s; } } diff --git a/docker/nginx.ssl.conf b/docker/nginx.ssl.conf index 862965b9a..92a370e2d 100644 --- a/docker/nginx.ssl.conf +++ b/docker/nginx.ssl.conf @@ -1,15 +1,18 @@ # #region docker.nginx.ssl.conf [C:3] [TYPE Module] [SEMANTICS nginx,ssl,https,enterprise] -# @BRIEF Nginx config с опциональным SSL — HTTP (порт 80) + HTTPS (порт 443). +# @BRIEF Nginx config с опциональным SSL — HTTP (порт 80) + HTTPS (порт 443), JSON access logs and trace forwarding. # Используется frontend.entrypoint.sh, если найдены server.crt + server.key. # @PRE SSL сертификаты: /etc/nginx/ssl/server.crt и /etc/nginx/ssl/server.key. # Внутренний Docker DNS: backend:8000 для API/WS прокси. # @POST HTTP (80) редиректит на HTTPS (443) если SSL включён. +# Proxy access records are JSON on stdout and include client X-Trace-ID/timings. # HTTPS (443) сервит статику + прокси /api/ и /ws/ на backend. # @LAYER Infrastructure # @RELATION DEPENDS_ON -> [docker/frontend.entrypoint.sh] # @RELATION DEPENDS_ON -> [docker/nginx.conf] # @INVARIANT Порт 80 всегда открыт (fallback для HTTP-only окружений). -# @INVARIANT /api/ и /ws/ проксируются на http://backend:8000 (без TLS внутри compose-сети). +# @INVARIANT /api/ и /ws/ проксируются на http://backend:8000 (без TLS внутри compose-сети), +# preserving Authorization and X-Trace-ID headers. +# @RATIONALE stdout JSON logs survive container log aggregation and make nginx/backend/agent event correlation queryable. # @REJECTED Прокси на HTTPS внутри compose-сети отвергнут — избыточно. # #endregion docker.nginx.ssl.conf @@ -18,6 +21,23 @@ map $http_upgrade $connection_upgrade { '' close; } +# Keep TLS deployment observability identical to HTTP-only deployment. +log_format observability_json escape=json '{' + '"ts":"$time_iso8601",' + '"trace_id":"$http_x_trace_id",' + '"request_id":"$request_id",' + '"remote_addr":"$remote_addr",' + '"method":"$request_method",' + '"uri":"$uri",' + '"status":$status,' + '"upstream_status":"$upstream_status",' + '"request_time":$request_time,' + '"upstream_response_time":"$upstream_response_time"' + '}'; + +access_log /dev/stdout observability_json; +error_log /dev/stderr warn; + server { listen 80 default_server; server_name _; @@ -56,6 +76,7 @@ server { proxy_set_header X-Real-IP $remote_addr; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_set_header X-Forwarded-Proto https; + proxy_set_header X-Trace-ID $http_x_trace_id; proxy_redirect http://agent:7860/ /api/agent/gradio/; proxy_read_timeout 300s; } @@ -68,6 +89,7 @@ server { proxy_set_header X-Real-IP $remote_addr; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_set_header X-Forwarded-Proto https; + proxy_set_header X-Trace-ID $http_x_trace_id; proxy_read_timeout 300s; } @@ -81,6 +103,7 @@ server { proxy_set_header X-Real-IP $remote_addr; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_set_header X-Forwarded-Proto https; + proxy_set_header X-Trace-ID $http_x_trace_id; proxy_read_timeout 300s; } } diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index ecba6730d..dd15cf9bd 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -119,9 +119,22 @@ function shouldSuppressApiErrorToast(endpoint: string, error: ApiError): boolean // #region wsUrlHelpers [C:2] [TYPE Block] [SEMANTICS websocket, url, task-logs, maintenance, translate] // @BRIEF WebSocket URL builders — each constructs an authenticated WS endpoint for a specific channel. // @PRE taskId / runId are non-empty strings where applicable. -// @POST Returns fully-qualified ws:// or wss:// URL with auth token as query parameter. +// @POST Returns fully-qualified ws:// or wss:// URL with auth and x-trace-id query parameters. // @SIDE_EFFECT Reads localStorage for auth_token on each call. -// @RATIONALE WebSocket API does not support custom headers in browser, so auth token is appended as query param. +// @RATIONALE WebSocket API does not support custom headers in browser, so auth and trace IDs are +// appended as query parameters for authentication and cross-stack correlation. + +function _appendWsCredentials(url: string): string { + const params = new URLSearchParams(); + if (typeof window !== 'undefined') { + const token = localStorage.getItem('auth_token'); + if (token) params.set('token', token); + const traceId = getTraceId(); + if (traceId && traceId !== 'no-trace') params.set('x-trace-id', traceId); + } + const query = params.toString(); + return query ? `${url}?${query}` : url; +} /** * Build a WebSocket URL for a task's log stream, including auth token. @@ -129,14 +142,7 @@ function shouldSuppressApiErrorToast(endpoint: string, error: ApiError): boolean export const getWsUrl = (taskId: string): string => { const protocol = typeof window !== 'undefined' && window.location.protocol === 'https:' ? 'wss:' : 'ws:'; const host = typeof window !== 'undefined' ? window.location.host : 'localhost:8000'; - let url = `${protocol}//${host}/ws/logs/${taskId}`; - if (typeof window !== 'undefined') { - const token = localStorage.getItem('auth_token'); - if (token) { - url += `?token=${encodeURIComponent(token)}`; - } - } - return url; + return _appendWsCredentials(`${protocol}//${host}/ws/logs/${taskId}`); }; /** @@ -145,14 +151,7 @@ export const getWsUrl = (taskId: string): string => { export const getTaskEventsWsUrl = (): string => { const protocol = typeof window !== 'undefined' && window.location.protocol === 'https:' ? 'wss:' : 'ws:'; const host = typeof window !== 'undefined' ? window.location.host : 'localhost:8000'; - let url = `${protocol}//${host}/ws/task-events`; - if (typeof window !== 'undefined') { - const token = localStorage.getItem('auth_token'); - if (token) { - url += `?token=${encodeURIComponent(token)}`; - } - } - return url; + return _appendWsCredentials(`${protocol}//${host}/ws/task-events`); }; /** @@ -162,14 +161,7 @@ export const getTaskEventsWsUrl = (): string => { export const getMaintenanceEventsWsUrl = (): string => { const protocol = typeof window !== 'undefined' && window.location.protocol === 'https:' ? 'wss:' : 'ws:'; const host = typeof window !== 'undefined' ? window.location.host : 'localhost:8000'; - let url = `${protocol}//${host}/ws/maintenance/events`; - if (typeof window !== 'undefined') { - const token = localStorage.getItem('auth_token'); - if (token) { - url += `?token=${encodeURIComponent(token)}`; - } - } - return url; + return _appendWsCredentials(`${protocol}//${host}/ws/maintenance/events`); }; // #endregion getMaintenanceEventsWsUrl @@ -180,14 +172,7 @@ export const getMaintenanceEventsWsUrl = (): string => { export const getTranslateRunWsUrl = (runId: string): string => { const protocol = typeof window !== 'undefined' && window.location.protocol === 'https:' ? 'wss:' : 'ws:'; const host = typeof window !== 'undefined' ? window.location.host : 'localhost:8000'; - let url = `${protocol}//${host}/ws/translate/run/${runId}`; - if (typeof window !== 'undefined') { - const token = localStorage.getItem('auth_token'); - if (token) { - url += `?token=${encodeURIComponent(token)}`; - } - } - return url; + return _appendWsCredentials(`${protocol}//${host}/ws/translate/run/${runId}`); }; // #endregion getTranslateRunWsUrl // #endregion wsUrlHelpers @@ -281,14 +266,18 @@ async function fetchApi(endpoint: string, options: FetchOptions = { // @PRE endpoint is a non-empty string path. // @POST Returns Promise with binary data. Throws ApiError on failure or 202 "in progress". // @SIDE_EFFECT Sends HTTP GET request. On failure (unless notifyError=false), dispatches error toast. +// Writes CoT log line on entry, success, and failure. // @RELATION DEPENDS_ON -> [buildApiError] // @RELATION DEPENDS_ON -> [notifyApiError] // @RELATION DEPENDS_ON -> [getAuthHeaders] // @RATIONALE 202 status is handled as a special case — the thumbnail generation may still be in progress. // The caller can retry after a delay. This is NOT treated as a server error. async function fetchApiBlob(endpoint: string, options: FetchOptions = {}): Promise { + const _start = performance.now(); const notifyError = options.notifyError !== false; + const _silent = _isSilentPolling(endpoint); try { + if (!_silent) log('ApiClient', 'REASON', 'GET blob', { endpoint }); const fetchInit: RequestInit = { headers: getAuthHeaders(options.headers || {}) }; if (options.signal) fetchInit.signal = options.signal; const response = await fetch(`${API_BASE_URL}${endpoint}`, fetchInit); @@ -299,9 +288,15 @@ async function fetchApiBlob(endpoint: string, options: FetchOptions = {}): Promi throw error; } if (!response.ok) throw await buildApiError(response); + _captureTraceId(response); + if (!_silent) log('ApiClient', 'REFLECT', 'GET blob completed', { + endpoint, status: response.status, + elapsed_ms: Math.round(performance.now() - _start), + }); return await response.blob(); } catch (error) { const apiError = error as ApiError; + log('ApiClient', 'EXPLORE', 'GET blob failed', { endpoint, status: apiError?.status }, apiError?.message || 'unknown'); if (notifyError) notifyApiError(apiError); throw error; } diff --git a/frontend/src/lib/api/__tests__/api.test.ts b/frontend/src/lib/api/__tests__/api.test.ts index 496868803..aee2494c8 100644 --- a/frontend/src/lib/api/__tests__/api.test.ts +++ b/frontend/src/lib/api/__tests__/api.test.ts @@ -7,22 +7,7 @@ // @TEST_CONTRACT: postApi -> Sends POST with JSON body, handles errors // @TEST_CONTRACT: deleteApi -> Sends DELETE, always notifies on error // @TEST_CONTRACT: requestApi -> Generic method with suppression heuristics -// @TEST_CONTRACT: wsUrl helpers -> Build correct ws:// URL with auth token -// @TEST_EDGE: network_failure -> fetch throws, error toast dispatched -// @TEST_EDGE: auth_token -> localStorage token injected as Bearer header -// @TEST_EDGE: suppress_toast -> Options.suppressToast suppresses error toast -// @TEST_EDGE: api_methods -> Registry methods call correct endpoints - -// #region ApiModuleTest [C:3] [TYPE Module] [SEMANTICS test, api, fetch, ws-url] -// @BRIEF Unit tests for the core API communication layer — fetch wrappers, WebSocket URL builders, -// endpoint registry, error normalization, and toast suppression. -// @LAYER Tests -// @RELATION BINDS_TO -> [ApiModule] -// @TEST_CONTRACT: fetchApi -> Returns typed JSON for 2xx, null for 204, throws ApiError on error -// @TEST_CONTRACT: postApi -> Sends POST with JSON body, handles errors -// @TEST_CONTRACT: deleteApi -> Sends DELETE, always notifies on error -// @TEST_CONTRACT: requestApi -> Generic method with suppression heuristics -// @TEST_CONTRACT: wsUrl helpers -> Build correct ws:// URL with auth token +// @TEST_CONTRACT: wsUrl helpers -> Build correct ws:// URL with auth token and trace ID // @TEST_EDGE: network_failure -> fetch throws, error toast dispatched // @TEST_EDGE: auth_token -> localStorage token injected as Bearer header // @TEST_EDGE: suppress_toast -> Options.suppressToast suppresses error toast @@ -65,34 +50,34 @@ describe('ApiModule — wsUrl helpers', () => { localStorage.setItem('auth_token', 'my-token-123'); const { getWsUrl } = await import('$lib/api.js'); const url = getWsUrl('task-42'); - expect(url).toBe('ws://localhost:5173/ws/logs/task-42?token=my-token-123'); + expect(url).toBe('ws://localhost:5173/ws/logs/task-42?token=my-token-123&x-trace-id=test-trace-id'); }); it('getWsUrl returns ws:// URL without token when auth_token missing', async () => { const { getWsUrl } = await import('$lib/api.js'); const url = getWsUrl('task-99'); - expect(url).toBe('ws://localhost:5173/ws/logs/task-99'); + expect(url).toBe('ws://localhost:5173/ws/logs/task-99?x-trace-id=test-trace-id'); }); it('getTaskEventsWsUrl builds correct URL with token', async () => { localStorage.setItem('auth_token', 'tok'); const { getTaskEventsWsUrl } = await import('$lib/api.js'); const url = getTaskEventsWsUrl(); - expect(url).toBe('ws://localhost:5173/ws/task-events?token=tok'); + expect(url).toBe('ws://localhost:5173/ws/task-events?token=tok&x-trace-id=test-trace-id'); }); it('getMaintenanceEventsWsUrl builds correct URL', async () => { localStorage.setItem('auth_token', 'maint-tok'); const { getMaintenanceEventsWsUrl } = await import('$lib/api.js'); const url = getMaintenanceEventsWsUrl(); - expect(url).toBe('ws://localhost:5173/ws/maintenance/events?token=maint-tok'); + expect(url).toBe('ws://localhost:5173/ws/maintenance/events?token=maint-tok&x-trace-id=test-trace-id'); }); it('getTranslateRunWsUrl builds correct URL with runId', async () => { localStorage.setItem('auth_token', 'run-tok'); const { getTranslateRunWsUrl } = await import('$lib/api.js'); const url = getTranslateRunWsUrl('run-1'); - expect(url).toBe('ws://localhost:5173/ws/translate/run/run-1?token=run-tok'); + expect(url).toBe('ws://localhost:5173/ws/translate/run/run-1?token=run-tok&x-trace-id=test-trace-id'); }); it('getWsUrl uses wss:// when window.location.protocol is https:', async () => { @@ -102,7 +87,7 @@ describe('ApiModule — wsUrl helpers', () => { localStorage.setItem('auth_token', 'sec-tok'); const { getWsUrl } = await import('$lib/api.js'); const url = getWsUrl('task-secure'); - expect(url).toBe('wss://app.example.com/ws/logs/task-secure?token=sec-tok'); + expect(url).toBe('wss://app.example.com/ws/logs/task-secure?token=sec-tok&x-trace-id=test-trace-id'); }); }); @@ -785,6 +770,146 @@ describe('ApiModule — fetch wrappers (global fetch mock)', () => { }); }); // #endregion captureTraceIdTests + + // #region fetchApiBlobTraceTests [C:2] [TYPE Test] [SEMANTICS test,api,blob,trace-id,cot] + // @BRIEF Verify fetchApiBlob captures x-trace-id and logs CoT markers on success/failure. + describe('fetchApiBlob — trace propagation & CoT logging', () => { + const testUuid = 'a1b2c3d4e5f67890abcdef1234567890'; + + beforeEach(() => { + vi.clearAllMocks(); + }); + + it('captures x-trace-id from response headers', async () => { + const blob = new Blob(['img'], { type: 'image/png' }); + vi.mocked(fetch).mockResolvedValue({ + ok: true, + status: 200, + blob: () => Promise.resolve(blob), + headers: { + get: vi.fn((name: string) => name === 'x-trace-id' ? testUuid : null), + }, + } as unknown as Response); + + const cotLogger = await import('$lib/cot-logger.js'); + cotLogger.setTraceId.mockClear(); + + const { api } = await import('$lib/api.js'); + await api.fetchApiBlob('/thumbnail'); + expect(cotLogger.setTraceId).toHaveBeenCalledWith(testUuid); + }); + + it('handles missing x-trace-id header gracefully', async () => { + const blob = new Blob(['img'], { type: 'image/png' }); + vi.mocked(fetch).mockResolvedValue({ + ok: true, + status: 200, + blob: () => Promise.resolve(blob), + headers: { get: vi.fn(() => null) }, + } as unknown as Response); + + const cotLogger = await import('$lib/cot-logger.js'); + cotLogger.setTraceId.mockClear(); + + const { api } = await import('$lib/api.js'); + await api.fetchApiBlob('/no-trace'); + expect(cotLogger.setTraceId).not.toHaveBeenCalled(); + }); + + it('logs REASON then REFLECT on success with elapsed_ms', async () => { + const blob = new Blob(['data'], { type: 'text/plain' }); + vi.mocked(fetch).mockResolvedValue({ + ok: true, + status: 200, + blob: () => Promise.resolve(blob), + headers: { get: vi.fn(() => null) }, + } as unknown as Response); + + const cotLogger = await import('$lib/cot-logger.js'); + cotLogger.log.mockClear(); + + const { api } = await import('$lib/api.js'); + await api.fetchApiBlob('/file'); + + const calls = cotLogger.log.mock.calls; + expect(calls.length).toBeGreaterThanOrEqual(2); + + const reasonCall = calls.find((c: unknown[]) => c[1] === 'REASON'); + expect(reasonCall).toBeTruthy(); + expect(reasonCall[0]).toBe('ApiClient'); + expect(reasonCall[2]).toBe('GET blob'); + expect(reasonCall[3]).toEqual({ endpoint: '/file' }); + + const reflectCall = calls.find((c: unknown[]) => c[1] === 'REFLECT'); + expect(reflectCall).toBeTruthy(); + expect(reflectCall[0]).toBe('ApiClient'); + expect(reflectCall[2]).toBe('GET blob completed'); + expect(reflectCall[3]).toMatchObject({ endpoint: '/file', status: 200 }); + expect(typeof reflectCall[3]?.elapsed_ms).toBe('number'); + }); + + it('logs EXPLORE on 500 with status in payload and error message', async () => { + vi.mocked(fetch).mockResolvedValue({ + ok: false, + status: 500, + json: () => Promise.resolve({ detail: 'Internal server error' }), + } as Response); + + const cotLogger = await import('$lib/cot-logger.js'); + cotLogger.log.mockClear(); + + const { api } = await import('$lib/api.js'); + await expect(api.fetchApiBlob('/fail')).rejects.toMatchObject({ status: 500 }); + + const exploreCall = cotLogger.log.mock.calls.find((c: unknown[]) => c[1] === 'EXPLORE'); + expect(exploreCall).toBeTruthy(); + expect(exploreCall[0]).toBe('ApiClient'); + expect(exploreCall[2]).toBe('GET blob failed'); + expect(exploreCall[3]).toMatchObject({ endpoint: '/fail', status: 500 }); + expect(exploreCall[4]).toMatch(/Internal server error/); + }); + + it('logs EXPLORE on 202 (resource being prepared)', async () => { + vi.mocked(fetch).mockResolvedValue({ + ok: true, + status: 202, + json: () => Promise.resolve({ message: 'Still generating' }), + } as Response); + + const cotLogger = await import('$lib/cot-logger.js'); + cotLogger.log.mockClear(); + + const { api } = await import('$lib/api.js'); + await expect(api.fetchApiBlob('/pending')).rejects.toMatchObject({ status: 202 }); + + const exploreCall = cotLogger.log.mock.calls.find((c: unknown[]) => c[1] === 'EXPLORE'); + expect(exploreCall).toBeTruthy(); + expect(exploreCall[2]).toBe('GET blob failed'); + // 202 is an error case for blob — status in payload + expect(exploreCall[3]).toMatchObject({ endpoint: '/pending', status: 202 }); + }); + + it('suppresses CoT REASON/REFLECT for silent polling endpoints, still logs EXPLORE on failure', async () => { + vi.mocked(fetch).mockResolvedValue({ + ok: false, + status: 500, + json: () => Promise.resolve({ detail: 'boom' }), + } as Response); + + const cotLogger = await import('$lib/cot-logger.js'); + cotLogger.log.mockClear(); + + const { api } = await import('$lib/api.js'); + await expect(api.fetchApiBlob('/health/summary')).rejects.toMatchObject({ status: 500 }); + + const calls = cotLogger.log.mock.calls; + const reasonCall = calls.find((c: unknown[]) => c[1] === 'REASON'); + expect(reasonCall).toBeUndefined(); + const exploreCall = calls.find((c: unknown[]) => c[1] === 'EXPLORE'); + expect(exploreCall).toBeTruthy(); + }); + }); + // #endregion fetchApiBlobTraceTests }); describe('ApiModule — deleteValidationTask query param', () => { diff --git a/frontend/src/lib/components/ui/MappingTable.svelte b/frontend/src/lib/components/ui/MappingTable.svelte index 18cd47776..107b0ac5c 100644 --- a/frontend/src/lib/components/ui/MappingTable.svelte +++ b/frontend/src/lib/components/ui/MappingTable.svelte @@ -2,6 +2,12 @@ + + + + + + -