From 32933a3d58b2576c1f5969e6d1091b335b81c6dc Mon Sep 17 00:00:00 2001 From: cj2026-bit <647646783@qq.com> Date: Sun, 20 Sep 2026 09:17:30 +0800 Subject: [PATCH 1/9] =?UTF-8?q?=F0=9F=90=9B=20Fix(evaluation):=20run=20tri?= =?UTF-8?q?als=20in=20runtime=20service=20(#3954)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(evaluation): run trials in runtime service Route trial evaluations through the authenticated Config-to-Runtime proxy and use Config's manager only for creation-stage preparation. Keep Agent execution and evaluator scoring in Runtime. Co-authored-by: Codex Generated-by: gpt-5 * test(evaluation): stub config thread manager Keep pure-logic service import tests aligned with the Config and Runtime thread-manager split. Co-authored-by: Codex Generated-by: gpt-5 * test(evaluation): stub runtime jwt helper * test(evaluation): cover trial proxy error paths --- backend/apps/agent_evaluation_app.py | 8 +- backend/apps/agent_evaluation_runtime_app.py | 48 ++++++- backend/services/agent_evaluation_service.py | 11 +- backend/services/runtime_proxy_service.py | 53 ++++++++ test/backend/app/test_agent_evaluation_app.py | 10 +- .../app/test_agent_evaluation_runtime_app.py | 54 +++++++- .../backend/app/test_evaluation_delete_app.py | 2 +- .../services/test_agent_evaluation_service.py | 8 +- .../services/test_evaluation_pure_logic.py | 1 + .../services/test_runtime_proxy_service.py | 128 ++++++++++++++++++ 10 files changed, 305 insertions(+), 18 deletions(-) diff --git a/backend/apps/agent_evaluation_app.py b/backend/apps/agent_evaluation_app.py index 45e54499b..23e872b0e 100644 --- a/backend/apps/agent_evaluation_app.py +++ b/backend/apps/agent_evaluation_app.py @@ -19,9 +19,9 @@ get_evaluation_stats_impl, list_agent_evaluation_cases_impl, list_agent_evaluations_by_agent_impl, - trial_run_evaluator_impl, ) from services.evaluation_report_service import generate_agent_evaluation_report_impl +from services.runtime_proxy_service import forward_agent_evaluation_trial_run from utils.auth_utils import get_current_user_id, get_current_user_info @@ -552,15 +552,15 @@ async def trial_run_api( """ try: user_id, tenant_id = get_current_user_id(authorization) - result = await trial_run_evaluator_impl( - tenant_id=tenant_id, - user_id=user_id, + result = await forward_agent_evaluation_trial_run( agent_id=payload.agent_id, agent_version_no=payload.agent_version_no, query=payload.query, judge_model_id=payload.judge_model_id, evaluator_ids=payload.evaluator_ids, language=payload.language, + user_id=user_id, + tenant_id=tenant_id, ) logger.info( "trial_run_api OK: tenant=%s user=%s agent_id=%s version=%s " diff --git a/backend/apps/agent_evaluation_runtime_app.py b/backend/apps/agent_evaluation_runtime_app.py index 3c098a92f..77aada355 100644 --- a/backend/apps/agent_evaluation_runtime_app.py +++ b/backend/apps/agent_evaluation_runtime_app.py @@ -5,6 +5,7 @@ from typing import Annotated from fastapi import APIRouter, Header, HTTPException +from nexent.core.concurrency import ManagedTaskSpec from pydantic import BaseModel, Field from consts.evaluation_status import EvalRunStatus @@ -13,7 +14,6 @@ claim_agent_evaluation_run, get_agent_evaluation, ) -from nexent.core.concurrency import ManagedTaskSpec from services.thread_lifecycle_service import runtime_thread_manager from utils.auth_utils import verify_internal_runtime_jwt @@ -28,6 +28,17 @@ class EvaluationRunRequest(BaseModel): agent_evaluation_id: int = Field(gt=0) +class TrialRunRequest(BaseModel): + """Payload used by Config service for a non-persistent trial evaluation.""" + + agent_id: int + agent_version_no: int = 1 + query: str + judge_model_id: int + evaluator_ids: list[int] | None = None + language: str = "zh" + + def _load_evaluation_executor(): """Load the evaluation service only when a runtime run is dispatched.""" from services.agent_evaluation_service import execute_agent_evaluation_run @@ -35,6 +46,41 @@ def _load_evaluation_executor(): return execute_agent_evaluation_run +def _load_trial_executor(): + """Load the trial executor only when Runtime receives a trial request.""" + from services.agent_evaluation_service import trial_run_evaluator_impl + + return trial_run_evaluator_impl + + +@router.post("/trial-run", include_in_schema=False) +async def trial_run_evaluation_api( + payload: TrialRunRequest, + authorization: Annotated[str | None, Header()] = None, +): + """Run one ad-hoc evaluation in the Runtime process.""" + try: + user_id, tenant_id = verify_internal_runtime_jwt(authorization) + except Exception as exc: + logger.warning("Rejected unauthenticated trial evaluation: %s", exc) + raise HTTPException( + status_code=HTTPStatus.UNAUTHORIZED, + detail="Invalid internal runtime authorization", + ) from exc + + trial_run_evaluator_impl = _load_trial_executor() + return await trial_run_evaluator_impl( + tenant_id=tenant_id, + user_id=user_id, + agent_id=payload.agent_id, + agent_version_no=payload.agent_version_no, + query=payload.query, + judge_model_id=payload.judge_model_id, + evaluator_ids=payload.evaluator_ids, + language=payload.language, + ) + + @router.post("/run", include_in_schema=False, status_code=HTTPStatus.ACCEPTED) async def dispatch_evaluation_run_api( payload: EvaluationRunRequest, diff --git a/backend/services/agent_evaluation_service.py b/backend/services/agent_evaluation_service.py index d32a1fb78..d1b55b72f 100644 --- a/backend/services/agent_evaluation_service.py +++ b/backend/services/agent_evaluation_service.py @@ -62,7 +62,10 @@ from database.evaluator_db import get_evaluator from management.services.agent.service import prepare_agent_run from services.evaluation_set_service import resolve_latest_published_version_no -from services.thread_lifecycle_service import runtime_thread_manager +from services.thread_lifecycle_service import ( + config_thread_manager, + runtime_thread_manager, +) from utils.llm_utils import call_llm_for_system_prompt from utils.prompt_template_utils import get_prompt_template @@ -1054,12 +1057,12 @@ def _check_run_limits(tenant_id: str) -> None: def _run_in_background( fn, *fn_args, tenant_id, user_id, agent_evaluation_id, language="zh" ): - """Submit fn to the Runtime evaluation lane and attach failure cleanup.""" - execution = runtime_thread_manager.submit( + """Submit Config-owned preparation or dispatch work in the background.""" + execution = config_thread_manager.submit( "evaluation", ManagedTaskSpec( task_name="agent-evaluation-run", - owner="runtime", + owner="config", run_id=str(agent_evaluation_id), ), fn, diff --git a/backend/services/runtime_proxy_service.py b/backend/services/runtime_proxy_service.py index 4a5b1d90d..cde7e8dae 100644 --- a/backend/services/runtime_proxy_service.py +++ b/backend/services/runtime_proxy_service.py @@ -29,6 +29,7 @@ _STREAM_TIMEOUT = httpx.Timeout(connect=10.0, read=None, write=30.0, pool=10.0) _REQUEST_TIMEOUT = httpx.Timeout(connect=10.0, read=30.0, write=30.0, pool=10.0) _EVALUATION_DISPATCH_TIMEOUT = httpx.Timeout(connect=10.0, read=30.0, write=30.0, pool=10.0) +_EVALUATION_TRIAL_TIMEOUT = httpx.Timeout(connect=10.0, read=None, write=30.0, pool=10.0) _RUNTIME_SERVICE_UNAVAILABLE_MESSAGE = "Runtime service is unavailable" @@ -101,6 +102,58 @@ def dispatch_agent_evaluation_run( return payload +async def forward_agent_evaluation_trial_run( + *, + agent_id: int, + agent_version_no: int, + query: str, + judge_model_id: int, + evaluator_ids: list[int] | None, + language: str, + user_id: str, + tenant_id: str, +) -> dict: + """Execute an ad-hoc evaluation in the Runtime process and return its result.""" + try: + async with create_httpx_client( + headers=_authorization_headers(user_id, tenant_id), + timeout=_EVALUATION_TRIAL_TIMEOUT, + ) as client: + response = await client.post( + _runtime_url("/agent-evaluations/internal/trial-run"), + json={ + "agent_id": agent_id, + "agent_version_no": agent_version_no, + "query": query, + "judge_model_id": judge_model_id, + "evaluator_ids": evaluator_ids, + "language": language, + }, + ) + except httpx.TimeoutException as exc: + raise RuntimeServiceTimeoutError("Runtime trial evaluation timed out") from exc + except httpx.RequestError as exc: + raise RuntimeServiceUnavailableError(_RUNTIME_SERVICE_UNAVAILABLE_MESSAGE) from exc + + if response.status_code >= 400: + raise RuntimeUpstreamError( + status_code=response.status_code, + content=response.content, + headers=_forwarded_headers(response.headers), + ) + try: + payload = response.json() + except ValueError as exc: + raise RuntimeServiceUnavailableError( + "Runtime trial evaluation response is not valid JSON" + ) from exc + if not isinstance(payload, dict): + raise RuntimeServiceUnavailableError( + "Runtime trial evaluation response is not a JSON object" + ) + return payload + + async def forward_agent_run( agent_request: AgentRequest, user_id: str, diff --git a/test/backend/app/test_agent_evaluation_app.py b/test/backend/app/test_agent_evaluation_app.py index b144df716..c01cfbb11 100644 --- a/test/backend/app/test_agent_evaluation_app.py +++ b/test/backend/app/test_agent_evaluation_app.py @@ -92,6 +92,7 @@ def _register_package(name: str) -> types.ModuleType: "services", "services.agent_evaluation_service", "services.evaluation_report_service", + "services.runtime_proxy_service", "database", "database.agent_evaluation_db", "utils", @@ -100,6 +101,7 @@ def _register_package(name: str) -> types.ModuleType: _register_package(_name) sys.modules["services.agent_evaluation_service"] = MagicMock(name="agent_eval_svc") sys.modules["services.evaluation_report_service"] = MagicMock(name="report_svc") +sys.modules["services.runtime_proxy_service"] = MagicMock(name="runtime_proxy_svc") sys.modules["database.agent_evaluation_db"] = MagicMock(name="agent_eval_db") sys.modules["utils.auth_utils"] = MagicMock(name="auth_utils") @@ -179,7 +181,7 @@ def _mock_impls(**overrides): return_value={"items": [], "total": 0} ), "list_agent_evaluations_by_agent_impl": MagicMock(return_value=[{"id": 1}]), - "trial_run_evaluator_impl": AsyncMock(return_value={"result": "ok"}), + "forward_agent_evaluation_trial_run": AsyncMock(return_value={"result": "ok"}), "generate_agent_evaluation_report_impl": MagicMock( return_value=(b"%PDF-1.4 fake", 0) ), @@ -593,7 +595,7 @@ def test_runs_trial(self, client): ) assert response.status_code == 200 assert response.json()["data"] == {"result": "ok"} - assert app.trial_run_evaluator_impl.call_args.kwargs["query"] == "hello" + assert app.forward_agent_evaluation_trial_run.call_args.kwargs["query"] == "hello" def test_401_on_unauthorized(self, client): from consts.exceptions import UnauthorizedError @@ -607,7 +609,7 @@ def test_401_on_unauthorized(self, client): def test_500_on_exception(self, client): _mock_impls( - trial_run_evaluator_impl=AsyncMock(side_effect=RuntimeError("boom")) + forward_agent_evaluation_trial_run=AsyncMock(side_effect=RuntimeError("boom")) ) response = client.post( "/agent-evaluations/trial-run", @@ -617,7 +619,7 @@ def test_500_on_exception(self, client): def test_app_exception_propagates(self, client): _mock_impls( - trial_run_evaluator_impl=AsyncMock( + forward_agent_evaluation_trial_run=AsyncMock( side_effect=_exc(_code("COMMON_RESOURCE_NOT_FOUND"), "missing") ) ) diff --git a/test/backend/app/test_agent_evaluation_runtime_app.py b/test/backend/app/test_agent_evaluation_runtime_app.py index e13c7f94b..3842438a5 100644 --- a/test/backend/app/test_agent_evaluation_runtime_app.py +++ b/test/backend/app/test_agent_evaluation_runtime_app.py @@ -1,6 +1,6 @@ """Tests for runtime-owned evaluation dispatch.""" -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import HTTPException @@ -8,12 +8,64 @@ import apps.agent_evaluation_runtime_app as runtime_app from apps.agent_evaluation_runtime_app import ( EvaluationRunRequest, + TrialRunRequest, dispatch_evaluation_run_api, + trial_run_evaluation_api, ) from consts.evaluation_status import EvalRunStatus from consts.exceptions import AppException +@pytest.mark.asyncio +async def test_trial_run_uses_internal_identity_and_runtime_executor(monkeypatch): + monkeypatch.setattr(runtime_app, "verify_internal_runtime_jwt", lambda _: ("u1", "t1")) + executor = AsyncMock(return_value={"answer": "ok", "scores": {"judge": 1.0}}) + monkeypatch.setattr(runtime_app, "_load_trial_executor", lambda: executor) + + result = await trial_run_evaluation_api( + TrialRunRequest( + agent_id=7, + agent_version_no=3, + query="hello", + judge_model_id=99, + evaluator_ids=[5], + ), + "internal-token", + ) + + assert result == {"answer": "ok", "scores": {"judge": 1.0}} + executor.assert_awaited_once_with( + tenant_id="t1", + user_id="u1", + agent_id=7, + agent_version_no=3, + query="hello", + judge_model_id=99, + evaluator_ids=[5], + language="zh", + ) + + +@pytest.mark.asyncio +async def test_trial_run_rejects_missing_internal_token_without_executing(monkeypatch): + monkeypatch.setattr( + runtime_app, + "verify_internal_runtime_jwt", + MagicMock(side_effect=ValueError("invalid token")), + ) + load_executor = MagicMock() + monkeypatch.setattr(runtime_app, "_load_trial_executor", load_executor) + + with pytest.raises(HTTPException) as exc_info: + await trial_run_evaluation_api( + TrialRunRequest(agent_id=7, query="hello", judge_model_id=99), + None, + ) + + assert exc_info.value.status_code == 401 + load_executor.assert_not_called() + + @pytest.mark.asyncio async def test_dispatch_claims_pending_run_and_submits_runtime_worker(monkeypatch): monkeypatch.setattr(runtime_app, "verify_internal_runtime_jwt", lambda _: ("u1", "t1")) diff --git a/test/backend/app/test_evaluation_delete_app.py b/test/backend/app/test_evaluation_delete_app.py index a554ecf77..1a53350e0 100644 --- a/test/backend/app/test_evaluation_delete_app.py +++ b/test/backend/app/test_evaluation_delete_app.py @@ -41,7 +41,7 @@ "update_evaluation_set_case_impl", ), "database.agent_evaluation_db": ("update_annotation_schema_ids",), - "utils.auth_utils": ("get_current_user_id", "get_current_user_info"), + "utils.auth_utils": ("get_current_user_id", "get_current_user_info", "generate_internal_runtime_jwt"), "utils.evaluation_set_excel_utils": ( "build_evaluation_set_excel_template_bytes", "parse_evaluation_cases_from_excel", ), diff --git a/test/backend/services/test_agent_evaluation_service.py b/test/backend/services/test_agent_evaluation_service.py index 2f94f6ae1..795fd6a27 100644 --- a/test/backend/services/test_agent_evaluation_service.py +++ b/test/backend/services/test_agent_evaluation_service.py @@ -448,6 +448,7 @@ def __init__(self, error_code=None, message=None, *args, **kwargs): "services.thread_lifecycle_service" ) _thread_lifecycle_service_module.runtime_thread_manager = MagicMock() +_thread_lifecycle_service_module.config_thread_manager = MagicMock() sys.modules["services.thread_lifecycle_service"] = _thread_lifecycle_service_module _services_pkg.thread_lifecycle_service = _thread_lifecycle_service_module @@ -1150,7 +1151,7 @@ def test_create_agent_evaluation_run_happy_path(service_module): pool_mock = MagicMock() future = MagicMock() pool_mock.submit.return_value = types.SimpleNamespace(future=future) - service_module.runtime_thread_manager = pool_mock + service_module.config_thread_manager = pool_mock run = service_module.create_agent_evaluation_run_impl( tenant_id="t1", @@ -1179,6 +1180,7 @@ def test_create_agent_evaluation_run_happy_path(service_module): assert len(kwargs["set_cases"]) == 3 pool_mock.submit.assert_called_once() + service_module.runtime_thread_manager.submit.assert_not_called() future.add_done_callback.assert_called_once() # Done-callback signature should be a callable wrapping the run id + tenant. callback = future.add_done_callback.call_args.args[0] @@ -1207,8 +1209,8 @@ def test_create_agent_evaluation_run_uses_resolved_version_no(service_module): """The published version number flows from ``resolve_latest_published_version_no``.""" create_mock = _wire_full_db_module(service_module) service_module.resolve_latest_published_version_no.return_value = 13 - service_module.runtime_thread_manager = MagicMock() - service_module.runtime_thread_manager.submit.return_value = types.SimpleNamespace( + service_module.config_thread_manager = MagicMock() + service_module.config_thread_manager.submit.return_value = types.SimpleNamespace( future=MagicMock() ) diff --git a/test/backend/services/test_evaluation_pure_logic.py b/test/backend/services/test_evaluation_pure_logic.py index 58a3d12f4..a75f7d2e6 100644 --- a/test/backend/services/test_evaluation_pure_logic.py +++ b/test/backend/services/test_evaluation_pure_logic.py @@ -371,6 +371,7 @@ def __init__(self, code: Any, msg: str = ""): _services_pkg.evaluation_set_service = _ess_mod _tls_mod = _mk_mod( "services.thread_lifecycle_service", + config_thread_manager=MagicMock(), runtime_thread_manager=MagicMock(), ) _services_pkg.thread_lifecycle_service = _tls_mod diff --git a/test/backend/services/test_runtime_proxy_service.py b/test/backend/services/test_runtime_proxy_service.py index 0c8b20f5b..ce43de09e 100644 --- a/test/backend/services/test_runtime_proxy_service.py +++ b/test/backend/services/test_runtime_proxy_service.py @@ -142,6 +142,134 @@ def test_authorization_headers_maps_missing_jwt_configuration(monkeypatch): proxy._authorization_headers("user-a", "tenant-a") +@pytest.mark.asyncio +async def test_forward_agent_evaluation_trial_run_posts_runtime_request(monkeypatch): + captured = {} + + async def handler(request: httpx.Request): + captured["request"] = request + return httpx.Response(200, json={"answer": "ok", "scores": {"judge": 1.0}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + monkeypatch.setattr(proxy, "RUNTIME_SERVICE_URL", "http://runtime:5014") + monkeypatch.setattr(proxy, "generate_internal_runtime_jwt", lambda *_: "jwt") + + def create_client(**kwargs): + client.headers.update(kwargs["headers"]) + return client + + monkeypatch.setattr(proxy, "create_httpx_client", create_client) + + result = await proxy.forward_agent_evaluation_trial_run( + agent_id=7, + agent_version_no=3, + query="hello", + judge_model_id=99, + evaluator_ids=[5], + language="zh", + user_id="user-a", + tenant_id="tenant-a", + ) + + assert result == {"answer": "ok", "scores": {"judge": 1.0}} + request = captured["request"] + assert str(request.url) == "http://runtime:5014/api/agent-evaluations/internal/trial-run" + assert request.headers["authorization"] == "Bearer jwt" + assert json.loads(request.content) == { + "agent_id": 7, + "agent_version_no": 3, + "query": "hello", + "judge_model_id": 99, + "evaluator_ids": [5], + "language": "zh", + } + assert client.is_closed is True + + +@pytest.mark.asyncio +async def test_forward_agent_evaluation_trial_run_maps_upstream_error(monkeypatch): + client = httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _: httpx.Response(500, content=b'{"detail":"failed"}') + ) + ) + monkeypatch.setattr(proxy, "generate_internal_runtime_jwt", lambda *_: "jwt") + monkeypatch.setattr(proxy, "create_httpx_client", lambda **_: client) + + with pytest.raises(RuntimeUpstreamError) as exc_info: + await proxy.forward_agent_evaluation_trial_run( + agent_id=7, + agent_version_no=3, + query="hello", + judge_model_id=99, + evaluator_ids=None, + language="zh", + user_id="user-a", + tenant_id="tenant-a", + ) + + assert exc_info.value.status_code == 500 + assert client.is_closed is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("transport_error", "expected_error"), + [ + (httpx.ReadTimeout("timed out"), RuntimeServiceTimeoutError), + (httpx.ConnectError("connection failed"), RuntimeServiceUnavailableError), + ], +) +async def test_forward_agent_evaluation_trial_run_maps_transport_errors( + monkeypatch, transport_error, expected_error, +): + async def handler(request: httpx.Request): + transport_error.request = request + raise transport_error + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + monkeypatch.setattr(proxy, "generate_internal_runtime_jwt", lambda *_: "jwt") + monkeypatch.setattr(proxy, "create_httpx_client", lambda **_: client) + + with pytest.raises(expected_error): + await proxy.forward_agent_evaluation_trial_run( + agent_id=7, + agent_version_no=3, + query="hello", + judge_model_id=99, + evaluator_ids=None, + language="zh", + user_id="user-a", + tenant_id="tenant-a", + ) + + assert client.is_closed is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("content", [b"not-json", b"[]"]) +async def test_forward_agent_evaluation_trial_run_rejects_invalid_success_payload(monkeypatch, content): + client = httpx.AsyncClient( + transport=httpx.MockTransport(lambda _: httpx.Response(200, content=content)) + ) + monkeypatch.setattr(proxy, "generate_internal_runtime_jwt", lambda *_: "jwt") + monkeypatch.setattr(proxy, "create_httpx_client", lambda **_: client) + + with pytest.raises(RuntimeServiceUnavailableError): + await proxy.forward_agent_evaluation_trial_run( + agent_id=7, + agent_version_no=3, + query="hello", + judge_model_id=99, + evaluator_ids=None, + language="zh", + user_id="user-a", + tenant_id="tenant-a", + ) + + assert client.is_closed is True + + @pytest.mark.asyncio async def test_forward_agent_run_streams_body_and_closes_resources(monkeypatch): stream = TrackingStream([b"data: one\n\n", b"data: two\n\n"]) From 47332025c6667d9f8e6789297c45c0f8efb3fd6e Mon Sep 17 00:00:00 2001 From: lijiayang619 <1170349871@qq.com> Date: Sun, 20 Sep 2026 09:50:13 +0800 Subject: [PATCH 2/9] Fix/override delete (#3958) * Fix: override dialog only shows override values, not model defaults (deleted params no longer reappear) * Fix: custom param deletion persists (null markers), per-agent capacity overrides take effect, and edit-dialog connectivity probe uses stored api_key * Fix: rename ModelRequest.model_id to probe_model_id - model_dump() is spread into INSERT column lists, so a model_id field injected an explicit NULL primary key and broke model creation * Fix: move probe_model_id to a dedicated ModelProbeRequest subclass - ModelRequest.model_dump() is spread into INSERT column lists, so any non-column field breaks model creation (Unconsumed column names) * Fix: editing/adding a model no longer steals the occupied default-model slot - persistCustomLocalConfig now only writes the slot when it is empty (onboarding) or the submitted model already occupies it * Fix: remove persistCustomLocalConfig - the frontend-cached-config guard could still steal an occupied default slot when the cache was stale/empty. Default-slot writes now come only from the server (create-time backfill for empty/dangling slots) * Revert "Fix: remove persistCustomLocalConfig - the frontend-cached-config guard could still steal an occupied default slot when the cache was stale/empty. Default-slot writes now come only from the server (create-time backfill for empty/dangling slots)" This reverts commit 2376aa50ccb0f170e5412a14c5aee33798ca1f06. * Fix: VLM connectivity probe never found the local test image - the gateway adapter's relative dirname chain resolved two levels short of the package root, so every probe silently fell back to a public URL that is unreachable in offline deployments. Anchor both probe copies on nexent.__file__ so the path survives module moves. --------- Co-authored-by: ljy --- backend/agents/create_agent_info.py | 74 ++++++++++++++++++- backend/apps/model_managment_app.py | 14 +++- backend/consts/model.py | 17 +++++ .../agents/components/agent-prompt.tsx | 38 ++++++---- .../components/model/ModelAddDialogV2.tsx | 18 +++++ .../model/ModelAdvancedSettings.tsx | 72 ++++++++++++++++++ frontend/services/modelService.ts | 6 ++ .../core/gateway/modality/vlm/openai.py | 9 ++- sdk/nexent/core/models/openai_vlm.py | 9 ++- 9 files changed, 234 insertions(+), 23 deletions(-) diff --git a/backend/agents/create_agent_info.py b/backend/agents/create_agent_info.py index 4a50f3c34..7c404a70d 100644 --- a/backend/agents/create_agent_info.py +++ b/backend/agents/create_agent_info.py @@ -274,6 +274,45 @@ def _operator_overrides_from_model_info(model_info: Optional[dict]) -> dict: return overrides +def _agent_capacity_overrides( + agent_info: Optional[dict], + model_id: Optional[int], +) -> Dict[str, Any]: + """Extract per-agent capacity overrides for the selected model. + + v2.6.0 model_params_override entries may carry capacity fields + (context_window_tokens / max_input_tokens / max_output_tokens / + default_output_reserve_tokens / tokenizer_family) next to inference + params. When present they win over the model-level capacity columns in + W1/W2 resolution, mirroring how temperature/top_p overrides win. + """ + if not isinstance(agent_info, dict) or model_id is None: + return {} + override_map = agent_info.get("model_params_override") + if not isinstance(override_map, dict): + return {} + entry = override_map.get(str(model_id)) + if not isinstance(entry, dict): + return {} + overrides: Dict[str, Any] = {} + for field in _OPERATOR_OVERRIDE_FIELDS: + value = entry.get(field) + if value is not None: + overrides[field] = value + # Per-agent override semantics: a filled value simply replaces the + # model-level value for THIS agent. "最大输出Token数" is the field users + # expect to control the actual per-request max_tokens, so mirror it into + # default_output_reserve_tokens unless the user set the reserve + # explicitly. Without this, the request would keep the model-level + # reserve (4096) and the filled cap alone would change nothing visible. + if ( + "max_output_tokens" in overrides + and "default_output_reserve_tokens" not in overrides + ): + overrides["default_output_reserve_tokens"] = overrides["max_output_tokens"] + return overrides + + def _dominant_capacity_source(field_sources: dict) -> Optional[str]: values = [value for value in field_sources.values() if value] if not values: @@ -363,6 +402,7 @@ def _resolve_context_budget( def _resolve_input_budget( model_info: Optional[dict], + capacity_overrides: Optional[Dict[str, Any]] = None, ) -> tuple[int, Optional[dict], Optional[ModelCapacitySnapshot]]: """Resolve the context-manager input budget for a model_record_t row. @@ -371,6 +411,9 @@ def _resolve_input_budget( Falls back to _TOKEN_THRESHOLD_LEGACY_FALLBACK with no snapshot when capacity is unknown - this is the migration-window behavior before all model rows are backfilled. + + capacity_overrides carries per-agent capacity fields (from + model_params_override) that win over the model-level columns. """ if not isinstance(model_info, dict): return _TOKEN_THRESHOLD_LEGACY_FALLBACK, None, None @@ -383,10 +426,13 @@ def _resolve_input_budget( "model_factory/provider is missing; capacity catalog matching is disabled" ) try: + operator_overrides = _operator_overrides_from_model_info(model_info) + if capacity_overrides: + operator_overrides.update(capacity_overrides) snapshot = resolve_capacity( model_id=model_id, provider=provider, - operator_overrides=_operator_overrides_from_model_info(model_info), + operator_overrides=operator_overrides, capability_profiles=CAPABILITY_CATALOG, ) logger.debug( @@ -1423,9 +1469,14 @@ async def create_agent_config( # W1 step 6: derive input budget via ModelCapacityResolver instead of # treating model_info["max_tokens"] (a deprecated output cap) as a # context threshold. Falls back to a safe constant when capacity is - # unknown during the migration window. + # unknown during the migration window. Per-agent capacity overrides + # (model_params_override) win over the model-level columns. input_budget, capacity_snapshot, resolved_capacity_snapshot = ( - _resolve_input_budget(model_info) + _resolve_input_budget( + model_info, + capacity_overrides=_agent_capacity_overrides( + agent_info, model_id_to_use), + ) ) else: model_name = "main_model" @@ -2297,12 +2348,27 @@ async def create_agent_run_info( mc.temperature = override_entry["temperature"] if override_entry.get("top_p") is not None: mc.top_p = override_entry["top_p"] + # v2.6.0: capacity fields are overridable per-agent. The + # authoritative consumer for request shaping is the W1/W2 + # resolution (agent-selected model), this keeps each + # ModelConfig consistent with its override entry. + for field in _OPERATOR_OVERRIDE_FIELDS: + if override_entry.get(field) is not None: + setattr(mc, field, override_entry[field]) override_extra = override_entry.get("extra_params") if override_extra and isinstance(override_extra, dict): merged = dict(mc.extra_body or {}) for k, v in override_extra.items(): if k == "__custom__" and isinstance(v, dict): - merged.update(v) + for custom_key, custom_value in v.items(): + # A null custom value is an explicit + # removal marker: the agent opts out of a + # model-level custom param instead of + # inheriting it. + if custom_value is None: + merged.pop(custom_key, None) + else: + merged[custom_key] = custom_value else: merged[k] = v mc.extra_body = merged if merged else None diff --git a/backend/apps/model_managment_app.py b/backend/apps/model_managment_app.py index fe31b87bf..2c02edcef 100644 --- a/backend/apps/model_managment_app.py +++ b/backend/apps/model_managment_app.py @@ -20,6 +20,7 @@ BatchCreateModelsRequest, CapacitySuggestionFields, ModelRequest, + ModelProbeRequest, ModelCapacitySuggestionRequest, ModelCapacitySuggestionResponse, ProviderModelRequest, @@ -64,6 +65,7 @@ from utils.auth_utils import get_current_user_id from consts.exceptions import TokenExpiredError from nexent.core.concurrency import run_blocking +from database.model_management_db import get_model_by_model_id # Model Catalog loader (with graceful fallback) try: @@ -580,7 +582,7 @@ async def check_model_health( @router.post("/temporary_healthcheck") async def check_temporary_model_health( - request: ModelRequest, authorization: Optional[str] = Header(None) + request: ModelProbeRequest, authorization: Optional[str] = Header(None) ): """Verify connectivity for the provided model configuration without persisting it. @@ -589,7 +591,15 @@ async def check_temporary_model_health( authorization: Bearer token header used to enforce authentication. """ try: - get_current_user_id(authorization) + _, tenant_id = get_current_user_id(authorization) + # Edit-dialog probes arrive without the api_key (the backend never + # returns the persisted key to the client, and the dialog leaves the + # field empty to "keep existing"). Fall back to the stored key so + # verifying does not require retyping it. + if request.probe_model_id is not None and request.api_key in (None, "", "sk-no-api-key"): + stored_model = get_model_by_model_id(request.probe_model_id, tenant_id=tenant_id) + if stored_model and stored_model.get("api_key"): + request.api_key = stored_model["api_key"] result = await verify_model_config_connectivity(request.model_dump()) if result.get("connectivity") is True: # suggest_capacity may now issue an LLM self-report HTTP call diff --git a/backend/consts/model.py b/backend/consts/model.py index 9a01ef055..8f2a9a787 100644 --- a/backend/consts/model.py +++ b/backend/consts/model.py @@ -580,6 +580,23 @@ class ModelRequest(BaseModel): accepted_capability_profile_version: Optional[str] = None +class ModelProbeRequest(ModelRequest): + """Request payload for POST /model/temporary_healthcheck only. + + Extends ModelRequest with probe-only fields. Never send this to the + create/update endpoints: they spread model_dump() straight into + INSERT/UPDATE column lists, so any field that is not a real + model_record_t column makes SQLAlchemy raise "Unconsumed column names" + (and a field matching a column name would inject wrong data — the + model_id NULL primary-key bug). + """ + # Edit-dialog connectivity probes omit the stored api_key (the backend + # never returns it to the client). When set, /temporary_healthcheck falls + # back to the persisted key for this model instead of probing with the + # "sk-no-api-key" placeholder. + probe_model_id: Optional[int] = None + + class CapacitySuggestionFields(BaseModel): context_window_tokens: Optional[int] = None max_input_tokens: Optional[int] = None diff --git a/frontend/app/[locale]/agents/components/agent-prompt.tsx b/frontend/app/[locale]/agents/components/agent-prompt.tsx index f99197827..4bc69effd 100644 --- a/frontend/app/[locale]/agents/components/agent-prompt.tsx +++ b/frontend/app/[locale]/agents/components/agent-prompt.tsx @@ -28,6 +28,8 @@ import { ModelAdvancedSettingsValue, advancedSettingsValueFromRecord, buildModelOverrideEntry, + diffCustomParamsForSave, + mergeCustomParamsForEditing, } from "../../models/components/model/ModelAdvancedSettings"; import type { ModelOverrideMap } from "../../models/components/model/ModelOverrideModal"; import { canManageModels } from "@/lib/auth"; @@ -208,22 +210,20 @@ export default function AgentPrompt() { setEditingOverrideValue(null); return; } - const modelDefaults: Record = { - temperature: (configuringModel as any).temperature, - top_p: (configuringModel as any).topP, - extra_params: (configuringModel as any).extraParams, - }; + // Only override values are editable in this dialog; model-level defaults + // show up as placeholders via inheritedDefaults. Merging model defaults + // into the editing state made "deleted" overrides reappear (the model + // value was being re-merged as if the user had set it in the override). const overrideEntry = modelParamsOverride[String(configuringModel.id)] ?? {}; - const formRecord: Record = { ...modelDefaults, ...overrideEntry }; - if (overrideEntry.extra_params && modelDefaults.extra_params) { - formRecord.extra_params = { - ...(modelDefaults.extra_params as Record), - ...(overrideEntry.extra_params as Record), - }; - } - setEditingOverrideValue( - advancedSettingsValueFromRecord(formRecord as any, inferenceSpecs, (configuringModel as any).type ?? "llm") + const base = advancedSettingsValueFromRecord(overrideEntry as any, inferenceSpecs, (configuringModel as any).type ?? "llm"); + // Custom params are the exception: model-level customs render as plain + // editable rows (merged display). Deletion persists as a null marker in + // the override entry, so deleted rows do NOT reappear on reopen. + base.__custom__ = mergeCustomParamsForEditing( + (overrideEntry as any).extra_params?.__custom__, + (configuringModel as any).extraParams?.__custom__ ); + setEditingOverrideValue(base); // eslint-disable-next-line react-hooks/exhaustive-deps }, [configuringModelId]); @@ -544,11 +544,21 @@ export default function AgentPrompt() { ); const diffValue: ModelAdvancedSettingsValue = {}; for (const [key, val] of Object.entries(editingOverrideValue)) { + // Custom params are diffed per-key below (deleting an + // inherited row must persist a removal, not drop the key). + if (key === "__custom__") continue; const modelVal = modelDefaults[key]; if (JSON.stringify(modelVal) !== JSON.stringify(val)) { diffValue[key] = val; } } + const customDiff = diffCustomParamsForSave( + editingOverrideValue.__custom__, + (configuringModel as any).extraParams?.__custom__ + ); + if (Object.keys(customDiff).length > 0) { + diffValue.__custom__ = customDiff; + } handleModelParamsOverrideChange(configuringModel.id, diffValue); } setConfiguringModelId(null); diff --git a/frontend/app/[locale]/models/components/model/ModelAddDialogV2.tsx b/frontend/app/[locale]/models/components/model/ModelAddDialogV2.tsx index e8bb94015..c534cf925 100644 --- a/frontend/app/[locale]/models/components/model/ModelAddDialogV2.tsx +++ b/frontend/app/[locale]/models/components/model/ModelAddDialogV2.tsx @@ -906,6 +906,10 @@ export const ModelAddDialogV2 = ({ modelType: resolvedModelType, baseUrl: customForm.url, apiKey: customForm.apiKey, + // Edit mode leaves apiKey empty to "keep existing"; pass the model + // id so the backend probes with the stored key instead of the + // "sk-no-api-key" placeholder. + modelId: model?.id, ...capacityPayload, ...inferencePayload, ...embeddingPayload, @@ -1062,6 +1066,20 @@ export const ModelAddDialogV2 = ({ ctx.resolvedModelType === MODEL_TYPES.MULTI_EMBEDDING ? "multiEmbedding" : ctx.resolvedModelType; + // Never steal an occupied default slot: adding or editing a model is + // not an intent to change the tenant default. Persist only when the + // slot is empty (add: deterministic onboarding default, mirroring the + // backend backfill) or when the submitted model already occupies the + // slot (edit: keep the slot's apiKey/url in sync with the model's new + // values). + const currentSlotDisplayName = modelConfig?.[configKey]?.displayName; + const slotIsFree = !currentSlotDisplayName; + const submitsCurrentSlotModel = + !!model && + currentSlotDisplayName === (model.displayName || model.name); + if (!slotIsFree && !submitsCurrentSlotModel) { + return; + } const existingApiKey = modelConfig?.[configKey]?.apiConfig?.apiKey || ""; const nextModelConfig: SingleModelConfig = { id: 0, diff --git a/frontend/app/[locale]/models/components/model/ModelAdvancedSettings.tsx b/frontend/app/[locale]/models/components/model/ModelAdvancedSettings.tsx index 8c0d31b23..dbbc03c4d 100644 --- a/frontend/app/[locale]/models/components/model/ModelAdvancedSettings.tsx +++ b/frontend/app/[locale]/models/components/model/ModelAdvancedSettings.tsx @@ -200,8 +200,13 @@ const formatCustomValueForEditing = (raw: unknown): string => { /** Convert the editing-state __custom__ entries array into a clean wire dict. * Empty keys are dropped and duplicates collapse (last-wins). + * A plain object passes through as-is: the override save path builds the + * final dict directly and uses null values as explicit removal markers. */ const buildCustomDict = (raw: unknown): Record => { + if (raw && typeof raw === "object" && !Array.isArray(raw)) { + return { ...(raw as Record) }; + } const entries = Array.isArray(raw) ? (raw as [string, string][]) : []; const dict: Record = {}; for (const [k, v] of entries) { @@ -213,6 +218,73 @@ const buildCustomDict = (raw: unknown): Record => { return dict; }; +/** + * Merge model-level and override custom params into the editing entries + * shown in the override dialog. Model-level params render as plain editable + * rows (pre-2.6.0 merged display); override values win; null override values + * are explicit removal markers and hide the key entirely. + */ +export const mergeCustomParamsForEditing = ( + overrideCustoms: Record | null | undefined, + modelCustoms: Record | null | undefined +): [string, string][] => { + const merged: Record = {}; + for (const [k, v] of Object.entries(modelCustoms ?? {})) { + merged[k] = v; + } + for (const [k, v] of Object.entries(overrideCustoms ?? {})) { + if (v === null || v === undefined) { + delete merged[k]; + } else { + merged[k] = v; + } + } + return Object.entries(merged).map( + ([k, v]) => [k, formatCustomValueForEditing(v)] as [string, string] + ); +}; + +/** + * Diff the editing entries against model-level custom params for saving. + * Returns the wire dict stored in the override entry: + * - explicit values for keys that differ from the model level + * - null for keys whose inherited row the user deleted (the runtime pops + * null-marked keys from the merged request params) + * - keys equal to the model level are dropped (pure inherit, no override) + */ +export const diffCustomParamsForSave = ( + editingEntries: unknown, + modelCustoms: Record | null | undefined +): Record => { + const entries = Array.isArray(editingEntries) + ? (editingEntries as [string, string][]) + : []; + const editingDict: Record = {}; + for (const [k, v] of entries) { + if (k === "") continue; + const parsed = parseCustomValue(String(v ?? "")); + if (parsed === undefined) continue; + editingDict[k] = parsed; + } + const modelDict = modelCustoms ?? {}; + const result: Record = {}; + const keys = new Set([...Object.keys(modelDict), ...Object.keys(editingDict)]); + for (const k of keys) { + if (k in editingDict) { + if ( + !(k in modelDict) || + JSON.stringify(editingDict[k]) !== JSON.stringify(modelDict[k]) + ) { + result[k] = editingDict[k]; + } + } else if (k in modelDict) { + // Inherited row deleted by the user: persist an explicit removal. + result[k] = null; + } + } + return result; +}; + export const buildInferenceParamsPayload = ( value: ModelAdvancedSettingsValue ): { diff --git a/frontend/services/modelService.ts b/frontend/services/modelService.ts index 4a035aa23..820b54bb7 100644 --- a/frontend/services/modelService.ts +++ b/frontend/services/modelService.ts @@ -832,6 +832,9 @@ export const modelService = { temperature?: number; topP?: number; extraParams?: Record; + // Edit-dialog probes: when set, the backend falls back to the stored + // api_key for this model when apiKey is empty/placeholder. + modelId?: number; }, signal?: AbortSignal ): Promise => { @@ -841,6 +844,9 @@ export const modelService = { model_type: config.modelType, api_key: config.apiKey || "sk-no-api-key", base_url: config.baseUrl || "", + ...(config.modelId !== undefined + ? { probe_model_id: config.modelId } + : {}), ...(config.maxTokens !== undefined ? { max_tokens: config.maxTokens } : {}), diff --git a/sdk/nexent/core/gateway/modality/vlm/openai.py b/sdk/nexent/core/gateway/modality/vlm/openai.py index ddb4f55bc..b4c0804d4 100644 --- a/sdk/nexent/core/gateway/modality/vlm/openai.py +++ b/sdk/nexent/core/gateway/modality/vlm/openai.py @@ -224,7 +224,14 @@ async def check_connectivity(self) -> bool: """ if self._model is None: self._build_model() - module_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + # Anchor on the nexent package root: the probe asset lives in + # nexent/assets/. A relative dirname chain silently broke when this + # module moved deeper into the package (3 levels up resolved to + # .../core/gateway instead of the package root), so the local probe + # image was never found and every check fell back to the public URL - + # which fails in offline deployments. + import nexent + module_dir = os.path.dirname(os.path.abspath(nexent.__file__)) test_image_path = os.path.join(module_dir, "assets", "git-flow.png") if os.path.exists(test_image_path): base64_image = self.encode_image(test_image_path) diff --git a/sdk/nexent/core/models/openai_vlm.py b/sdk/nexent/core/models/openai_vlm.py index 983d24fce..44233a563 100644 --- a/sdk/nexent/core/models/openai_vlm.py +++ b/sdk/nexent/core/models/openai_vlm.py @@ -45,8 +45,13 @@ async def check_connectivity(self) -> bool: Returns: bool: True if the model responds successfully, otherwise False. """ - # Use local test image from images folder - use absolute path based on module location - module_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + # Use local test image from images folder - anchor on the nexent + # package root (asset lives in nexent/assets/). A relative dirname + # chain worked here only by coincidence of module depth and broke in + # the gateway copy of this probe; anchoring on the package makes both + # robust against future module moves. + import nexent + module_dir = os.path.dirname(os.path.abspath(nexent.__file__)) test_image_path = os.path.join(module_dir, "assets", "git-flow.png") if os.path.exists(test_image_path): base64_image = self.encode_image(test_image_path) From 3adf163aff91c6015cce8d60cccdcb26a48fb0f4 Mon Sep 17 00:00:00 2001 From: bernard1234 <840646206@qq.com> Date: Sun, 20 Sep 2026 14:07:33 +0800 Subject: [PATCH 3/9] cherry-pick: HITL bugfixes from PR #3948 into hotfix/v2.6.1 (#3959) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Bubfix: guarantee event order, harden chunk buffer, SSE-subscribe controller, and break adapter on terminal human_run (#3948) * fix(hitl): preserve correct event ordering between observer chunks and human_interaction requests Root cause: the worker thread writes human_interaction events synchronously via SQLAlchemy in ask_user, while model_output_thinking/parse observer messages flow through the async consumer and are flushed only on a batched threshold (32 chunks or 250ms). When the worker suspends before that flush fires, human_interaction gains a lower event_seq number than the already-buffered observer chunks, causing the SSE replay stream to show them in the wrong order. Fix: replace the plain async-for consumer loop with a manual asyncio.wait iterator using a 50ms timeout. Once the worker finishes producing model output (i.e. right before ask_user), the loop times out and flushes any buffered observer chunks to the DB first, guaranteeing they precede the subsequent human_interaction row. Empty queue idle periods are essentially zero-cost; overall DB write frequency stays on par with the original. * fix(hitl): guarantee observer chunks precede human_interaction in DB event order When the agent invokes ask_user, two independent write paths caused the human_interaction row to be persisted BEFORE model_output_thinking / parse chunks, breaking the SSE replay ordering: the worker thread writes HITL events synchronously via SQLAlchemy, while observer messages flow through the async consumer which only flushes on a batched threshold. Fix: introduce a thread-safe shared chunk buffer on RuntimeInteractionPort (port.add_chunk / port.take_chunks). The async consumer pushes every processed chunk there; the worker thread calls flush_chunks_until_idle() before dispatching any HITL event — it polls the shared buffer and waits for the async loop to drain the observer queue (20ms idle window, 500ms max wait), then persists every chunk in its own transaction. This guarantees chunk event_seq < human_interaction event_seq regardless of async scheduling latency. Also fix ImportError: openai 2.50 removed the httpx2 module. OpenAIModel now falls back from httpx2 to httpx at import time. * test(hitl): cover shared chunk buffer, flush_chunks_until_idle, and httpx2-fallback paths Add unit tests for the RuntimeInteractionPort thread-safe chunk buffer and the flush_chunks_until_idle poll loop that guarantees observer chunks precede human_interaction events in DB order. All five HITL entry points (dispatch / boundary / receipt / finish / _wait_until_ready) are verified to invoke the idle flush before opening their transaction. Also add two tests for the openai_llm httpx2 → httpx ImportError fallback introduced to support openai >= 2.50 where the httpx2 shim was removed: one covers the fallback path, one confirms httpx2 still wins when present. * fix(hitl): address 4 review comments — hard deadline, emit_in_flight, peek_chunks, try/except safety Fix 4 real issues flagged by github-code-review: 1. Non-resettable hard_deadline in flush_chunks_until_idle — previously reset on every drain, meaning a model that kept producing chunks could stall the worker forever. deadline is now computed once at entry and the sleep call clips to hard_deadline - now. 2. _emit_in_flight Event bridges the async emit path and the worker's idle poll. Without this, buffer-empty = 'persisted' was confused with buffer-empty = 'taken for emit but still in run_blocking queue'. The worker now checks both 'buffer empty for settle_ms' AND 'no emit in flight' before deciding the async side is truly idle. 3. peek_chunks() replaces the take-put-back pattern in _flush_if_due. Previously the async loop drained the buffer, decided it was not yet due, then put everything back. That transiently-empty window (16 us normally, arbitrarily long under GIL/GC/preemption) was enough for the worker's 20 ms poll to mis-fire. We now peek (read count, no drain) and only take_chunks when we actually intend to persist. 4. emit_chunks wrapped in try/except that puts drained chunks back into the shared buffer before re-raising, and finish() wraps its flush call in try/except: pass. Guarantees (a) no chunk loss on DB failure and (b) the terminal human_run row is always written even if the flush step fails. Tests added: - 12 pure-mock unit tests in test_runtime_port_chunk_buffer.py cover hard_deadline, _emit_in_flight, peek_chunks, begin_emit/end_emit, try/except path, and every HITL entry-point's flush-before-transaction. - 1 async execute_attempt integration test in new test_application_execute_attempt.py drives the full consumer loop through _flush_if_due (peek → take → begin/end_emit) and the final flush, verifying that every patch line added in application.py is hit. * perf(hitl): stop polling while SSE stream is active, fallback to 5s when disconnected When isRunning=true the EventSource already pushes human_interaction and human_execution events in real time — the 1.5s polling loop duplicated that work, hitting the DB and re-rendering the frontend for every tick. Disable polling entirely while the SSE stream is alive, and drop to 5s intervals only when the stream is closed (e.g. page load before the first run, or after a run finishes) so we can still discover WAITING_HUMAN requests that were created while the client was disconnected. Add isRunning to the useEffect dependency array so the polling cadence resets immediately when the SSE connection state changes. * perf(hitl): replace polling with SSE subscription and move snapshot off the write path Frontend — /conversation polling → /{run_id}/events SSE: - Replace the 5s conversation snapshot polling with a native EventSource subscription to the backend's /{run_id}/events SSE stream. Discovery is now one-shot: conversationId change and the agent stream pause (isRunning true→false), the exact moment a HITL run is most likely to exist. The SSE stream then keeps run state live with native auto-reconnect. - Add dual guards inside refresh() to absorb the thundering herd from adapter.onHumanInteractionEvent (fires once per HITL SSE chunk) plus our own SSE effect: (1) in-flight dedupe — one snapshot absorbs all concurrent callers and returns cached state; (2) 3s minimum interval so bursts after the in-flight resolves do not immediately re-hit DB. - Use a runRef mirror so refresh() stays stable and downstream effects do not re-run on every snapshot. - Detect terminal status inside SSE onmessage and proactively es.close() to prevent EventSource from reconnecting forever against COMPLETED runs. Backend — snapshot off the write path: - Add repository.read_only() context manager: plain SELECT without WITH FOR UPDATE, no transaction, no flush, no _expire scan. Pure reads must not contend with worker writes on the same row lock. - Add service.light_snapshot() using read_only. Retain snapshot() as a writer-path API for any future lock-held callers. - Route conversation_snapshot, run snapshot endpoint, and both snapshot calls inside stream_run() through light_snapshot. - Move expiration to the writer path: decide() still calls _expire inline before processing each request, and expire_waiting() remains the periodic scheduler sweep. Impact: conversation snapshot calls drop from 12+/min (polling) or 10+/s (burst from adapter + SSE) to at most one every 3s. Each call is now two plain SELECTs instead of a lock-held transaction with a possible write from _expire. Read and write paths are fully decoupled. * fix(hitl): detect terminal human_run in adapter and break stream so isRunning flips false After a HITL run reaches FAILED/COMPLETED, Assistant-UI's isRunning stayed true — the stop button remained visible and new messages went into the queue buffer instead of being sent normally. The root cause is that isRunning is driven entirely by the ChatModelRun generator lifetime, which only returns when the backend SSE HTTP connection closes (reader.read() -> done=true). The backend stream_run loop can hang on heartbeat even after the run is terminal when the SSE was opened during WAITING_HUMAN with attempt_active=true: the break condition requires both cursor >= event_seq AND (terminal status OR WAITING_HUMAN with attempt_active=false and empty rows). If continueHitl fires mid-flight with a stale after_event, the cursor never catches up, so the SSE stays alive forever and the generator never returns. Stop depending on the backend closing first. Inside the adapter's SSE chunk loop, detect a terminal human_run event (status in COMPLETED, FAILED, STOPPED, EXPIRED), set a hitlTerminal flag, break the inner for-loop, and let the outer while-loop exit via the same flag on the next iteration. Assistant-UI sees the generator return and flips isRunning false immediately. Only affects HITL streams — the normal non-HITL agent path never emits human_run events so this branch is never taken. * test(hitl): update mock from snapshot to light_snapshot after read-path refactor test_human_interaction_app.py still mocked service.snapshot after commit 288ae4e69 moved conversation_snapshot and the run snapshot endpoint to service.light_snapshot (read-only path, no lock, no _expire). The fixture return_value and the two assert_called_once_with/assert_not_called assertions all referenced the old method name, causing CI to fail because MagicMock.snapshot was never called. * test(hitl): raise diff coverage above the 90% merge gate Codecov reported 70.43% patch coverage (target 90%) because new error and race paths in the HITL changes had no tests. Add mocked unit tests for: leftover chunk flush in execute_attempt's finally block before the failed finish, CancelledError scope and stop-event fallbacks, RunTerminated finish race, recovery-required outcome, and chunk iterator aclose failure tolerance; runtime_port in-flight emit busy detection, chunk restoration when emit_chunks raises, and terminal status persistence on flush failure; light_snapshot/read_only service behavior with signed tenant and user scoping; and the httpx fallback when openai._base_client.httpx2 is absent. Measured locally with CI-equivalent per-file pytest isolation: patch coverage 202/202 = 100%. * style(hitl): unify comment style across HITL changes Merge explanatory inline comments into docstrings, keep single-line comments for inline notes, convert TypeScript block notes to JSDoc, and drop banner/separator lines. Comment-level changes only, no behavior change. * style(hitl): unify comment style across HITL changes Merge explanatory inline comments into docstrings, keep single-line comments for inline notes, convert TypeScript block notes to JSDoc, and drop banner/separator lines. Comment-level changes only, no behavior change. * refactor(hitl-test): dedupe fake port setup to satisfy SonarCloud duplication gate SonarCloud failed the quality gate with new_duplicated_lines_density=5.2% (threshold 3%), caused solely by test_application_execute_attempt.py: the inline _Port stub in the flush test and the one in _run_execute_attempt duplicated ~69 lines (2 CPD blocks, 14.4% file density). Extract a shared _build_port_class/_make_port_factory plus a _patched_application context manager and _execute_attempt_args so both call sites reuse a single definition; drop dead code (last_flush, install/monkeypatches, unused imports) and fix the latent bare-contextmanager NameError by using contextlib.contextmanager. No behavioral change; all 8 tests pass. * fix(sdk): restore httpx2 → httpx ImportError fallback in openai_llm openai >= 2.50 removed the httpx2 shim from openai._base_client. The bare import httpx2 causes ImportError in CI and on systems with recent openai versions. This restores the try/except fallback introduced in PR #3948 commit 9521b934 and later accidentally reverted by commit d086da259. * Revert "fix(sdk): restore httpx2 → httpx ImportError fallback in openai_llm" This reverts commit 8e06ce348dfc8c34baf42c6bfa8211883d0a2c61. --- backend/apps/human_interaction_app.py | 4 +- backend/database/human_interaction_db.py | 10 + .../services/human_interaction/application.py | 93 +++- .../human_interaction/runtime_port.py | 89 ++++ backend/services/human_interaction/service.py | 24 +- .../adapter/remote-chat-model-adapter.ts | 13 +- .../useHumanInteractionController.ts | 145 +++++- .../backend/app/test_human_interaction_app.py | 15 +- .../test_application_execute_attempt.py | 421 ++++++++++++++++++ .../test_human_interaction_snapshot.py | 150 +++++++ .../services/test_human_interaction_stream.py | 6 +- .../test_runtime_port_chunk_buffer.py | 356 +++++++++++++++ 12 files changed, 1280 insertions(+), 46 deletions(-) create mode 100644 test/backend/services/test_application_execute_attempt.py create mode 100644 test/backend/services/test_human_interaction_snapshot.py create mode 100644 test/backend/services/test_runtime_port_chunk_buffer.py diff --git a/backend/apps/human_interaction_app.py b/backend/apps/human_interaction_app.py index 37a73dc65..921e1340b 100644 --- a/backend/apps/human_interaction_app.py +++ b/backend/apps/human_interaction_app.py @@ -50,12 +50,12 @@ async def capabilities(identity=Depends(identity_dependency)): async def conversation_snapshot(conversation_id: int, identity=Depends(identity_dependency)): def read(service, tenant_id, user_id): run_id = service.repository.latest(tenant_id, user_id, conversation_id) - return service.snapshot(run_id, tenant_id, user_id) if run_id else None + return service.light_snapshot(run_id, tenant_id, user_id) if run_id else None return await _call(identity, read) @result.get("/{run_id}") async def snapshot(run_id: str, identity=Depends(identity_dependency)): - return await _call(identity, lambda service, tenant, user: service.snapshot(run_id, tenant, user)) + return await _call(identity, lambda service, tenant, user: service.light_snapshot(run_id, tenant, user)) @result.get("/{run_id}/events") async def events(run_id: str, after_event: int = Query(0, ge=0), identity=Depends(identity_dependency)): diff --git a/backend/database/human_interaction_db.py b/backend/database/human_interaction_db.py index 6269f5743..19575bc4e 100644 --- a/backend/database/human_interaction_db.py +++ b/backend/database/human_interaction_db.py @@ -168,6 +168,16 @@ def transaction(self, run_id, tenant_id=None, user_id=None): if tx: tx.flush() + @contextmanager + def read_only(self, run_id, tenant_id=None, user_id=None): + """Snapshot reads must not contend with write paths. No lock, no writes.""" + with self.session_factory() as session: + conditions = [HumanRun.run_id == run_id, HumanRun.delete_flag == "N"] + if tenant_id is not None: + conditions.extend([HumanRun.tenant_id == tenant_id, HumanRun.user_id == user_id]) + run = session.scalar(select(HumanRun).where(*conditions)) + yield RunTransaction(session, run, self.validator, user_id) if run else None + @contextmanager def creation(self, run): with self.session_factory() as session, session.no_autoflush: diff --git a/backend/services/human_interaction/application.py b/backend/services/human_interaction/application.py index 05ea8bbd8..d2251fa32 100644 --- a/backend/services/human_interaction/application.py +++ b/backend/services/human_interaction/application.py @@ -3,6 +3,7 @@ import asyncio import json import time +from contextlib import suppress from functools import lru_cache from fastapi.responses import StreamingResponse @@ -162,25 +163,83 @@ def authorize(): # Deployments opt into the conservative approval gate independently. port.allowed_tools = _allowed_tool_names(config.tools) run_info.human_interaction = runtime_type(port) - buffered_chunks = [] + + _FLUSH_INTERVAL = 0.05 + _FLUSH_BATCH = 16 + last_flush = time.monotonic() - async for chunk in _stream_agent_chunks( + + async def _flush_if_due() -> None: + """Flush buffered chunks on the interval/batch trigger to keep DB order. + + The async loop flushes on its own; the worker flushes again before + each HITL transaction. Peek first and never take-put-back — that + would open a transiently empty window the worker's idle poll can + mistake for "async is done". + """ + nonlocal last_flush + if (time.monotonic() - last_flush >= _FLUSH_INTERVAL + or port.peek_chunks() >= _FLUSH_BATCH): + buffered = port.take_chunks() + if buffered: + try: + port.begin_emit() + await run_blocking( + "hitl-port-emit_chunks", port.emit_chunks, buffered, + lane="control-io", owner=__name__, + ) + last_flush = time.monotonic() + finally: + port.end_emit() + + chunk_iter = _stream_agent_chunks( agent_request=request, user_id=identity["user_id"], tenant_id=identity["tenant_id"], agent_run_info=run_info, memory_ctx=memory_context, - ): - buffered_chunks.append(chunk) - if len(buffered_chunks) >= 32 or time.monotonic() - last_flush >= 0.25: - await run_blocking( - "hitl-port-emit_chunks", port.emit_chunks, buffered_chunks, lane="control-io", - owner=__name__, + ).__aiter__() + + anext_task: asyncio.Task | None = None + try: + anext_task = asyncio.create_task(chunk_iter.__anext__()) + while True: + # On timeout the task stays pending; the worker may resume it later. + done, pending = await asyncio.wait( + {anext_task}, + timeout=_FLUSH_INTERVAL, + return_when=asyncio.FIRST_COMPLETED, ) - buffered_chunks = [] - last_flush = time.monotonic() - if buffered_chunks: - await run_blocking( - "hitl-port-emit_chunks", port.emit_chunks, buffered_chunks, lane="control-io", - owner=__name__, - ) + if pending: + # No new chunk yet; flush what we have to keep DB order. + await _flush_if_due() + continue + + task = done.pop() + try: + chunk = task.result() + except StopAsyncIteration: + break + port.add_chunk(chunk) + await _flush_if_due() + anext_task = asyncio.create_task(chunk_iter.__anext__()) + finally: + if anext_task is not None: + anext_task.cancel() + with suppress(asyncio.CancelledError): + await anext_task + try: + await chunk_iter.aclose() + except Exception: + pass + # Final flush so leftover chunks precede finish() in DB order. + leftover = port.take_chunks() + if leftover: + try: + port.begin_emit() + await run_blocking( + "hitl-port-emit_chunks", port.emit_chunks, leftover, + lane="control-io", owner=__name__, + ) + finally: + port.end_emit() await run_blocking( "hitl-port-finish", port.finish, run_info.attempt_outcome or "failed", lane="control-io", owner=__name__, @@ -264,7 +323,7 @@ async def release(self, job_id, owner_id): async def stream_run(run_id, tenant_id, user_id, *, after=0): service = require_enabled() snapshot = await run_blocking( - "hitl-service-snapshot", service.snapshot, run_id, tenant_id, user_id, lane="control-io", + "hitl-service-snapshot", service.light_snapshot, run_id, tenant_id, user_id, lane="control-io", owner=__name__, ) if after > snapshot["event_seq"]: @@ -276,7 +335,7 @@ async def events(): while True: # Re-check ownership and deadlines for reconnecting subscribers. current = await run_blocking( - "hitl-service-snapshot", service.snapshot, run_id, tenant_id, user_id, lane="control-io", + "hitl-service-snapshot", service.light_snapshot, run_id, tenant_id, user_id, lane="control-io", owner=__name__, ) rows = await run_blocking( diff --git a/backend/services/human_interaction/runtime_port.py b/backend/services/human_interaction/runtime_port.py index 373e065aa..62b01b417 100644 --- a/backend/services/human_interaction/runtime_port.py +++ b/backend/services/human_interaction/runtime_port.py @@ -1,5 +1,6 @@ """Fenced tool dispatch adapter. Approval and STARTED are consumed under one run lock.""" +import threading import time from contextlib import contextmanager from datetime import timezone @@ -17,6 +18,14 @@ class RuntimeInteractionPort: def __init__(self, service, identity, owner_id, authorize, allowed_tools=(), *, live_resume=False, stop_event=None): + """Bind the port and set up async-worker chunk coordination. + + ``_chunk_buffer`` stages processed observer chunks: the async consumer + appends via add_chunk and the worker flushes before each HITL event so + chunk rows precede human_interaction rows in DB event order. + ``_emit_in_flight`` marks an in-flight hand-off to the DB thread so an + empty buffer is not mistaken for "async is idle". + """ self.service = service self.repository = service.repository self.cipher = service.cipher @@ -29,10 +38,73 @@ def __init__(self, service, identity, owner_id, authorize, allowed_tools=(), *, self.allowed_tools = frozenset(allowed_tools) self.live_resume = live_resume self.stop_event = stop_event + self._chunk_buffer: list[str] = [] + self._chunk_buffer_lock = threading.Lock() + self._emit_in_flight = threading.Event() with self.transaction() as tx: self.checkpoint = self.cipher.open(tx.run.checkpoint) self.request_payload = self.cipher.open(tx.run.request_payload) + def add_chunk(self, chunk: str) -> None: + """Append a processed chunk from the async consumer to the shared buffer.""" + with self._chunk_buffer_lock: + self._chunk_buffer.append(chunk) + + def take_chunks(self) -> list[str]: + """Atomically drain the shared buffer for persistence.""" + with self._chunk_buffer_lock: + chunks = self._chunk_buffer + self._chunk_buffer = [] + return chunks + + def peek_chunks(self) -> int: + """Return the number of buffered chunks without draining them.""" + with self._chunk_buffer_lock: + return len(self._chunk_buffer) + + def begin_emit(self) -> None: + """Signal that the async consumer is handing chunks to run_blocking(emit_chunks).""" + self._emit_in_flight.set() + + def end_emit(self) -> None: + """Signal that the async consumer's run_blocking(emit_chunks) has returned.""" + self._emit_in_flight.clear() + + def flush_chunks_until_idle(self, *, max_wait_ms: int = 500, settle_ms: int = 20) -> None: + """Wait for the async consumer to drain the observer queue, then persist. + + Worker-thread only. Returns once the buffer has stayed empty for + ``settle_ms`` with no in-flight emit; ``max_wait_ms`` is a hard + upper bound measured from entry and never resets. + """ + hard_deadline = time.monotonic() + max_wait_ms / 1000.0 + idle_since: float | None = None + + while True: + now = time.monotonic() + if now >= hard_deadline: + break + + chunks = self.take_chunks() + if chunks: + try: + self.emit_chunks(chunks) + except Exception: + # Never lose drained chunks; put them back for retry. + for chunk in chunks: + self.add_chunk(chunk) + raise + idle_since = None + elif self._emit_in_flight.is_set(): + # Empty buffer is not idle while an emit is in flight. + idle_since = None + elif idle_since is None: + idle_since = now + elif now - idle_since >= settle_ms / 1000.0: + break + + time.sleep(min(settle_ms / 1000.0, hard_deadline - now)) + @contextmanager def transaction(self, *, receipt=False): with self.repository.transaction(self.run_id, self.tenant_id, self.user_id) as tx: @@ -75,6 +147,8 @@ def boundary(self, checkpoint): self.authorize() suspended = False feedback = None + # Flush so buffered chunks precede the human_run status row in DB order. + self.flush_chunks_until_idle() with self.transaction() as tx: tx.run.checkpoint = self.cipher.seal(checkpoint) if tx.run.pause_requested: @@ -139,6 +213,8 @@ def _wait_until_ready(self): raise RunTerminated("Execution lease is no longer valid") self.service._expire(tx) if tx.run.status == "READY": + # Flush so lingering chunks precede the human_run row. + self.flush_chunks_until_idle() tx.run.status = "RUNNING" tx.emit({"type": "human_run", "content": { "run_id": self.run_id, "status": "RUNNING", @@ -156,6 +232,9 @@ def dispatch(self, slot, tool, arguments, *, interaction=None): suspended = False steering_requested = False outcome = None + # Flush before the transaction so chunk rows get lower seq numbers + # than any subsequent human_interaction row. + self.flush_chunks_until_idle() with self.transaction() as tx: if tx.run.pause_requested: self.service._request_steering(tx) @@ -234,6 +313,8 @@ def dispatch(self, slot, tool, arguments, *, interaction=None): return outcome def receipt(self, slot, result, *, uncertain=False): + # Flush so the DB event order is chunks → human_execution. + self.flush_chunks_until_idle() with self.transaction(receipt=True) as tx: execution = tx.execution(slot) if execution is None or execution.status != "STARTED": @@ -274,6 +355,14 @@ def visible_guidance(self): ] def finish(self, outcome): + """Write the terminal run status, flushing buffered chunks first. + + A flush failure must not prevent the terminal status row. + """ + try: + self.flush_chunks_until_idle() + except Exception: + pass with self.transaction(receipt=True) as tx: if tx.run.status in {"STOPPED", "EXPIRED"}: return diff --git a/backend/services/human_interaction/service.py b/backend/services/human_interaction/service.py index 5a2453d40..4e5f1caba 100644 --- a/backend/services/human_interaction/service.py +++ b/backend/services/human_interaction/service.py @@ -96,13 +96,29 @@ def _expire(self, tx): return expired def snapshot(self, run_id, tenant_id, user_id): + """Writer-path snapshot: takes a lock and may expire stale requests. + + Kept for the SSE stream's initial handshake where a terminal-status + request needs a lock-held check. Prefer :meth:`light_snapshot` for + read-only callers. + """ with self.repository.transaction(run_id, tenant_id, user_id) as tx: run = self.require(tx) self._expire(tx) - return {"run_id": run.run_id, "conversation_id": run.conversation_id, "status": run.status, - "event_seq": run.event_seq, "pause_requested": bool(run.pause_requested), - "attempt_active": bool(run.lock_until and run.lock_until > utcnow()), - "requests": [self._project_request(item, run.run_id) for item in tx.requests() if item.status == "PENDING"]} + return self._build_snapshot(run, tx) + + def light_snapshot(self, run_id, tenant_id, user_id): + """Read-only snapshot: no lock, no writes, no expiration scan.""" + with self.repository.read_only(run_id, tenant_id, user_id) as tx: + if tx is None: + raise InteractionError("Human interaction run was not found", 404) + return self._build_snapshot(tx.run, tx) + + def _build_snapshot(self, run, tx): + return {"run_id": run.run_id, "conversation_id": run.conversation_id, "status": run.status, + "event_seq": run.event_seq, "pause_requested": bool(run.pause_requested), + "attempt_active": bool(run.lock_until and run.lock_until > utcnow()), + "requests": [self._project_request(item, run.run_id) for item in tx.requests() if item.status == "PENDING"]} def request(self, tx, *, kind, slot, action_digest, payload): request = HumanRequest( diff --git a/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts b/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts index e40eead40..c80bf3819 100644 --- a/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts +++ b/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts @@ -2034,11 +2034,12 @@ export const remoteChatModelAdapter: ChatModelAdapter = { let firstTokenTime: number | undefined; let toolCallCount = 0; let storedTiming: ReturnType | null = null; + let hitlTerminal: boolean | undefined = undefined; try { while (true) { const { done, value } = await reader.read(); - if (done) break; + if (done || hitlTerminal) break; buffer += decoder.decode(value, { stream: true }); @@ -2058,6 +2059,16 @@ export const remoteChatModelAdapter: ChatModelAdapter = { if (value && typeof value.run_id === "string") humanRunId = value.run_id; custom?.onHumanInteractionEvent?.(); + // Terminal HITL status: stop reading so isRunning flips false + // without waiting for the backend to close the stream. + if ( + value && + typeof value.status === "string" && + ["COMPLETED", "FAILED", "STOPPED", "EXPIRED"].includes(value.status) + ) { + hitlTerminal = true; + break; + } continue; } if ( diff --git a/frontend/features/humanInteraction/useHumanInteractionController.ts b/frontend/features/humanInteraction/useHumanInteractionController.ts index fd5e0aa1f..9c45849c4 100644 --- a/frontend/features/humanInteraction/useHumanInteractionController.ts +++ b/frontend/features/humanInteraction/useHumanInteractionController.ts @@ -2,6 +2,7 @@ import { useCallback, useEffect, useMemo, useRef, useState } from "react"; +import { API_BASE_URL } from "@/services/api"; import { humanInteractionClient, type HumanDecision, @@ -18,6 +19,13 @@ const ACTIVE_STATUSES = new Set([ "RECOVERY_REQUIRED", ]); const STREAM_RECONNECT_STATUSES = new Set(["READY", "RUNNING"]); +const TERMINAL_STATUSES = new Set([ + "COMPLETED", + "FAILED", + "STOPPED", + "EXPIRED", + "RECOVERY_REQUIRED", +]); export interface HumanInteractionController { available: boolean; @@ -59,8 +67,16 @@ export function useHumanInteractionController({ new Map() ); const onEnabledChangeRef = useRef(onEnabledChange); + // Rate-limit snapshots: dedupe in-flight callers and enforce a 3s min interval. + const refreshInFlight = useRef(false); + const lastSnapshotAt = useRef(0); + const MIN_SNAPSHOT_INTERVAL_MS = 3000; + // Mirror of the latest `run` state; keeps `refresh` dependencies stable. + const runRef = useRef(null); activeConversation.current = conversationId; + const [activeRunId, setActiveRunId] = useState(null); + useEffect(() => { onEnabledChangeRef.current = onEnabledChange; }, [onEnabledChange]); @@ -84,11 +100,24 @@ export function useHumanInteractionController({ const refresh = useCallback(async () => { const requestedConversation = conversationId; - const sequence = ++refreshSequence.current; if (!conversationId) { + runRef.current = null; setRun(null); + setActiveRunId(null); return null; } + // Guard 1: coalesce while a snapshot is in flight. + if (refreshInFlight.current) { + return runRef.current; + } + // Guard 2: respect the minimum snapshot interval. + const now = Date.now(); + if (now - lastSnapshotAt.current < MIN_SNAPSHOT_INTERVAL_MS) { + return runRef.current; + } + refreshInFlight.current = true; + lastSnapshotAt.current = now; + const sequence = ++refreshSequence.current; try { const value = await humanInteractionClient.conversation(conversationId); if ( @@ -97,10 +126,13 @@ export function useHumanInteractionController({ ) { return null; } + runRef.current = value; setRun(value); setError(""); - if (!value || !STREAM_RECONNECT_STATUSES.has(value.status)) { - resume.current = null; + if (value) { + if (!STREAM_RECONNECT_STATUSES.has(value.status)) { + resume.current = null; + } } return value; } catch (cause) { @@ -112,29 +144,107 @@ export function useHumanInteractionController({ } setError(cause instanceof Error ? cause.message : String(cause)); return null; + } finally { + refreshInFlight.current = false; } }, [conversationId]); + /** + * Discovery: one-shot snapshots at conversation change and agent pause. + * After a snapshot reveals a run, the SSE subscription below keeps it live. + */ + useEffect(() => { + runRef.current = null; setRun(null); setError(""); resume.current = null; decisions.current.clear(); + setActiveRunId(null); if (!available || !conversationId) return; + void refresh(); + // One-shot fire on conversation change; do not re-run when refresh changes. + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [available, conversationId]); - let active = true; - let timer: ReturnType; - const poll = async () => { - if (!active) return; - await refresh(); - if (active) timer = setTimeout(poll, 1500); + useEffect(() => { + // Agent stream paused → a HITL run may have just been created (ask_user). + if (isRunning || !available || !conversationId) return; + void refresh(); + }, [isRunning, available, conversationId, refresh]); + + // Retry tick for SSE resubscribe after transient network disconnects; + // terminal closes must not trigger a reconnect. + const [retryTick, setRetryTick] = useState(0); + + /** + * SSE subscription to the run events stream; every event triggers one + * rate-limited refresh(). Native auto-reconnect handles transient drops. + */ + + useEffect(() => { + if (!activeRunId) return; + + let tornDown = false; + let stoppedByUs = false; + const url = `${API_BASE_URL}/agent/human-interactions/${activeRunId}/events?after_event=${run?.event_seq ?? 0}`; + const es = new EventSource(url); + + es.onmessage = (ev) => { + if (tornDown || stoppedByUs) return; + if (!ev.data) return; + // Close proactively on terminal status so EventSource stops + // auto-reconnecting against a run that already finished. + try { + const parsed = JSON.parse(ev.data); + if ( + parsed.type === "human_run" && + parsed.content && + typeof parsed.content === "object" && + TERMINAL_STATUSES.has(parsed.content.status) + ) { + stoppedByUs = true; + es.close(); + void refresh(); + return; + } + } catch { + // Not JSON — refresh without parsing. + } + void refresh(); + }; + + es.onerror = () => { + if (tornDown || stoppedByUs) return; + // Let native auto-reconnect handle transient drops; fall back to a + // manual resubscribe once the browser gives up and the run is alive. + void refresh().then((latest) => { + if (tornDown || stoppedByUs) return; + if (!latest || TERMINAL_STATUSES.has(latest.status)) return; + es.close(); + stoppedByUs = true; + setTimeout(() => { + if (!tornDown) setRetryTick((c) => c + 1); + }, 2000); + }); }; - void poll(); + return () => { - active = false; - clearTimeout(timer); + tornDown = true; + stoppedByUs = true; + es.close(); }; - }, [available, conversationId, refresh]); + }, [activeRunId, retryTick]); + + // Subscribe to a discovered run's events stream; null tears it down. + + useEffect(() => { + if (!run) { + setActiveRunId((prev) => (prev ? null : prev)); + return; + } + setActiveRunId((prev) => (prev !== run.run_id ? run.run_id : prev)); + }, [run]); useEffect(() => { if (isRunning || !resume.current) return; @@ -168,6 +278,7 @@ export function useHumanInteractionController({ action ); refreshSequence.current += 1; + runRef.current = nextRun; setRun(nextRun); } catch (cause) { setError(cause instanceof Error ? cause.message : String(cause)); @@ -210,13 +321,15 @@ export function useHumanInteractionController({ runId: currentRun.run_id, after: currentRun.event_seq, }; - setRun({ + const nextRun = { ...currentRun, - status: "READY", + status: "READY" as const, requests: currentRun.requests.filter( (request) => request.request_id !== item.request_id ), - }); + }; + runRef.current = nextRun; + setRun(nextRun); } } catch (cause) { const message = cause instanceof Error ? cause.message : String(cause); diff --git a/test/backend/app/test_human_interaction_app.py b/test/backend/app/test_human_interaction_app.py index 3844075c4..671f66e75 100644 --- a/test/backend/app/test_human_interaction_app.py +++ b/test/backend/app/test_human_interaction_app.py @@ -19,7 +19,7 @@ @pytest.fixture def boundary(monkeypatch): service = MagicMock() - service.snapshot.return_value = {"run_id": RUN, "status": "WAITING_HUMAN"} + service.light_snapshot.return_value = {"run_id": RUN, "status": "WAITING_HUMAN"} service.repository.latest.return_value = RUN service.decide.return_value = {"accepted": True} service.control.return_value = {"run_id": RUN} @@ -39,7 +39,16 @@ def test_snapshot_and_conversation_use_signed_tenant_and_user(boundary): assert response.status_code == 200 verify.assert_called_once_with("Bearer internal-token") service.repository.latest.assert_called_once_with("tenant-a", "owner", 7) - service.snapshot.assert_called_once_with(RUN, "tenant-a", "owner") + service.light_snapshot.assert_called_once_with(RUN, "tenant-a", "owner") + + +def test_run_snapshot_uses_signed_tenant_and_user(boundary): + client, service, verify = boundary + response = client.get(BASE + f"/{RUN}", headers={"Authorization": "Bearer internal-token"}) + assert response.status_code == 200 + assert response.json() == {"run_id": RUN, "status": "WAITING_HUMAN"} + verify.assert_called_once_with("Bearer internal-token") + service.light_snapshot.assert_called_once_with(RUN, "tenant-a", "owner") @pytest.mark.parametrize("path", ["/capabilities", f"/{RUN}", "/conversation/7", f"/{RUN}/events"]) @@ -47,7 +56,7 @@ def test_invalid_internal_token_is_rejected(boundary, path): client, service, verify = boundary verify.side_effect = UnauthorizedError("Invalid internal runtime token") assert client.get(BASE + path).status_code == 401 - service.snapshot.assert_not_called() + service.light_snapshot.assert_not_called() @pytest.mark.parametrize("status", [404, 409, 410, 422, 503]) diff --git a/test/backend/services/test_application_execute_attempt.py b/test/backend/services/test_application_execute_attempt.py new file mode 100644 index 000000000..cff6064a2 --- /dev/null +++ b/test/backend/services/test_application_execute_attempt.py @@ -0,0 +1,421 @@ +"""Unit test for application.py execute_attempt async consumer loop. + +Covers the _flush_if_due batch/threshold flush (peek-then-take, begin_emit / +end_emit lifecycle) and the final flush in the ``finally`` block. No real +database or agent — every collaborator is mocked. This lets us exercise the +patch lines added in the HITL reorder PR without depending on the full +Postgres-backed test suite. +""" + +import asyncio +import contextlib +import threading +import types +from unittest.mock import MagicMock, patch + +import pytest + + +async def _run_blocking_mock(*args, **kwargs): + """Drop-in replacement for ``run_blocking`` that calls ``fn(*a, **kw)`` + in the current thread and awaits nothing extra. This MUST be async so + ``await run_blocking(...)`` works when patched in. + """ + fn = args[1] + rest = args[2:] + return fn(*rest) + + +async def _authorize_mock(*args, **kwargs): + return None + + +def _build_port_class(on_finish): + """Build a DB-free ``RuntimeInteractionPort`` subclass for consumer-loop tests. + + Storage collaborators are stubbed in-memory: chunk emissions are tracked + on ``self._emits`` and terminal outcomes are forwarded to ``on_finish`` + instead of hitting the database. + """ + from services.human_interaction.runtime_port import RuntimeInteractionPort + + class _Port(RuntimeInteractionPort): + def __init__(self, *a, **kw): + self._chunk_buffer = [] + self._chunk_buffer_lock = threading.Lock() + self._emit_in_flight = threading.Event() + self.service = MagicMock() + self.run_id = "run-1" + self.tenant_id = "tenant-1" + self.user_id = "user-1" + self.owner_id = "worker" + self.fence = "fence-1" + self.request_payload = { + "runtime_mode": "native-live-v1", + "request": {"agent_id": 1, "conversation_id": 7, "query": "hi", "enable_hitl": True}, + "language": "en", + "runtime_metadata": {}, + "runtime_metadata_version": 1, + "runtime_knowledge_context": None, + } + self.checkpoint = None + self.allowed_tools = frozenset() + self.live_resume = True + self.stop_event = None + self._emits: list[list[str]] = [] + + def transaction(self, *a, **kw): + return contextlib.contextmanager(lambda: iter([MagicMock()]))() + + def context_snapshot(self, items): + return [] + + def bind_catalog(self, catalog): + return None + + def emit_chunks(self, chunks): + self._emits.append(list(chunks)) + + def finish(self, outcome): + on_finish(outcome) + + return _Port + + +def _make_port_factory(fake_info, port_ref, on_finish, port_hook=None): + """Build the ``RuntimeInteractionPort`` replacement used by the tests. + + The returned factory instantiates the fake port, applies the optional + ``port_hook`` for per-test tweaking, wires it into ``fake_info`` and + records the instance in ``port_ref`` for later assertions. + """ + port_cls = _build_port_class(on_finish) + + def make_port(*_a, **_kw): + p = port_cls() + if port_hook is not None: + port_hook(p) + port_ref["p"] = p + fake_info.human_interaction = MagicMock() + fake_info.human_interaction.port = p + return p + + return make_port + + +def _execute_attempt_args(): + """Standard job/payload and owner handed to ``execute_attempt``.""" + return ( + types.SimpleNamespace( + job_id="j", + payload={ + "run_id": "run-1", + "tenant_id": "tenant-1", + "user_id": "user-1", + "conversation_id": 7, + }, + ), + types.SimpleNamespace(owner_id="worker"), + ) + + +@contextlib.contextmanager +def _patched_application(make_port, fake_stream, prepare_mock): + """Patch ``execute_attempt``'s full dependency chain; yields the module.""" + from services.human_interaction import application + + with patch.object(application, "get_service", lambda: MagicMock()), \ + patch.object(application, "RuntimeInteractionPort", make_port), \ + patch.object(application, "authorize_run", _authorize_mock), \ + patch("nexent.core.concurrency.run_blocking", _run_blocking_mock), \ + patch("management.services.agent.run.prepare_agent_run", prepare_mock), \ + patch("management.services.agent.run._stream_agent_chunks", fake_stream), \ + patch("management.services.agent.run._unregister_agent_run_after_execution", + lambda *_a, **_kw: None), \ + patch("agents.agent_run_manager.agent_run_manager.unregister_agent_run", + lambda *_: None): + yield application + + +async def test_flush_if_due_uses_peek_then_take_without_transiently_empty_buffer(): + """_flush_if_due never does take-put-back — it peeks first, then only + drains when actually ready to persist. Eliminates the race where the + worker's idle poll sees an empty buffer between take and put-back. + """ + calls: list[tuple] = [] + + async def fake_stream(**_): + for t in range(3): + yield f"chunk-{t}" + # A small sleep lets the timeout path fire once. + await asyncio.sleep(0.07) + yield "final" + + class FakeInfo: + human_interaction = None + agent_config = MagicMock() + agent_config.tools = [] + context_input = MagicMock(items=()) + model_config_list = [] + runtime_metadata = {} + attempt_outcome = None + + fake_info = FakeInfo() + port_ref: dict = {} + make_port = _make_port_factory( + fake_info, port_ref, lambda outcome: calls.append(("finish", outcome))) + + async def prepare_mock(**kwargs): + return fake_info, None + + with _patched_application(make_port, fake_stream, prepare_mock) as application: + try: + await application.execute_attempt(*_execute_attempt_args()) + except StopAsyncIteration: + # The async consumer loop hit the StopAsyncIteration re-raised + # from the finally block — that's fine, we still ran through + # the entire consumer loop including _flush_if_due + final flush. + pass + except Exception as exc: + pytest.fail(f"execute_attempt raised unexpected {type(exc).__name__}: {exc!r}") + + port = port_ref["p"] + total_persisted = sum(len(batch) for batch in port._emits) + # Timing-sensitive — the sleep in fake_stream may cause the "final" chunk + # to land in either the timed _flush_if_due path or the final-flush path. + # Either way, every chunk that was added must be accounted for. + assert total_persisted >= 3, ( + f"Expected >=3 chunks persisted, got {total_persisted} batches={port._emits}" + ) + + # begin_emit must have been paired with end_emit — otherwise we would + # have seen _emit_in_flight still set after the loop. + assert not port._emit_in_flight.is_set(), "begin_emit without end_emit leaked" + + +# --- exception / finally path coverage --------------------------------------- + +class _FakeInfo: + human_interaction = None + agent_config = MagicMock() + agent_config.tools = [] + context_input = MagicMock(items=()) + model_config_list = [] + runtime_metadata = {} + attempt_outcome = None + cancellation_scope = None + stop_event = None + + +async def _run_execute_attempt(fake_stream, fake_info, *, port_hook=None, refs=None): + """Run execute_attempt with the full dependency chain patched. + + Returns ``(port, finish_calls)`` on success. ``refs`` (optional dict) is + populated BEFORE execution so callers can inspect state even when + execute_attempt raises. + """ + finish_calls: list = [] + port_ref: dict = {} + make_port = _make_port_factory(fake_info, port_ref, finish_calls.append, port_hook=port_hook) + + async def prepare_mock(**kwargs): + return fake_info, None + + if refs is not None: + refs["port_ref"] = port_ref + refs["finish_calls"] = finish_calls + + with _patched_application(make_port, fake_stream, prepare_mock) as application: + await application.execute_attempt(*_execute_attempt_args()) + + return port_ref["p"], finish_calls + + +async def test_leftover_chunks_flush_in_finally_before_failed_finish(): + """Chunks buffered when the loop errors mid-stream must be persisted by the + finally-block leftover flush BEFORE finish("failed") writes the status row. + """ + from services.human_interaction import application + + info = _FakeInfo() + + async def fake_stream(**_): + yield "chunk-0" + yield "chunk-1" + yield "chunk-2" + + def fail_third_add(port): + original_add = port.add_chunk + seen = {"n": 0} + + def add_chunk(chunk): + seen["n"] += 1 + if seen["n"] == 3: + raise RuntimeError("mid-loop failure") + original_add(chunk) + + port.add_chunk = add_chunk + + refs: dict = {} + with pytest.raises(RuntimeError, match="mid-loop failure"): + await _run_execute_attempt(fake_stream, info, port_hook=fail_third_add, refs=refs) + + port = refs["port_ref"]["p"] + # The two buffered chunks were flushed by the finally block, before finish. + assert port._emits == [["chunk-0", "chunk-1"]], port._emits + assert refs["finish_calls"] == ["failed"] + assert not port._emit_in_flight.is_set(), "begin_emit without end_emit leaked" + + +async def test_cancelled_error_cancels_scope_and_reraises_without_finish(): + """Cancellation propagates: the run scope is cancelled and no terminal + status row is written — the run stays resumable. + """ + from services.human_interaction import application + + info = _FakeInfo() + info.cancellation_scope = MagicMock() + + async def fake_stream(**_): + raise asyncio.CancelledError() + yield # pragma: no cover + + refs: dict = {} + with pytest.raises(asyncio.CancelledError): + await _run_execute_attempt(fake_stream, info, refs=refs) + + info.cancellation_scope.cancel.assert_called_once() + assert refs["finish_calls"] == [] + + +async def test_cancelled_error_falls_back_to_stop_event_when_scope_missing(): + """Without a cancellation scope, the stop event is the cancellation signal.""" + import threading + + info = _FakeInfo() + info.cancellation_scope = None + info.stop_event = threading.Event() + + async def fake_stream(**_): + raise asyncio.CancelledError() + yield # pragma: no cover + + refs: dict = {} + with pytest.raises(asyncio.CancelledError): + await _run_execute_attempt(fake_stream, info, refs=refs) + + assert info.stop_event.is_set() + assert refs["finish_calls"] == [] + + +async def test_run_terminated_finishes_stopped_and_swallows_finish_race(): + """RunTerminated writes the "stopped" terminal row; a RunTerminated raised + by the terminal write itself (status raced the stop) is swallowed so no + exception escapes execute_attempt. + """ + from services.human_interaction import application + + info = _FakeInfo() + run_terminated = application.RunTerminated + + async def fake_stream(**_): + raise run_terminated("terminated") + yield # pragma: no cover + + refs: dict = {} + + def racy_finish(port): + original_finish = port.finish + + def finish(outcome): + original_finish(outcome) # records into refs["finish_calls"] + raise run_terminated("terminal write raced the stop") + + port.finish = finish + + port, finish_calls = await _run_execute_attempt( + fake_stream, info, port_hook=racy_finish, refs=refs) + assert refs["finish_calls"] == ["stopped"] + assert finish_calls == ["stopped"] + + +async def test_run_terminated_finishes_stopped(): + """RunTerminated → finish("stopped") and no exception escapes.""" + from services.human_interaction import application + + info = _FakeInfo() + + async def fake_stream(**_): + raise application.RunTerminated("terminated") + yield # pragma: no cover + + refs: dict = {} + port, finish_calls = await _run_execute_attempt(fake_stream, info, refs=refs) + assert finish_calls == ["stopped"] + + +async def test_recovery_required_finishes_with_recovery_outcome(): + """RecoveryRequired → finish("recovery_required") so the run can be retried.""" + from services.human_interaction import application + + info = _FakeInfo() + + async def fake_stream(**_): + raise application.RecoveryRequired("lease lost") + yield # pragma: no cover + + port, finish_calls = await _run_execute_attempt(fake_stream, info) + assert finish_calls == ["recovery_required"] + + +async def test_finally_survives_chunk_iterator_aclose_failure(): + """A raising chunk_iter.aclose() in the finally block is swallowed and + must not prevent the leftover flush or the failed-finish from running. + """ + from services.human_interaction import application + + info = _FakeInfo() + + class _AcloseRaises: + """Async iterator that yields one chunk then fails on aclose().""" + + def __init__(self): + self._inner = self._gen() + + async def _gen(self): + yield "chunk-0" + yield "chunk-1" + + def __aiter__(self): + return self + + async def __anext__(self): + return await self._inner.__anext__() + + async def aclose(self): + raise RuntimeError("aclose failed") + + def fail_second_add(port): + original_add = port.add_chunk + seen = {"n": 0} + + def add_chunk(chunk): + seen["n"] += 1 + if seen["n"] == 2: + raise RuntimeError("mid-loop failure") + original_add(chunk) + + port.add_chunk = add_chunk + + def fake_stream(**_): + return _AcloseRaises() + + refs: dict = {} + with pytest.raises(RuntimeError, match="mid-loop failure"): + await _run_execute_attempt(fake_stream, info, port_hook=fail_second_add, refs=refs) + + port = refs["port_ref"]["p"] + # The buffered chunk was still flushed by the finally block despite the + # aclose() failure right before it. + assert port._emits == [["chunk-0"]] + assert refs["finish_calls"] == ["failed"] diff --git a/test/backend/services/test_human_interaction_snapshot.py b/test/backend/services/test_human_interaction_snapshot.py new file mode 100644 index 000000000..d3f36ea9a --- /dev/null +++ b/test/backend/services/test_human_interaction_snapshot.py @@ -0,0 +1,150 @@ +"""Unit tests for the read-path snapshot split (snapshot / light_snapshot) +and the repository read_only context manager. + +The Postgres-backed suite in test_human_interaction.py skips without a real +database, so these mocks keep the read-path lines covered in CI. No real +database is required. +""" + +from contextlib import contextmanager +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from database.human_interaction_db import HumanInteractionRepository, RunTransaction +from services.human_interaction.models import InteractionError +from services.human_interaction.service import HumanInteractionService + + +def _future_lock(): + return datetime.now(timezone.utc) + timedelta(hours=1) + + +def _make_run(**overrides): + run = SimpleNamespace( + run_id="run-1", + conversation_id=7, + status="WAITING_HUMAN", + event_seq=4, + pause_requested=True, + lock_until=_future_lock(), + ) + for key, value in overrides.items(): + setattr(run, key, value) + return run + + +def _yielding_ctx(value): + """Fresh single-use context manager yielding ``value`` (exception-safe).""" + + @contextmanager + def ctx(): + yield value + + return ctx() + + +def _make_service(tx): + repo = MagicMock() + repo.read_only.side_effect = lambda *a, **k: _yielding_ctx(tx) + repo.transaction.side_effect = lambda *a, **k: _yielding_ctx(tx) + service = HumanInteractionService(repository=repo, cipher=MagicMock(), wait_seconds=60) + return service, repo + + +def _snapshot_tx(run): + tx = MagicMock() + tx.run = run + tx.requests.return_value = [] + return tx + + +# --- service.light_snapshot --------------------------------------------------- + +def test_light_snapshot_builds_snapshot_without_writer_path(): + run = _make_run() + tx = _snapshot_tx(run) + service, repo = _make_service(tx) + + snap = service.light_snapshot("run-1", "tenant-1", "user-1") + + assert snap["run_id"] == "run-1" + assert snap["conversation_id"] == 7 + assert snap["status"] == "WAITING_HUMAN" + assert snap["event_seq"] == 4 + assert snap["pause_requested"] is True + assert snap["attempt_active"] is True + # Read path only: no write transaction, no expiration side effects. + repo.read_only.assert_called_once_with("run-1", "tenant-1", "user-1") + repo.transaction.assert_not_called() + + +def test_light_snapshot_missing_run_raises_404(): + service, repo = _make_service(None) + + with pytest.raises(InteractionError, match="not found"): + service.light_snapshot("run-1", "tenant-1", "user-1") + + +# --- service.snapshot (writer path) ------------------------------------------- + +def test_snapshot_takes_write_transaction_and_expires(): + run = _make_run() + tx = _snapshot_tx(run) + service, repo = _make_service(tx) + + with patch.object(service, "_expire") as expire: + snap = service.snapshot("run-1", "tenant-1", "user-1") + + assert snap["run_id"] == "run-1" + expire.assert_called_once() + repo.transaction.assert_called_once_with("run-1", "tenant-1", "user-1") + repo.read_only.assert_not_called() + + +# --- repository.read_only ------------------------------------------------------ + +def _repo_with_session(session): + factory = MagicMock(side_effect=lambda: _yielding_ctx(session)) + return HumanInteractionRepository(session_factory=factory), factory + + +def test_read_only_yields_transaction_without_lock_or_writes(): + run = _make_run() + session = MagicMock() + session.scalar.return_value = run + repo, factory = _repo_with_session(session) + + with repo.read_only("run-1", "tenant-1", "user-1") as tx: + assert isinstance(tx, RunTransaction) + assert tx.run is run + assert tx.actor == "user-1" + + # Plain SELECT: no FOR UPDATE lock row, no flush, no commit. + session.flush.assert_not_called() + session.commit.assert_not_called() + factory.assert_called_once_with() + + +def test_read_only_missing_run_yields_none(): + session = MagicMock() + session.scalar.return_value = None + repo, _ = _repo_with_session(session) + + with repo.read_only("run-1") as tx: + assert tx is None + + +def test_read_only_works_with_and_without_tenant_scope(): + run = _make_run() + session = MagicMock() + session.scalar.return_value = run + repo, _ = _repo_with_session(session) + + with repo.read_only("run-1", "tenant-1", "user-1") as scoped: + assert scoped is not None + with repo.read_only("run-1") as unscoped: + assert unscoped is not None + assert session.scalar.call_count == 2 diff --git a/test/backend/services/test_human_interaction_stream.py b/test/backend/services/test_human_interaction_stream.py index 07ae283e9..b0515e35c 100644 --- a/test/backend/services/test_human_interaction_stream.py +++ b/test/backend/services/test_human_interaction_stream.py @@ -87,7 +87,7 @@ async def test_durable_event_ids_allow_replay_without_skipping_snapshot_backlog( snapshot = {"run_id": "run", "conversation_id": 7, "status": "COMPLETED", "event_seq": 3} service = MagicMock() - service.snapshot.return_value = snapshot + service.light_snapshot.return_value = snapshot service.repository.events.return_value = [ {"seq": 2, "payload": {"type": "human_decision", "content": {"status": "DECIDED"}}}, {"seq": 3, "payload": {"chunk_cipher": "encrypted"}}, @@ -102,7 +102,7 @@ async def test_durable_event_ids_allow_replay_without_skipping_snapshot_backlog( assert chunks[2] == 'id: 3\ndata: {"type":"final_answer","content":"done"}\n\n' assert chunks[-1].startswith("data: ") service.repository.events.assert_called_once_with("run", 1) - assert all(call.args == ("run", "tenant", "owner") for call in service.snapshot.call_args_list) + assert all(call.args == ("run", "tenant", "owner") for call in service.light_snapshot.call_args_list) @pytest.mark.asyncio @@ -112,7 +112,7 @@ async def test_future_event_cursor_is_rejected_before_streaming(monkeypatch): from services.human_interaction.models import InteractionError service = MagicMock() - service.snapshot.return_value = {"event_seq": 3} + service.light_snapshot.return_value = {"event_seq": 3} monkeypatch.setattr(application, "require_enabled", lambda: service) with pytest.raises(InteractionError) as error: await application.stream_run("run", "tenant", "owner", after=4) diff --git a/test/backend/services/test_runtime_port_chunk_buffer.py b/test/backend/services/test_runtime_port_chunk_buffer.py new file mode 100644 index 000000000..e7eb40f08 --- /dev/null +++ b/test/backend/services/test_runtime_port_chunk_buffer.py @@ -0,0 +1,356 @@ +"""Unit tests for RuntimeInteractionPort shared chunk buffer and idle flush. + +These tests do NOT require a real database — collaborators are mocked. +They cover the thread-safe add_chunk / take_chunks buffer, the +flush_chunks_until_idle poll loop, and verify that every HITL entry point +(dispatch / boundary / receipt / finish / _wait_until_ready) invokes the +idle flush before opening its own transaction. +""" + +import threading +import time +import types +from contextlib import contextmanager +from unittest.mock import MagicMock, patch + +import pytest + + +def _make_port(): + """Build a RuntimeInteractionPort with every DB-dependent collaborator mocked. + + We patch ``transaction`` as a context manager yielding a fake tx, so the + real SQLAlchemy path is never hit. The buffer + idle-flush methods under + test operate entirely on shared state and mock ``emit_chunks``. + """ + from services.human_interaction.runtime_port import RuntimeInteractionPort + + service = MagicMock() + service.repository = MagicMock() + service.cipher = MagicMock() + service.cipher.open.return_value = {} + service.cipher.seal.return_value = "sealed" + + identity = { + "run_id": "run-1", "tenant_id": "tenant-1", "user_id": "user-1", + "fence": "fence-1", + } + + def authorize(): + return None + + with patch.object(RuntimeInteractionPort, "transaction") as mock_tx: + fake_tx = MagicMock() + fake_tx.run.checkpoint = "{}" + fake_tx.run.request_payload = "{}" + ctx = contextmanager(lambda: iter([fake_tx]))() + mock_tx.return_value = ctx + port = RuntimeInteractionPort( + service, identity, "worker", authorize, allowed_tools=("tool_a",), + live_resume=True, + ) + + # Replace the transaction contextmanager with one we control so tests + # that exercise dispatch/boundary/etc. stay DB-free. + port.transaction = MagicMock() + port.transaction.return_value = ctx + port.emit_chunks = MagicMock() + port.authorize = MagicMock() + + return port + + +# --- add_chunk / take_chunks ------------------------------------------------ + +def test_add_chunk_appends_to_buffer(): + port = _make_port() + port.add_chunk("chunk-1") + port.add_chunk("chunk-2") + assert port.take_chunks() == ["chunk-1", "chunk-2"] + + +def test_take_chunks_clears_buffer(): + port = _make_port() + port.add_chunk("x") + assert port.take_chunks() == ["x"] + assert port.take_chunks() == [] + + +def test_add_chunk_take_chunks_thread_safety(): + """Multiple producers must not lose chunks or see torn buffers.""" + port = _make_port() + n_threads = 8 + per_thread = 50 + barrier = threading.Barrier(n_threads) + + def producer(tid): + barrier.wait() + for i in range(per_thread): + port.add_chunk(f"t{tid}-{i}") + + threads = [threading.Thread(target=producer, args=(t,)) for t in range(n_threads)] + for t in threads: + t.start() + for t in threads: + t.join() + + chunks = port.take_chunks() + assert len(chunks) == n_threads * per_thread + + +# --- flush_chunks_until_idle ------------------------------------------------ + +def test_flush_until_idle_empty_buffer_settles_immediately(): + """If the buffer is already empty, two short polls + settle_ms break.""" + port = _make_port() + # settle_ms=20, so first poll sees empty → idle_since set; second poll + # 20ms later sees empty → break. Total sleep = settle_ms. + port.flush_chunks_until_idle(max_wait_ms=100, settle_ms=20) + port.emit_chunks.assert_not_called() + + +def test_flush_until_idle_persists_then_waits_for_settle(): + """Chunks present → emit_chunks called; then buffer stays empty until settle.""" + port = _make_port() + port.add_chunk("a") + port.add_chunk("b") + + # Use a tight settle to keep the test fast. + port.flush_chunks_until_idle(max_wait_ms=200, settle_ms=10) + + port.emit_chunks.assert_called_once_with(["a", "b"]) + # Buffer should now be drained (take_chunks was called inside flush). + assert port.take_chunks() == [] + + +def test_flush_until_idle_hard_deadline_never_resets(): + """max_wait_ms is a hard absolute cap measured from function entry. + + Even if new chunks keep arriving, the worker must break within + max_wait_ms. Before the fix the deadline was re-set after every + drained chunk — that could keep the worker spinning indefinitely. + """ + port = _make_port() + port.add_chunk("initial") + + # monotonic timeline: 0.0 (start), 0.010, 0.020, 0.030 (past 30ms cap), 0.030 + fake_monotonic = iter([0.0, 0.010, 0.020, 0.030, 0.030]) + + # emit_chunks keeps re-seeding the buffer so it never empties — + # without a hard deadline the loop would never break. + def fake_emit(chunks): + port.add_chunk("still-more") + port.emit_chunks.side_effect = fake_emit + + with patch("services.human_interaction.runtime_port.time.monotonic", side_effect=lambda: next(fake_monotonic)), \ + patch("services.human_interaction.runtime_port.time.sleep") as mock_sleep: + # Hard cap 30ms, settle 500ms. settle is unreachable because the + # buffer is never empty — we rely purely on the hard_deadline. + port.flush_chunks_until_idle(max_wait_ms=30, settle_ms=500) + + # emit_chunks was called a few times but NOT infinitely — the hard + # deadline cut it off. + assert port.emit_chunks.call_count >= 1 + # The last sleep call should be clipped to the remaining time (≤30ms). + if mock_sleep.call_args_list: + last_sleep = mock_sleep.call_args_list[-1].args[0] + assert last_sleep <= 0.031 # ≤ 30ms with tiny float slack + + +def test_flush_until_idle_timeout_wins_over_settle(): + """When max_wait_ms expires before an idle settle, we still break. + + Mock monotonic so time appears to advance past the deadline on the + second poll — settle_ms is long but max_wait is short, so timeout wins. + """ + port = _make_port() + port.add_chunk("initial") + + # monotonic sequence: 0.0, 0.01 (still inside 30ms window), 0.05 (past deadline) + fake_monotonic = iter([0.0, 0.010, 0.050, 0.050]) + + with patch("services.human_interaction.runtime_port.time.monotonic", side_effect=lambda: next(fake_monotonic)), \ + patch("services.human_interaction.runtime_port.time.sleep"): + # max_wait=30ms, settle=500ms. First call sees chunks → emit, deadline=0.030. + # Second poll: now=0.050 > 0.030 → timeout break. + port.flush_chunks_until_idle(max_wait_ms=30, settle_ms=500) + + # emit_chunks called exactly once — we drained the initial chunk, then timeout fired. + assert port.emit_chunks.call_count == 1 + assert port.emit_chunks.call_args.args[0] == ["initial"] + + +# --- HITL entry points invoke flush_chunks_until_idle ----------------------- + +def test_dispatch_calls_flush_before_transaction(): + port = _make_port() + # Make dispatch not actually open a transaction — we only care about flush. + port.flush_chunks_until_idle = MagicMock() + + # Patch the transaction contextmanager with one that yields a mock tx + # where pause_requested is False and we never hit ask_user path. + fake_tx = MagicMock() + fake_tx.run.pause_requested = False + fake_tx.execution.side_effect = MagicMock(status="SUCCEEDED") + fake_tx.requests.return_value = [] + port.transaction = MagicMock() + port.transaction.return_value = contextmanager(lambda: iter([fake_tx]))() + + # Interaction path needs payload with questions for the early-return + # "already answered" branches — let's use no interaction so we fall + # through to the "not suspended → emit STARTED" case. + from nexent.core.human_interaction.contracts import AttemptSuspended + + # We expect either normal return or AttemptSuspended; both are fine as + # long as flush_chunks_until_idle is called before any DB work. + try: + port.dispatch(0, "tool_a", {"x": 1}) + except AttemptSuspended: + pass + except Exception: + # Any other exception from the mocked tx is acceptable — we only + # care that flush was called. + pass + + port.flush_chunks_until_idle.assert_called_once() + + +def test_finish_calls_flush_before_transaction(): + port = _make_port() + port.flush_chunks_until_idle = MagicMock() + port.transaction = MagicMock() + fake_tx = MagicMock() + fake_tx.run.status = "RUNNING" + port.transaction.return_value = contextmanager(lambda: iter([fake_tx]))() + + port.finish("COMPLETED") + port.flush_chunks_until_idle.assert_called_once() + + +def test_boundary_calls_flush_before_transaction(): + port = _make_port() + port.flush_chunks_until_idle = MagicMock() + port.transaction = MagicMock() + fake_tx = MagicMock() + fake_tx.run.pause_requested = False + port.transaction.return_value = contextmanager(lambda: iter([fake_tx]))() + + try: + port.boundary({}) + except AttemptSuspended: + pass + + port.flush_chunks_until_idle.assert_called_once() + + +def test_receipt_calls_flush_before_transaction(): + port = _make_port() + port.flush_chunks_until_idle = MagicMock() + port.transaction = MagicMock() + fake_tx = MagicMock() + fake_exec = MagicMock() + fake_exec.status = "STARTED" + fake_tx.execution.return_value = fake_exec + port.transaction.return_value = contextmanager(lambda: iter([fake_tx]))() + + port.receipt(0, {"ok": True}) + port.flush_chunks_until_idle.assert_called_once() + + +def test_wait_until_ready_calls_flush_before_status_transition(): + """_wait_until_ready → flush_chunks_until_idle before transitioning to RUNNING.""" + from contextlib import contextmanager + + port = _make_port() + port.flush_chunks_until_idle = MagicMock() + + fake_tx = MagicMock() + fake_tx.run.status = "READY" + fake_tx.run.fence = "fence-1" + fake_tx.run.lock_owner = "worker" + + @contextmanager + def fake_repo_transaction(*args, **kwargs): + yield fake_tx + + port.repository.transaction = fake_repo_transaction + port.service = MagicMock() # _wait_until_ready calls service._expire(tx) + + # Patch utcnow so that the "lock_until is None or <= utcnow" check passes. + from datetime import datetime, timezone + fake_tx.run.lock_until = datetime(2099, 1, 1, tzinfo=timezone.utc) + + with patch("services.human_interaction.runtime_port.utcnow", return_value=datetime(2025, 1, 1, tzinfo=timezone.utc)), \ + patch("services.human_interaction.runtime_port.time.sleep"): + port._wait_until_ready() + + port.flush_chunks_until_idle.assert_called_once() + assert fake_tx.run.status == "RUNNING" + + +# --- in-flight emit and failure recovery -------------------------------------- + +def test_flush_until_idle_treats_in_flight_emit_as_busy(): + """An empty buffer is NOT idle while an async emit is handing chunks to + the DB thread — flush keeps polling until the emit completes. + """ + port = _make_port() + port.emit_chunks = MagicMock() + port._emit_in_flight.set() + result: dict = {} + + def clear_soon(): + time.sleep(0.06) + port._emit_in_flight.clear() + + threading.Thread(target=clear_soon, daemon=True).start() + + def run_flush(): + started = time.monotonic() + port.flush_chunks_until_idle(max_wait_ms=2000, settle_ms=10) + result["elapsed"] = time.monotonic() - started + + worker = threading.Thread(target=run_flush, daemon=True) + worker.start() + worker.join(timeout=5) + + assert not worker.is_alive(), "flush_chunks_until_idle blocked past hard deadline" + # Without the in-flight guard the flush would settle at ~10ms on the + # empty buffer; observing >= 50ms proves it waited for the emit. + assert result["elapsed"] >= 0.05, result + assert not port._emit_in_flight.is_set() + + +def test_flush_until_idle_restores_chunks_when_emit_chunks_raises(): + """A failed DB emit must not lose drained chunks — they go back to the + shared buffer for the next flush attempt. + """ + port = _make_port() + port.emit_chunks = MagicMock(side_effect=RuntimeError("db down")) + port.add_chunk("c1") + port.add_chunk("c2") + + with pytest.raises(RuntimeError, match="db down"): + port.flush_chunks_until_idle(max_wait_ms=200, settle_ms=10) + + assert port.peek_chunks() == 2 + assert not port._emit_in_flight.is_set(), "begin_emit without end_emit leaked" + + +def test_finish_flush_failure_still_writes_terminal_status(): + """A chunk-flush failure inside finish must not prevent the terminal + status row (and its human_run event) from being written. + """ + port = _make_port() + port.flush_chunks_until_idle = MagicMock(side_effect=RuntimeError("flush failed")) + port.transaction = MagicMock() + fake_tx = MagicMock() + fake_tx.run.status = "RUNNING" + fake_tx.requests.return_value = [] + port.transaction.return_value = contextmanager(lambda: iter([fake_tx]))() + + port.finish("COMPLETED") + + port.flush_chunks_until_idle.assert_called_once() + assert fake_tx.run.status == "COMPLETED" + fake_tx.emit.assert_called_once() From 5ed736d2711e04ae38fdc4f677c217a39b6d3429 Mon Sep 17 00:00:00 2001 From: panyehong <91180085+YehongPan@users.noreply.github.com> Date: Sun, 20 Sep 2026 15:04:08 +0800 Subject: [PATCH 4/9] =?UTF-8?q?=F0=9F=90=9B=20Bugfix:=20Fixed=20an=20issue?= =?UTF-8?q?=20where=20the=20sandbox=20container=20user=20lacked=20the=20pe?= =?UTF-8?q?rmissions=20to=20create=20folders=20and=20files.=20(#3963)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- deploy/common/start-backend.sh | 7 +++++++ deploy/docker/compose/docker-compose.prod.yml | 8 ++++---- deploy/docker/compose/docker-compose.yml | 8 ++++---- sdk/nexent/core/agents/nexent_agent.py | 5 ++++- test/sdk/core/agents/test_nexent_agent.py | 6 +++++- 5 files changed, 24 insertions(+), 10 deletions(-) diff --git a/deploy/common/start-backend.sh b/deploy/common/start-backend.sh index a49d77661..0c2e4d5aa 100755 --- a/deploy/common/start-backend.sh +++ b/deploy/common/start-backend.sh @@ -2,6 +2,13 @@ set -euo pipefail +requested_umask="${UMASK:-0022}" +if [[ ! "$requested_umask" =~ ^0?[0-7]{3}$ ]]; then + printf '[start-backend] ERROR: unsupported UMASK: %s\n' "$requested_umask" >&2 + exit 1 +fi +umask "$requested_umask" + SQL_STARTUP_MODE="${NEXENT_SQL_STARTUP_MODE:-off}" if [ -z "${NEXENT_SQL_STARTUP_MODE+x}" ] && [ -n "${NEXENT_RUN_SQL_MIGRATIONS:-}" ]; then diff --git a/deploy/docker/compose/docker-compose.prod.yml b/deploy/docker/compose/docker-compose.prod.yml index 7a20d6a2d..9bc9ceb0c 100644 --- a/deploy/docker/compose/docker-compose.prod.yml +++ b/deploy/docker/compose/docker-compose.prod.yml @@ -85,7 +85,7 @@ services: NEXENT_SQL_STARTUP_MODE: migrate NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env - ../../env/monitoring.env @@ -120,7 +120,7 @@ services: NEXENT_SQL_STARTUP_MODE: wait NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env - ../../env/monitoring.env @@ -152,7 +152,7 @@ services: NEXENT_SQL_STARTUP_MODE: wait NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env user: root @@ -184,7 +184,7 @@ services: NEXENT_SQL_STARTUP_MODE: wait NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env user: root diff --git a/deploy/docker/compose/docker-compose.yml b/deploy/docker/compose/docker-compose.yml index 93dbfd10a..1ee521ad6 100644 --- a/deploy/docker/compose/docker-compose.yml +++ b/deploy/docker/compose/docker-compose.yml @@ -98,7 +98,7 @@ services: NEXENT_SQL_STARTUP_MODE: migrate NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env - ../../env/monitoring.env @@ -135,7 +135,7 @@ services: NEXENT_SQL_STARTUP_MODE: wait NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env - ../../env/monitoring.env @@ -170,7 +170,7 @@ services: NEXENT_SQL_STARTUP_MODE: wait NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env user: root @@ -203,7 +203,7 @@ services: NEXENT_SQL_STARTUP_MODE: wait NEXENT_SQL_FILES_CHECKSUM: ${NEXENT_SQL_FILES_CHECKSUM:-} skip_proxy: "true" - UMASK: 0022 + UMASK: "0022" env_file: - ../../env/.env user: root diff --git a/sdk/nexent/core/agents/nexent_agent.py b/sdk/nexent/core/agents/nexent_agent.py index 163c6fc20..de37fb04b 100644 --- a/sdk/nexent/core/agents/nexent_agent.py +++ b/sdk/nexent/core/agents/nexent_agent.py @@ -1380,7 +1380,7 @@ def _initialize_sandbox_workspaces(self) -> None: @staticmethod def _grant_sandbox_output_access(container: Any, workspace: Path) -> None: - """Allow the sandbox user to read and write the exact run workspace.""" + """Allow sandbox traversal of the user directory and writes in the run workspace.""" gid_result = container.exec_run(["id", "-g"]) gid_exit_code = getattr(gid_result, "exit_code", None) gid_output = getattr(gid_result, "output", b"") @@ -1394,7 +1394,10 @@ def _grant_sandbox_output_access(container: Any, workspace: Path) -> None: raise RuntimeError("Sandbox user returned an invalid group ID") workspace_dir = str(workspace) + workspace_parent_dir = str(workspace.parent) commands = ( + ["chgrp", sandbox_gid, workspace_parent_dir], + ["chmod", "g+xs", workspace_parent_dir], ["chgrp", "-R", sandbox_gid, workspace_dir], ["chmod", "-R", "g+rwX", workspace_dir], ["find", workspace_dir, "-type", "d", "-exec", "chmod", "g+s", "{}", "+"], diff --git a/test/sdk/core/agents/test_nexent_agent.py b/test/sdk/core/agents/test_nexent_agent.py index 452736a8b..80b369298 100644 --- a/test/sdk/core/agents/test_nexent_agent.py +++ b/test/sdk/core/agents/test_nexent_agent.py @@ -4805,7 +4805,7 @@ def test_non_shared_workspace_pushes_archive( else: grant.assert_not_called() - def test_grant_sandbox_output_access_uses_sandbox_group(self, tmp_path): + def test_grant_sandbox_output_access_grants_parent_traversal(self, tmp_path): workspace = tmp_path / "tenant" / "user" / "run-1" input_dir = workspace / "inputs" output_dir = workspace / "outputs" @@ -4817,12 +4817,16 @@ def test_grant_sandbox_output_access_uses_sandbox_group(self, tmp_path): MagicMock(exit_code=0, output=b""), MagicMock(exit_code=0, output=b""), MagicMock(exit_code=0, output=b""), + MagicMock(exit_code=0, output=b""), + MagicMock(exit_code=0, output=b""), ] NexentAgent._grant_sandbox_output_access(container, workspace) assert container.exec_run.call_args_list == [ call(["id", "-g"]), + call(["chgrp", "1000", str(workspace.parent)], user="0"), + call(["chmod", "g+xs", str(workspace.parent)], user="0"), call(["chgrp", "-R", "1000", str(workspace)], user="0"), call(["chmod", "-R", "g+rwX", str(workspace)], user="0"), call( From 293a42263c63a5df356417a95f3f878811be981f Mon Sep 17 00:00:00 2001 From: lijiayang619 <1170349871@qq.com> Date: Sun, 20 Sep 2026 15:21:50 +0800 Subject: [PATCH 5/9] =?UTF-8?q?Fix:=20dispatch=20ModelEngine=20provider=20?= =?UTF-8?q?listing=20to=20the=20dedicated=20ModelEngi=E2=80=A6=20(#3962)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Fix: dispatch ModelEngine provider listing to the dedicated ModelEngineProvider - get_provider_models routed every provider through the OpenAI-compatible adapter, so ModelEngine batch import failed (wrong endpoint path /open/router/v1/models, self-signed cert, custom type taxonomy, missing per-model base_url). The dedicated class existed but was never wired in. * chore: ModelEngine catalog base_url placeholder - preset public URL is wrong for private deployments, placeholder communicates the required /open/router/v1 path format --------- Co-authored-by: ljy --- backend/configs/model_catalog.json | 2 +- backend/services/model_provider_service.py | 20 +++++++++++++------ .../services/test_model_provider_service.py | 8 ++++---- 3 files changed, 19 insertions(+), 11 deletions(-) diff --git a/backend/configs/model_catalog.json b/backend/configs/model_catalog.json index 7e0527a5a..3740d86a9 100644 --- a/backend/configs/model_catalog.json +++ b/backend/configs/model_catalog.json @@ -384,7 +384,7 @@ }, "modelengine": { "display_name": "ModelEngine", - "base_url": "https://api.modelengine-ai.net/v1/", + "base_url": "https:///open/router/v1", "models": { "qwen3-8b": { "model_type": "llm", diff --git a/backend/services/model_provider_service.py b/backend/services/model_provider_service.py index 828b209fa..11e9424a7 100644 --- a/backend/services/model_provider_service.py +++ b/backend/services/model_provider_service.py @@ -29,8 +29,14 @@ async def get_provider_models(model_data: dict) -> List[dict]: """ Get model list based on provider. - All providers are queried via the standard OpenAI-compatible - GET {base_url}/models endpoint. + ModelEngine is dispatched to its dedicated provider class: its models + endpoint lives at {host}/open/router/v1/models (not the OpenAI + /models path), the response uses ModelEngine's own type taxonomy + ("chat"/"embed"/"multimodal"/...) that must be mapped to internal + types, each model must carry its host so prepare_model_dict can build + the full base_url later, and its endpoints serve self-signed + certificates (ssl=False). All other providers are queried via the + standard OpenAI-compatible GET {base_url}/models endpoint. Args: model_data: Model data containing provider information @@ -38,10 +44,12 @@ async def get_provider_models(model_data: dict) -> List[dict]: Returns: List of models from the specified provider """ - provider = OpenAICompatibleProvider() - model_list = await provider.get_models(model_data) - - return model_list + provider_key = (model_data.get("provider") or "").lower() + if provider_key == ProviderEnum.MODELENGINE.value: + client: AbstractModelProvider = ModelEngineProvider() + else: + client = OpenAICompatibleProvider() + return await client.get_models(model_data) # ============================================================================= diff --git a/test/backend/services/test_model_provider_service.py b/test/backend/services/test_model_provider_service.py index 22b6187db..f1f1fb028 100644 --- a/test/backend/services/test_model_provider_service.py +++ b/test/backend/services/test_model_provider_service.py @@ -2123,7 +2123,7 @@ async def test_prepare_model_dict_modelengine_base_url_stripping(): @pytest.mark.asyncio async def test_get_provider_models_modelengine_success(): - """ModelEngine provider models are fetched via the unified OpenAI-compatible adapter.""" + """ModelEngine provider models are fetched via the dedicated ModelEngineProvider adapter.""" model_data = {"provider": "modelengine", "model_type": "llm"} expected_models = [ @@ -2136,7 +2136,7 @@ async def test_get_provider_models_modelengine_success(): ] with mock.patch( - "backend.services.model_provider_service.OpenAICompatibleProvider" + "backend.services.model_provider_service.ModelEngineProvider" ) as mock_provider_class: mock_provider_instance = mock.AsyncMock() mock_provider_instance.get_models.return_value = expected_models @@ -2151,11 +2151,11 @@ async def test_get_provider_models_modelengine_success(): @pytest.mark.asyncio async def test_get_provider_models_modelengine_empty_result(): - """Should handle empty result from the unified adapter for ModelEngine provider.""" + """Should handle empty result from the ModelEngine adapter for ModelEngine provider.""" model_data = {"provider": "modelengine", "model_type": "embedding"} with mock.patch( - "backend.services.model_provider_service.OpenAICompatibleProvider" + "backend.services.model_provider_service.ModelEngineProvider" ) as mock_provider_class: mock_provider_instance = mock.AsyncMock() mock_provider_instance.get_models.return_value = [] From 47712528d17d47cc6b06a76cf43419634c04408c Mon Sep 17 00:00:00 2001 From: Jason Wang <56037774+JasonW404@users.noreply.github.com> Date: Sun, 20 Sep 2026 16:23:41 +0800 Subject: [PATCH 6/9] [codex] fix(agent): silently retry transient model errors (#3965) * fix(agent): retry transient model failures silently * fix(model): support OpenAI httpx2 timeout client * test(model): add deterministic OpenAI-compatible mock * fix(agent): keep stream runtime within line budget --- backend/management/services/agent/run.py | 55 +-- backend/utils/agent_stream_utils.py | 41 ++ .../chat/streaming/chatStreamHandler.tsx | 103 ++++- .../adapter/remote-chat-model-adapter.ts | 89 +++- frontend/lib/reasoningAccumulator.ts | 31 ++ frontend/tests/reasoningAccumulator.test.ts | 43 ++ sdk/nexent/core/agents/core_agent.py | 14 + sdk/nexent/core/agents/nexent_agent.py | 10 + sdk/nexent/core/agents/run_agent.py | 4 + sdk/nexent/core/model_errors.py | 92 ++++ sdk/nexent/core/models/openai_llm.py | 154 ++++++- sdk/nexent/core/models/retry.py | 57 ++- sdk/nexent/core/utils/observer.py | 61 ++- .../test_agent_model_attempt_persistence.py | 172 +++++++ test/common/openai_compatible_mock_server.py | 430 ++++++++++++++++++ .../test_openai_compatible_mock_server.py | 172 +++++++ test/sdk/core/agents/test_core_agent.py | 52 +++ test/sdk/core/agents/test_nexent_agent.py | 26 ++ test/sdk/core/agents/test_subagent_wrapper.py | 21 + .../core/models/test_model_silent_retry.py | 37 ++ test/sdk/core/models/test_openai_llm.py | 161 +++++-- .../utils/test_observer_model_attempts.py | 66 +++ 22 files changed, 1774 insertions(+), 117 deletions(-) create mode 100644 sdk/nexent/core/model_errors.py create mode 100644 test/backend/services/test_agent_model_attempt_persistence.py create mode 100644 test/common/openai_compatible_mock_server.py create mode 100644 test/common/test_openai_compatible_mock_server.py create mode 100644 test/sdk/core/models/test_model_silent_retry.py create mode 100644 test/sdk/core/utils/test_observer_model_attempts.py diff --git a/backend/management/services/agent/run.py b/backend/management/services/agent/run.py index 6d7fdd48f..9cbe6b7fb 100644 --- a/backend/management/services/agent/run.py +++ b/backend/management/services/agent/run.py @@ -101,7 +101,10 @@ from utils.agent_stream_utils import ( enrich_file_uploads_with_presigned_urls as _enrich_file_uploads_with_presigned_urls, extract_json_objects_from_text as _extract_json_objects_from_text, + finalize_buffered_unit_fragments as _finalize_buffered_unit_fragments, + is_stream_unit_continuation as _is_continuation, process_skill_file_uploads as _process_skill_file_uploads, + rollback_model_attempt_units as _rollback_model_attempt_units, safe_agent_stream_error_chunk as _safe_agent_stream_error_chunk, serialize_stream_unit_content as _serialize_stream_unit_content, transform_skill_files_to_standard_format as _transform_skill_files_to_standard_format, @@ -129,19 +132,6 @@ _fa_extraction_tasks: set[asyncio.Task[None]] = set() -def _finalize_buffered_unit_fragments(message_units: list[dict[str, Any]]) -> int: - """Join mergeable unit fragments once and return finalized UTF-8 bytes.""" - finalized_bytes = 0 - for unit in message_units: - fragments = unit.pop("_content_fragments", None) - if fragments is not None: - content = "".join(fragments) - unit["content"] = content - unit["unit_content"] = content - finalized_bytes += len(str(unit.get("unit_content", "")).encode("utf-8")) - return finalized_bytes - - def _unregister_agent_run_after_execution( conversation_id: int | str, user_id: str, @@ -511,17 +501,11 @@ async def _iter_run_chunks(): chunk_type = data.get("type") chunk_content = data.get("content", "") or "" - # Add unit_index to the chunk data for frontend resume skip logic. - # This allows frontend to accurately skip chunks that were already persisted. - # For mergeable types (continuing chunks), use the current unit's index. - # For new units, use the next_unit_index that will be assigned. + # Use the current unit index for continuations and the next index + # for new units so the frontend can skip persisted resume chunks. if streaming_message_id is not None and chunk_type: mergeable = chunk_type in _MERGEABLE_TYPES - if ( - current_unit is not None - and mergeable - and current_unit.get("type") == chunk_type - ): + if _is_continuation(current_unit, mergeable, chunk_type, data): # Continuing chunk - use current unit's index data["unit_index"] = current_unit["unit_index"] elif chunk_type not in ("search_content_placeholder",): @@ -540,6 +524,16 @@ async def _iter_run_chunks(): yield f"data: {chunk}\n\n" continue + if chunk_type == "model_attempt_control": + phase = data.get("phase") + attempt_id = data.get("attempt_id") + current_unit = None + if phase == "rollback" and isinstance(attempt_id, str): + _rollback_model_attempt_units(buffered_units, attempt_id) + await channel.publish(f"data: {chunk}\n\n") + yield f"data: {chunk}\n\n" + continue + if chunk_type == ProcessType.SKILL_ARTIFACT.value: artifact_content = data.get("content") if isinstance(artifact_content, str): @@ -633,10 +627,8 @@ async def _iter_run_chunks(): # stream reaches a terminal state. if streaming_message_id is not None and chunk_type: mergeable = chunk_type in _MERGEABLE_TYPES - is_continuation = ( - current_unit is not None - and mergeable - and current_unit.get("type") == chunk_type + is_continuation = _is_continuation( + current_unit, mergeable, chunk_type, data ) if is_continuation: @@ -763,6 +755,7 @@ async def _iter_run_chunks(): "unit_content": persisted_content, "tool_call_id": data.get("tool_call_id"), "invocation_id": data.get("invocation_id"), + "_attempt_id": data.get("attempt_id"), "mergeable": mergeable, } if mergeable: @@ -815,9 +808,17 @@ async def _iter_run_chunks(): else "failed" ) outcome = getattr(agent_run_info, "attempt_outcome", None) - if getattr(agent_run_info, "human_interaction", None) is not None and isinstance(outcome, str): + if ( + getattr(agent_run_info, "human_interaction", None) is not None + and isinstance(outcome, str) + ): terminal_status = outcome if stream_completed_normally else "recovery_required" agent_run_info.attempt_outcome = terminal_status + elif outcome in {"failed", "stopped"}: + # A typed terminal model error is delivered as a normal observer + # ``error`` chunk, so the async iterator can finish normally while + # the worker outcome still authoritatively marks the run failed. + terminal_status = outcome try: skill_file_payloads = list(captured_skill_files.values()) diff --git a/backend/utils/agent_stream_utils.py b/backend/utils/agent_stream_utils.py index 728939fc9..7e0a07483 100644 --- a/backend/utils/agent_stream_utils.py +++ b/backend/utils/agent_stream_utils.py @@ -12,6 +12,47 @@ logger = logging.getLogger(__name__) +def finalize_buffered_unit_fragments(message_units: list[dict[str, Any]]) -> int: + """Join mergeable unit fragments once and return finalized UTF-8 bytes.""" + finalized_bytes = 0 + for unit in message_units: + unit.pop("_attempt_id", None) + fragments = unit.pop("_content_fragments", None) + if fragments is not None: + content = "".join(fragments) + unit["content"] = content + unit["unit_content"] = content + finalized_bytes += len(str(unit.get("unit_content", "")).encode("utf-8")) + return finalized_bytes + + +def rollback_model_attempt_units( + message_units: list[dict[str, Any]], attempt_id: str +) -> int: + """Remove uncommitted model fragments for one physical model attempt.""" + original_count = len(message_units) + message_units[:] = [ + unit for unit in message_units if unit.get("_attempt_id") != attempt_id + ] + return original_count - len(message_units) + + +def is_stream_unit_continuation( + current_unit: dict[str, Any] | None, + mergeable: bool, + chunk_type: str, + data: dict[str, Any], +) -> bool: + """Return whether a chunk can extend the current persisted stream unit.""" + return bool( + current_unit is not None + and mergeable + and current_unit.get("type") == chunk_type + and current_unit.get("_attempt_id") == data.get("attempt_id") + and current_unit.get("invocation_id") == data.get("invocation_id") + ) + + def extract_json_objects_from_text(text: str) -> list[dict]: """Extract all JSON objects embedded in a text blob.""" if not text: diff --git a/frontend/app/[locale]/chat/streaming/chatStreamHandler.tsx b/frontend/app/[locale]/chat/streaming/chatStreamHandler.tsx index 79df8b66b..c0519e6e7 100644 --- a/frontend/app/[locale]/chat/streaming/chatStreamHandler.tsx +++ b/frontend/app/[locale]/chat/streaming/chatStreamHandler.tsx @@ -83,6 +83,8 @@ interface JsonData { last_unit_index?: number; replay_chunk_count?: number; conversation_id?: number; + attempt_id?: string; + phase?: "begin" | "rollback" | "commit"; } // Reconstruct streaming state from persisted units (for tab-switch recovery) @@ -460,6 +462,10 @@ export const handleStreamResponse = async ( let finalAnswer = ""; let lastModelOutputIndex = -1; let lastContentType: string | null = null; + const attemptBlockCheckpoints = new Map< + string, + { originalLengths: Map; createdIds: Set } + >(); if (resumeConfig) { const recovered = reconstructFromStreamingMessage( @@ -549,6 +555,42 @@ export const handleStreamResponse = async ( } } + if ( + jsonData.type === "model_attempt_control" && + jsonData.attempt_id && + jsonData.phase + ) { + if (jsonData.phase === "begin") { + attemptBlockCheckpoints.set(jsonData.attempt_id, { + originalLengths: new Map(), + createdIds: new Set(), + }); + } else if (jsonData.phase === "rollback") { + const checkpoints = attemptBlockCheckpoints.get( + jsonData.attempt_id + ); + currentStep.contents = currentStep.contents.filter((item) => { + if (checkpoints?.createdIds.has(item.id)) { + return false; + } + const originalLength = checkpoints?.originalLengths.get( + item.id + ); + if (originalLength !== undefined) { + item.content = item.content.slice(0, originalLength); + } + return true; + }); + attemptBlockCheckpoints.delete(jsonData.attempt_id); + lastModelOutputIndex = currentStep.contents.length - 1; + lastContentType = + currentStep.contents[lastModelOutputIndex]?.type ?? null; + } else { + attemptBlockCheckpoints.delete(jsonData.attempt_id); + } + continue; + } + if (jsonData.type && jsonData.content) { const messageType = jsonData.type; @@ -686,21 +728,40 @@ export const handleStreamResponse = async ( lastContentBlock && lastContentBlock.type === messageType; if (shouldAppend) { + if (jsonData.attempt_id) { + const checkpoints = attemptBlockCheckpoints.get( + jsonData.attempt_id + ); + if ( + !checkpoints?.originalLengths.has(lastContentBlock.id) + ) { + checkpoints?.originalLengths.set( + lastContentBlock.id, + lastContentBlock.content.length + ); + } + } // Same type - append to existing block lastContentBlock.content += messageContent; } else { // Different type or no existing block - create new content block // This ensures thinking and deep_thinking are shown as separate nodes + const blockId = `model-${Date.now()}-${Math.random() + .toString(36) + .substring(2, 7)}`; currentStep.contents.push({ - id: `model-${Date.now()}-${Math.random() - .toString(36) - .substring(2, 7)}`, + id: blockId, type: messageType, subType, content: messageContent, expanded: true, timestamp: Date.now(), }); + if (jsonData.attempt_id) { + attemptBlockCheckpoints + .get(jsonData.attempt_id) + ?.createdIds.add(blockId); + } lastModelOutputIndex = currentStep.contents.length - 1; } @@ -741,19 +802,37 @@ export const handleStreamResponse = async ( lastModelOutputIndex >= 0 && currentStep.contents[lastModelOutputIndex] ) { - currentStep.contents[lastModelOutputIndex].content += - processedContent; + const codeBlock = + currentStep.contents[lastModelOutputIndex]; + if (jsonData.attempt_id) { + const checkpoints = attemptBlockCheckpoints.get( + jsonData.attempt_id + ); + if (!checkpoints?.originalLengths.has(codeBlock.id)) { + checkpoints?.originalLengths.set( + codeBlock.id, + codeBlock.content.length + ); + } + } + codeBlock.content += processedContent; } else { // Create new main content block for code + const blockId = `model-code-${Date.now()}-${Math.random() + .toString(36) + .substring(2, 7)}`; currentStep.contents.push({ - id: `model-code-${Date.now()}-${Math.random() - .toString(36) - .substring(2, 7)}`, + id: blockId, type: chatConfig.messageTypes.MODEL_OUTPUT_CODE, content: processedContent, expanded: true, timestamp: Date.now(), }); + if (jsonData.attempt_id) { + attemptBlockCheckpoints + .get(jsonData.attempt_id) + ?.createdIds.add(blockId); + } lastModelOutputIndex = currentStep.contents.length - 1; } @@ -769,8 +848,9 @@ export const handleStreamResponse = async ( } // If it does not exist, add one + const blockId = `generating-code-${stepIdCounter.current}`; const newGeneratingItem = { - id: `generating-code-${stepIdCounter.current}`, + id: blockId, type: chatConfig.messageTypes.GENERATING_CODE, content: t("chatStreamHandler.callingTool"), expanded: true, @@ -779,6 +859,11 @@ export const handleStreamResponse = async ( }; currentStep.contents.push(newGeneratingItem); + if (jsonData.attempt_id) { + attemptBlockCheckpoints + .get(jsonData.attempt_id) + ?.createdIds.add(blockId); + } // Mark as code generation type lastContentType = chatConfig.contentTypes.GENERATING_CODE; diff --git a/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts b/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts index c80bf3819..fcc13ce29 100644 --- a/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts +++ b/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts @@ -65,6 +65,9 @@ interface SseChunk { // frontend can route streaming content to the matching card even when // sibling sub-agents execute in parallel. invocation_id?: string; + attempt_id?: string; + phase?: "begin" | "rollback" | "commit"; + attempt?: number; path?: string; block_id?: string; origin_type?: string; @@ -202,9 +205,7 @@ export interface Nl2aResourceCandidate { } export type Nl2aInstallationFormKind = - | "SKILL_CONFIG" - | "MCP_REMOTE" - | "MCP_CONTAINER"; + "SKILL_CONFIG" | "MCP_REMOTE" | "MCP_CONTAINER"; export interface Nl2aResourceInstallationOption { option_id: string; @@ -1462,8 +1463,7 @@ export const remoteChatModelAdapter: ChatModelAdapter = { const history = historyMessages.map((msg) => { const customMetadata = isNl2Agent ? (msg.metadata?.custom as - | { nl2agentCardAction?: Nl2AgentCardAction } - | undefined) + { nl2agentCardAction?: Nl2AgentCardAction } | undefined) : undefined; const text = customMetadata?.nl2agentCardAction ? JSON.stringify(customMetadata.nl2agentCardAction) @@ -1563,8 +1563,7 @@ export const remoteChatModelAdapter: ChatModelAdapter = { if (abortHandled) return; abortHandled = true; const abortReason = abortSignal?.reason as - | { detach?: boolean } - | undefined; + { detach?: boolean } | undefined; if (abortReason?.detach) { log.log( `[ChatModelAdapter] Local stream detached from conversation ${backendConversationId ?? "unknown"}` @@ -1592,8 +1591,7 @@ export const remoteChatModelAdapter: ChatModelAdapter = { }; let agentResponse: - | ReadableStreamDefaultReader - | { type: "json"; data: unknown }; + ReadableStreamDefaultReader | { type: "json"; data: unknown }; let returnedRuntimeMetadataVersion: number | undefined; try { agentResponse = await conversationService.runAgent( @@ -1876,6 +1874,70 @@ export const remoteChatModelAdapter: ChatModelAdapter = { return resolved; }; + type SubAgentAttemptCheckpoint = { + invocationId: string; + reasoningIdx: number | null; + textLength: number; + }; + const subAgentAttemptCheckpoints = new Map< + string, + SubAgentAttemptCheckpoint + >(); + const removeContentPart = (index: number) => { + contentParts.splice(index, 1); + for (const slot of invocationSlots.values()) { + if (slot.reasoningIdx === index) slot.reasoningIdx = null; + else if (slot.reasoningIdx !== null && slot.reasoningIdx > index) { + slot.reasoningIdx -= 1; + } + } + }; + const handleModelAttemptControl = (chunk: SseChunk): boolean => { + if ( + chunk.type !== "model_attempt_control" || + !chunk.attempt_id || + !chunk.phase + ) { + return false; + } + const top = resolveSubAgent(chunk.invocation_id); + if (!top) { + if (chunk.phase === "begin") { + parentReasoning.beginAttempt(chunk.attempt_id); + } else if (chunk.phase === "rollback") { + parentReasoning.rollbackAttempt(chunk.attempt_id); + } else { + parentReasoning.commitAttempt(chunk.attempt_id); + } + return true; + } + + if (chunk.phase === "begin") { + const idx = top.slot.reasoningIdx; + subAgentAttemptCheckpoints.set(chunk.attempt_id, { + invocationId: top.invocationId, + reasoningIdx: idx, + textLength: idx === null ? 0 : (contentParts[idx]?.text?.length ?? 0), + }); + return true; + } + + const checkpoint = subAgentAttemptCheckpoints.get(chunk.attempt_id); + subAgentAttemptCheckpoints.delete(chunk.attempt_id); + if (chunk.phase !== "rollback" || !checkpoint) return true; + const slot = slotForInvocation(checkpoint.invocationId); + if (!slot || slot.reasoningIdx === null) return true; + if (checkpoint.reasoningIdx === null) { + removeContentPart(slot.reasoningIdx); + } else { + const part = contentParts[slot.reasoningIdx]; + if (part?.type === "reasoning") { + part.text = part.text.slice(0, checkpoint.textLength); + } + } + return true; + }; + const flushOpenReasoning = (specificInvocationId?: string | null) => { if (specificInvocationId) { const entry = activeSubAgents.get(specificInvocationId); @@ -2051,6 +2113,11 @@ export const remoteChatModelAdapter: ChatModelAdapter = { const chunk = parseSseChunk(line); if (!chunk) continue; + if (handleModelAttemptControl(chunk)) { + yield buildStreamResult(contentParts); + continue; + } + if (chunk.type === "human_run") { const value = typeof chunk.content === "string" @@ -2665,7 +2732,9 @@ export const remoteChatModelAdapter: ChatModelAdapter = { flushOpenReasoning(); } const partType = - chunk.type === "step_count" ? "reasoning" : mapChunkType(chunk.type); + chunk.type === "step_count" + ? "reasoning" + : mapChunkType(chunk.type); if (chunk.type === "parse") { flushOpenReasoning(chunk.invocation_id); if (chunk.content.trim()) { diff --git a/frontend/lib/reasoningAccumulator.ts b/frontend/lib/reasoningAccumulator.ts index 95c80f893..16de12103 100644 --- a/frontend/lib/reasoningAccumulator.ts +++ b/frontend/lib/reasoningAccumulator.ts @@ -7,6 +7,7 @@ interface ReasoningPart { /** Keep one reasoning block in its original position until a model/tool boundary. */ export function createReasoningAccumulator(parts: unknown[]) { let current: ReasoningPart | null = null; + const attempts = new Map(); const replace = (next: ReasoningPart) => { const index = current ? parts.indexOf(current) : -1; @@ -29,5 +30,35 @@ export function createReasoningAccumulator(parts: unknown[]) { replace({ ...current, status: { type: "done" } }); current = null; }, + beginAttempt(attemptId: string) { + attempts.set(attemptId, { + hadCurrent: current !== null, + text: current?.text ?? "", + }); + }, + rollbackAttempt(attemptId: string) { + const checkpoint = attempts.get(attemptId); + if (!checkpoint) return; + attempts.delete(attemptId); + if (checkpoint.hadCurrent) { + if (current) { + const index = parts.indexOf(current); + if (index < 0) return; + const restored = { + ...current, + text: checkpoint.text, + }; + parts[index] = restored; + current = restored; + } + } else if (current) { + const index = parts.indexOf(current); + if (index >= 0) parts.splice(index, 1); + current = null; + } + }, + commitAttempt(attemptId: string) { + attempts.delete(attemptId); + }, }; } diff --git a/frontend/tests/reasoningAccumulator.test.ts b/frontend/tests/reasoningAccumulator.test.ts index 6ed3a1a36..f36f67fb5 100644 --- a/frontend/tests/reasoningAccumulator.test.ts +++ b/frontend/tests/reasoningAccumulator.test.ts @@ -90,6 +90,49 @@ test("empty reasoning cannot move guidance received before the model starts", () ); }); +test("CMSR-003 rollback removes only the failed model attempt", () => { + const parts: unknown[] = []; + const reasoning = createReasoningAccumulator(parts); + reasoning.append("stable prefix"); + reasoning.beginAttempt("attempt-one"); + reasoning.append(" leaked partial"); + reasoning.rollbackAttempt("attempt-one"); + + assert.deepEqual(parts, [ + { + type: "reasoning", + text: "stable prefix", + status: { type: "running" }, + }, + ]); + + reasoning.beginAttempt("attempt-two"); + reasoning.append(" recovered"); + reasoning.commitAttempt("attempt-two"); + assert.equal((parts[0] as { text: string }).text, "stable prefix recovered"); +}); + +test("CMSR-003 rollback preserves interleaved sibling output", () => { + const parts: unknown[] = []; + const reasoning = createReasoningAccumulator(parts); + reasoning.append("parent prefix"); + reasoning.beginAttempt("parent-attempt"); + reasoning.append(" leaked parent token"); + + const sibling = { type: "reasoning", text: "sibling output" }; + parts.unshift(sibling); + reasoning.rollbackAttempt("parent-attempt"); + + assert.deepEqual(parts, [ + sibling, + { + type: "reasoning", + text: "parent prefix", + status: { type: "running" }, + }, + ]); +}); + test("a new model step starts below guidance even when no tool ran", () => { const parts: unknown[] = []; const reasoning = createReasoningAccumulator(parts); diff --git a/sdk/nexent/core/agents/core_agent.py b/sdk/nexent/core/agents/core_agent.py index acb9e1653..086be61be 100644 --- a/sdk/nexent/core/agents/core_agent.py +++ b/sdk/nexent/core/agents/core_agent.py @@ -25,6 +25,7 @@ from ...monitor import get_monitoring_manager +from ..model_errors import ModelInvocationTerminalError from ..utils.observer import MessageObserver, ProcessType from jinja2 import Template, StrictUndefined @@ -910,7 +911,11 @@ def rebuild_after_provider_overflow(): self.logger.log_markdown( content=model_output, title="MODEL OUTPUT", level=LogLevel.INFO) + except ModelInvocationTerminalError: + raise except Exception as e: + if self.stop_event.is_set(): + raise RunTerminated() from e raise AgentGenerationError( f"Error in generating model output:\n{e}", self.logger) from e @@ -1492,6 +1497,13 @@ def _run_stream( except StepSteered: interrupted = True continue + except ModelInvocationTerminalError: + # The model adapter has already exhausted its complete physical + # call budget (or classified the failure as non-retryable). + # Do not persist this incomplete step or let the ReAct loop + # turn it into a subsequent model invocation. + interrupted = True + raise except (AttemptSuspended, RecoveryRequired, RunTerminated): interrupted = True raise @@ -1703,6 +1715,8 @@ def rebuild_final_after_provider_overflow(): total_input_tokens = chat_message.token_usage.input_tokens total_output_tokens = chat_message.token_usage.output_tokens + except ModelInvocationTerminalError: + raise except Exception as e: # Fallback to error message if streaming fails model_output = f"Error in generating final LLM output: {e}" diff --git a/sdk/nexent/core/agents/nexent_agent.py b/sdk/nexent/core/agents/nexent_agent.py index de37fb04b..77746eef0 100644 --- a/sdk/nexent/core/agents/nexent_agent.py +++ b/sdk/nexent/core/agents/nexent_agent.py @@ -21,6 +21,7 @@ from ...monitor import AgentRunMetadata, get_agent_monitoring_context, get_monitoring_manager from ..models.openai_llm import OpenAIModel +from ..model_errors import ModelInvocationTerminalError from ..tools import * # Used for tool creation, do not delete!!! from ..utils.constants import THINK_PREFIX_PATTERN, THINK_TAG_PATTERN from ..utils.observer import MessageObserver, ProcessType @@ -1177,6 +1178,15 @@ def agent_run_with_observer( if self.agent.stop_event.is_set(): observer.add_message(self.agent.agent_name, ProcessType.WARNING, "Agent execution interrupted by external stop signal") + except ModelInvocationTerminalError as e: + observer.add_message( + agent_name=self.agent.agent_name, + process_type=ProcessType.ERROR, + content=e.safe_message(getattr(observer, "lang", "en")), + error_code=e.error_code.value, + retryable=False, + ) + raise except Exception as e: observer.add_message(agent_name=self.agent.agent_name, process_type=ProcessType.ERROR, content=f"Error in interaction: {str(e)}") diff --git a/sdk/nexent/core/agents/run_agent.py b/sdk/nexent/core/agents/run_agent.py index f3d5ea1e3..2aa267028 100644 --- a/sdk/nexent/core/agents/run_agent.py +++ b/sdk/nexent/core/agents/run_agent.py @@ -17,6 +17,7 @@ shutdown_fallback_thread_manager, ) from ..human_interaction.contracts import AttemptSuspended, RecoveryRequired, RunTerminated +from ..model_errors import ModelInvocationTerminalError from .agent_model import AgentRunInfo from .managed_mcp import ManagedMCPToolCollection from .nexent_agent import NexentAgent, ProcessType, cleanup_run_workspace @@ -363,6 +364,9 @@ def agent_run_thread(agent_run_info: AgentRunInfo): "Please start a new task." ) agent_run_info.observer.add_message("", ProcessType.ERROR, message) + except ModelInvocationTerminalError: + agent_run_info.attempt_outcome = "failed" + raise except Exception as e: agent_run_info.attempt_outcome = "failed" if "Couldn't connect to the MCP server" in str(e): diff --git a/sdk/nexent/core/model_errors.py b/sdk/nexent/core/model_errors.py new file mode 100644 index 000000000..eae0b61b6 --- /dev/null +++ b/sdk/nexent/core/model_errors.py @@ -0,0 +1,92 @@ +"""Typed terminal failures shared by model adapters and Agent boundaries. + +This module deliberately lives outside ``core.models`` so Agent modules can +depend on the error contract without importing every model implementation. +""" + +from __future__ import annotations + +from enum import Enum + + +class ModelErrorCode(str, Enum): + """Stable error categories exposed by the Agent stream boundary.""" + + RATE_LIMIT_EXHAUSTED = "model_rate_limit_exhausted" + SERVICE_UNAVAILABLE = "model_service_unavailable" + TIMEOUT = "model_timeout" + CONNECTION_ERROR = "model_connection_error" + AUTHENTICATION_ERROR = "model_authentication_error" + NOT_FOUND = "model_not_found" + INVALID_REQUEST = "model_invalid_request" + CONTEXT_OVERFLOW = "model_context_overflow" + EMPTY_RESPONSE_EXHAUSTED = "model_empty_response_exhausted" + UNKNOWN_ERROR = "model_unknown_error" + + +_SAFE_ERROR_MESSAGES = { + ModelErrorCode.RATE_LIMIT_EXHAUSTED: { + "en": "The model is receiving too many requests. Please try again later.", + "zh": "模型请求过于频繁,请稍后重试。", + }, + ModelErrorCode.SERVICE_UNAVAILABLE: { + "en": "The model service is temporarily unavailable. Please try again later.", + "zh": "模型服务暂时不可用,请稍后重试。", + }, + ModelErrorCode.TIMEOUT: { + "en": "The model request timed out. Please try again later.", + "zh": "模型请求超时,请稍后重试。", + }, + ModelErrorCode.CONNECTION_ERROR: { + "en": "The model service connection was interrupted. Please try again later.", + "zh": "模型服务连接中断,请稍后重试。", + }, + ModelErrorCode.AUTHENTICATION_ERROR: { + "en": "The model service credentials are invalid or unauthorized.", + "zh": "模型服务凭据无效或未获授权。", + }, + ModelErrorCode.NOT_FOUND: { + "en": "The selected model or model endpoint was not found.", + "zh": "未找到所选模型或模型服务地址。", + }, + ModelErrorCode.INVALID_REQUEST: { + "en": "The model request is invalid and cannot be retried.", + "zh": "模型请求无效,无法重试。", + }, + ModelErrorCode.CONTEXT_OVERFLOW: { + "en": "The conversation is too long for the selected model.", + "zh": "当前会话内容超出所选模型的上下文限制。", + }, + ModelErrorCode.EMPTY_RESPONSE_EXHAUSTED: { + "en": "The model repeatedly returned an empty response. Please try again later.", + "zh": "模型连续返回空响应,请稍后重试。", + }, + ModelErrorCode.UNKNOWN_ERROR: { + "en": "The model request failed. Please try again later.", + "zh": "模型请求失败,请稍后重试。", + }, +} + + +class ModelInvocationTerminalError(RuntimeError): + """Terminal model failure that must not be repaired by another Agent step.""" + + def __init__( + self, + error_code: ModelErrorCode, + attempts: int, + *, + cause: BaseException | None = None, + ) -> None: + # Keep provider detail internal for logs and tracing. UI boundaries + # must use ``safe_message`` instead of ``str``. + super().__init__(str(cause) if cause is not None else error_code.value) + self.error_code = error_code + self.attempts = attempts + self.retryable = False + if cause is not None: + self.__cause__ = cause + + def safe_message(self, lang: str = "en") -> str: + messages = _SAFE_ERROR_MESSAGES[self.error_code] + return messages.get(lang, messages["en"]) diff --git a/sdk/nexent/core/models/openai_llm.py b/sdk/nexent/core/models/openai_llm.py index d15ae7cc9..295d02c6b 100644 --- a/sdk/nexent/core/models/openai_llm.py +++ b/sdk/nexent/core/models/openai_llm.py @@ -11,9 +11,11 @@ import logging import threading import asyncio +import importlib import time import json import httpx +import uuid from typing import List, Optional, Dict, Any from openai.types.chat.chat_completion_message import ChatCompletionMessage @@ -43,6 +45,8 @@ from .retry import ( DEFAULT_MODEL_RETRY, ModelRetryConfig, + ModelErrorCode, + ModelInvocationTerminalError, classify_model_error, get_retry_after_seconds, ) @@ -74,6 +78,33 @@ class EmptyModelResponseError(RuntimeError): """Raised when a completed provider stream contains no user-visible content.""" +def _build_compatible_http_timeout( + default_http_client_type: type, + *, + connect: float, + read: float, + write: float, + pool: float, +): + """Build a timeout owned by the HTTP implementation used by OpenAI. + + OpenAI 3.x may use ``httpx2`` internally while Nexent still imports the + public ``httpx`` package for its own exception handling. Passing an + ``httpx.Timeout`` into an ``httpx2.Client`` nests incompatible timeout + objects and fails before the first provider request. Resolve the timeout + class from ``DefaultHttpxClient``'s public base class instead, falling + back to ``httpx`` for older OpenAI releases and test doubles. + """ + for base in getattr(default_http_client_type, "__mro__", ()): + module_root = getattr(base, "__module__", "").partition(".")[0] + if not module_root.startswith("httpx"): + continue + timeout_type = getattr(importlib.import_module(module_root), "Timeout", None) + if timeout_type is not None: + return timeout_type(connect=connect, read=read, write=write, pool=pool) + return httpx.Timeout(connect=connect, read=read, write=write, pool=pool) + + def _is_timeout_error(exc: BaseException) -> bool: """Return whether an exception chain represents a network or caller timeout.""" current: BaseException | None = exc @@ -188,13 +219,19 @@ def __init__(self, observer: MessageObserver = MessageObserver, temperature=0.2, # Keep every streaming HTTP phase finite. Callers can still inject a # custom client through client_kwargs when they own its lifecycle. client_kwargs = kwargs.get("client_kwargs", {}) + # The Agent retry budget counts physical provider requests. Disable + # the OpenAI client's hidden transport retries by default so one + # adapter attempt cannot fan out into multiple uncounted HTTP calls. + # A fully injected client remains under its caller's ownership, but + # clients constructed here must never exceed this adapter's budget. + client_kwargs["max_retries"] = 0 if "http_client" not in client_kwargs: from openai import DefaultHttpxClient - from openai._base_client import httpx2 http_client = DefaultHttpxClient( verify=ssl_verify, - timeout=httpx2.Timeout( + timeout=_build_compatible_http_timeout( + DefaultHttpxClient, connect=connect_timeout_seconds, read=self.read_timeout_seconds, write=write_timeout_seconds, @@ -240,6 +277,7 @@ def __call__(self, messages: List[Dict[str, Any]], stop_sequences: Optional[List response_format: dict[str, str] | None = None, tools_to_call_from: Optional[List[Tool]] = None, _token_tracker=None, context_budget_snapshot: Optional[ContextBudgetSnapshot] = None, context_rebuild=None, _overflow_recovery_ordinal: int = 0, + _model_attempts_used: int = 0, **kwargs, ) -> ChatMessage: _monitoring_operation.set("chat_completion") @@ -282,6 +320,7 @@ def __call__(self, messages: List[Dict[str, Any]], stop_sequences: Optional[List context_budget_snapshot=context_budget_snapshot, context_rebuild=context_rebuild, _overflow_recovery_ordinal=_overflow_recovery_ordinal, + _model_attempts_used=_model_attempts_used, **kwargs, ) @@ -413,13 +452,22 @@ def __call__(self, messages: List[Dict[str, Any]], stop_sequences: Optional[List } ) - for attempt in range(1, self.retry_config.max_attempts + 1): + for attempt in range(_model_attempts_used + 1, self.retry_config.max_attempts + 1): first_token_received = False if self.stop_event.is_set(): if token_tracker: self._monitoring.add_span_event("model_stopped", { "reason": "stop_event_set"}) raise RuntimeError(STOP_EVENT_INTERRUPTED_MESSAGE) + attempt_id = uuid.uuid4().hex + begin_attempt = getattr(self.observer, "begin_model_attempt", None) + if callable(begin_attempt): + begin_attempt(attempt_id, attempt) + self._monitoring.add_span_event("model_attempt_begin", { + "attempt_id": attempt_id, + "attempt": attempt, + "max_attempts": self.retry_config.max_attempts, + }) current_request = None stream_token = None close_stream_once = None @@ -673,6 +721,13 @@ def _close_stream_once(): ) message.raw = current_request message.role = MessageRole.ASSISTANT + commit_attempt = getattr(self.observer, "commit_model_attempt", None) + if callable(commit_attempt): + commit_attempt(attempt_id, attempt) + self._monitoring.add_span_event("model_attempt_commit", { + "attempt_id": attempt_id, + "attempt": attempt, + }) return message except Exception as e: @@ -681,22 +736,29 @@ def _close_stream_once(): e).__name__, "error_message": str(e)}) raise e - except EmptyModelResponseError: - # Some reasoning-capable OpenAI-compatible providers - # occasionally finish with ``stop`` after emitting only - # reasoning chunks. Retry once inside the model adapter so an - # otherwise transient malformed stream does not consume a - # visible agent step. A ``length`` finish is deterministic - # truncation and must still surface immediately. - empty_retry_limit = min(self.retry_config.max_attempts, 2) - if self.last_finish_reason not in (None, "stop") or attempt >= empty_retry_limit: - raise + except EmptyModelResponseError as empty_error: + rollback_attempt = getattr(self.observer, "rollback_model_attempt", None) + if callable(rollback_attempt): + rollback_attempt(attempt_id, attempt) + self._monitoring.add_span_event("model_attempt_rollback", { + "attempt_id": attempt_id, + "attempt": attempt, + "reason": "empty_response", + }) + # Empty ``stop`` responses share the normal model attempt + # budget. Deterministic truncation (``length``) fails fast. + if self.last_finish_reason not in (None, "stop") or attempt >= self.retry_config.max_attempts: + raise ModelInvocationTerminalError( + ModelErrorCode.EMPTY_RESPONSE_EXHAUSTED, + attempt, + cause=empty_error, + ) from empty_error backoff = self.retry_config.calculate_backoff(attempt) logger.warning( "event=retry_empty_model_response attempt=%d/%d finish_reason=%s " "retrying_after_seconds=%.2f", attempt, - empty_retry_limit, + self.retry_config.max_attempts, self.last_finish_reason, backoff, ) @@ -706,6 +768,14 @@ def _close_stream_once(): self.stop_event.wait(backoff) continue except Exception as e: + rollback_attempt = getattr(self.observer, "rollback_model_attempt", None) + if callable(rollback_attempt): + rollback_attempt(attempt_id, attempt) + self._monitoring.add_span_event("model_attempt_rollback", { + "attempt_id": attempt_id, + "attempt": attempt, + "error_type": type(e).__name__, + }) if self.stop_event.is_set() or self.cancellation_scope.cancelled: raise RuntimeError(STOP_EVENT_INTERRUPTED_MESSAGE) from e if isinstance(e, ModelConcurrencyExceeded): @@ -717,13 +787,29 @@ def _close_stream_once(): }) if is_provider_context_overflow(e): if first_token_received or context_rebuild is None: - raise ProviderContextOverflowRetryUnsafe( + overflow_error = ProviderContextOverflowRetryUnsafe( "Provider context overflow cannot be safely rebuilt: " f"{e}" + ) + raise ModelInvocationTerminalError( + ModelErrorCode.CONTEXT_OVERFLOW, + attempt, + cause=overflow_error, + ) from e + if attempt >= self.retry_config.max_attempts: + raise ModelInvocationTerminalError( + ModelErrorCode.CONTEXT_OVERFLOW, + attempt, + cause=e, ) from e if _overflow_recovery_ordinal >= 2: - raise ProviderContextOverflowRetryExhausted( + overflow_error = ProviderContextOverflowRetryExhausted( "Provider context overflow persisted after two recovery dispatches" + ) + raise ModelInvocationTerminalError( + ModelErrorCode.CONTEXT_OVERFLOW, + attempt, + cause=overflow_error, ) from e rebuilt = context_rebuild() rebuilt_messages = getattr(rebuilt, "messages", rebuilt) @@ -752,6 +838,7 @@ def _close_stream_once(): context_budget_snapshot=trusted_budget_snapshot, context_rebuild=context_rebuild, _overflow_recovery_ordinal=_overflow_recovery_ordinal + 1, + _model_attempts_used=attempt, **kwargs, ) is_timeout = _is_timeout_error(e) @@ -769,15 +856,28 @@ def _close_stream_once(): received_chunk_count, type(e).__name__, ) - if classify_model_error(e) != "retryable": - raise + classification = classify_model_error(e) + if not classification.retryable: + raise ModelInvocationTerminalError( + classification.error_code, + attempt, + cause=e, + ) from e if attempt >= self.retry_config.max_attempts: if not is_timeout: - logging.exception( - "Model call failed after %d attempts: %s", - attempt, str(e), + logger.error( + "event=model_retry_exhausted attempt=%d/%d " + "error_type=%s error_code=%s", + attempt, + self.retry_config.max_attempts, + type(e).__name__, + classification.error_code.value, ) - raise + raise ModelInvocationTerminalError( + classification.error_code, + attempt, + cause=e, + ) from e backoff = self.retry_config.calculate_backoff(attempt) retry_after = get_retry_after_seconds(e) if retry_after is not None: @@ -796,9 +896,13 @@ def _close_stream_once(): ) else: logger.warning( - "Model call attempt %d/%d failed with retryable error (%s); " - "retrying after %.2fs", - attempt, self.retry_config.max_attempts, str(e), backoff, + "event=model_retry attempt=%d/%d error_type=%s " + "error_code=%s retrying_after_seconds=%.2f", + attempt, + self.retry_config.max_attempts, + type(e).__name__, + classification.error_code.value, + backoff, ) self.last_retry_count = attempt if self.stop_event.is_set(): diff --git a/sdk/nexent/core/models/retry.py b/sdk/nexent/core/models/retry.py index 17491db54..b48b86ee2 100644 --- a/sdk/nexent/core/models/retry.py +++ b/sdk/nexent/core/models/retry.py @@ -13,8 +13,9 @@ network / connection / timeout failures). Authentication errors, not-found, invalid payloads and context-length errors are treated as non-retryable so we fail fast instead of burning backoff on hopeless requests. -* Empty responses (stream completed without user-visible content) are handled - by the caller (the agent step loop / summary truncation), **not** here. +* Empty ``stop`` responses share the same physical-call budget as provider + failures. The model adapter detects them and raises the typed terminal error + defined here after the budget is exhausted. """ from __future__ import annotations @@ -22,6 +23,8 @@ import random from dataclasses import dataclass +from ..model_errors import ModelErrorCode, ModelInvocationTerminalError + @dataclass class ModelRetryConfig: @@ -39,7 +42,7 @@ class ModelRetryConfig: across many concurrent clients. """ - max_attempts: int = 6 + max_attempts: int = 5 backoff_base_seconds: float = 2.0 max_backoff_seconds: float = 30.0 jitter: bool = True @@ -63,7 +66,13 @@ def calculate_backoff(self, attempt: int) -> float: DEFAULT_MODEL_RETRY = ModelRetryConfig() -def classify_model_error(exc: BaseException) -> str: +@dataclass(frozen=True) +class ModelErrorClassification: + retryable: bool + error_code: ModelErrorCode + + +def classify_model_error(exc: BaseException) -> ModelErrorClassification: """Classify an exception raised by a model invocation. Returns ``"retryable"`` for transient errors (rate limiting, server-side @@ -78,9 +87,17 @@ def classify_model_error(exc: BaseException) -> str: """ status = getattr(exc, "status_code", None) if isinstance(status, int): - if status == 429 or 500 <= status < 600: - return "retryable" - return "non_retryable" + if status == 429: + return ModelErrorClassification(True, ModelErrorCode.RATE_LIMIT_EXHAUSTED) + if 500 <= status < 600: + return ModelErrorClassification(True, ModelErrorCode.SERVICE_UNAVAILABLE) + if status in (401, 403): + return ModelErrorClassification(False, ModelErrorCode.AUTHENTICATION_ERROR) + if status == 404: + return ModelErrorClassification(False, ModelErrorCode.NOT_FOUND) + if status in (400, 422): + return ModelErrorClassification(False, ModelErrorCode.INVALID_REQUEST) + return ModelErrorClassification(False, ModelErrorCode.UNKNOWN_ERROR) msg = str(exc).lower() non_retryable_markers = ( @@ -89,9 +106,19 @@ def classify_model_error(exc: BaseException) -> str: "invalid", "api key", "authentication", "context_length", "context length", "token limit", ) - for marker in non_retryable_markers: - if marker in msg: - return "non_retryable" + if any(marker in msg for marker in ("context_length", "context length", "token limit")): + return ModelErrorClassification(False, ModelErrorCode.CONTEXT_OVERFLOW) + if any(marker in msg for marker in ("401", "unauthorized", "403", "forbidden", "api key", "authentication")): + return ModelErrorClassification(False, ModelErrorCode.AUTHENTICATION_ERROR) + if any(marker in msg for marker in ("404", "not found")): + return ModelErrorClassification(False, ModelErrorCode.NOT_FOUND) + if any(marker in msg for marker in non_retryable_markers): + return ModelErrorClassification(False, ModelErrorCode.INVALID_REQUEST) + + if isinstance(exc, TimeoutError) or "timeout" in type(exc).__name__.lower(): + return ModelErrorClassification(True, ModelErrorCode.TIMEOUT) + if isinstance(exc, ConnectionError): + return ModelErrorClassification(True, ModelErrorCode.CONNECTION_ERROR) retryable_markers = ( "429", "rate limit", "rate_limit", @@ -105,10 +132,16 @@ def classify_model_error(exc: BaseException) -> str: ) for marker in retryable_markers: if marker in msg: - return "retryable" + if "timeout" in marker or "timed out" in marker or "time out" in marker or marker == "etimedout": + return ModelErrorClassification(True, ModelErrorCode.TIMEOUT) + if marker in {"429", "rate limit", "rate_limit"}: + return ModelErrorClassification(True, ModelErrorCode.RATE_LIMIT_EXHAUSTED) + if marker in {"500", "502", "503", "504", "server error", "service unavailable", "temporarily unavailable", "try again", "gateway timeout", "bad gateway"}: + return ModelErrorClassification(True, ModelErrorCode.SERVICE_UNAVAILABLE) + return ModelErrorClassification(True, ModelErrorCode.CONNECTION_ERROR) # Unknown error: prefer failing fast over retrying blindly. - return "non_retryable" + return ModelErrorClassification(False, ModelErrorCode.UNKNOWN_ERROR) def get_retry_after_seconds(exc: BaseException) -> float | None: diff --git a/sdk/nexent/core/utils/observer.py b/sdk/nexent/core/utils/observer.py index 08e7962d7..0a39bde28 100644 --- a/sdk/nexent/core/utils/observer.py +++ b/sdk/nexent/core/utils/observer.py @@ -40,6 +40,7 @@ class ProcessType(Enum): MODEL_OUTPUT_THINKING = "model_output_thinking" # model streaming output, thinking content MODEL_OUTPUT_DEEP_THINKING = "model_output_deep_thinking" # model streaming output, deep thinking content MODEL_OUTPUT_CODE = "model_output_code" # model streaming output, code content + MODEL_ATTEMPT_CONTROL = "model_attempt_control" # hidden begin/rollback/commit boundary STEP_COUNT = "step_count" # current step of agent PARSE = "parse" # code parsing result @@ -177,6 +178,9 @@ def __init__(self, lang="zh", enable_nl2a_wrapper=False): self._current_invocation_id: ContextVar[str | None] = ContextVar( "current_invocation_id", default=None ) + self._model_attempt_id: ContextVar[str | None] = ContextVar( + "model_attempt_id", default=None + ) @property def token_buffer(self) -> deque: @@ -248,6 +252,7 @@ def _init_message_transformers(self): ProcessType.PLAN: default_transformer, ProcessType.PLAN_STEP_UPDATE: default_transformer, ProcessType.AUTOMATION_PROPOSAL: default_transformer, + ProcessType.MODEL_ATTEMPT_CONTROL: default_transformer, } def _active_subagent(self) -> tuple | None: @@ -274,6 +279,7 @@ def _emit( invocation_id: str | None = None, explicit_agent_id: bool = False, explicit_invocation_id: bool = False, + metadata: dict[str, Any] | None = None, ) -> None: """Append a ``Message`` with the current sub-agent context auto-stamped. @@ -305,9 +311,51 @@ def _emit( depth=resolved_depth, tool_call_id=tool_call_id, invocation_id=resolved_invocation, + attempt_id=( + self._model_attempt_id.get() + if process_type in { + ProcessType.MODEL_OUTPUT_THINKING, + ProcessType.MODEL_OUTPUT_DEEP_THINKING, + ProcessType.MODEL_OUTPUT_CODE, + } + else None + ), + metadata=metadata, ).to_json() ) + def _reset_model_stream_state(self) -> None: + self.token_buffer.clear() + self.think_buffer.clear() + self.current_mode = ProcessType.MODEL_OUTPUT_THINKING + self.in_think_mode = False + + def begin_model_attempt(self, attempt_id: str, attempt: int) -> None: + self._reset_model_stream_state() + self._model_attempt_id.set(attempt_id) + self._emit( + ProcessType.MODEL_ATTEMPT_CONTROL, + "", + metadata={"phase": "begin", "attempt_id": attempt_id, "attempt": attempt}, + ) + + def rollback_model_attempt(self, attempt_id: str, attempt: int) -> None: + self._reset_model_stream_state() + self._emit( + ProcessType.MODEL_ATTEMPT_CONTROL, + "", + metadata={"phase": "rollback", "attempt_id": attempt_id, "attempt": attempt}, + ) + self._model_attempt_id.set(None) + + def commit_model_attempt(self, attempt_id: str, attempt: int) -> None: + self._emit( + ProcessType.MODEL_ATTEMPT_CONTROL, + "", + metadata={"phase": "commit", "attempt_id": attempt_id, "attempt": attempt}, + ) + self._model_attempt_id.set(None) + def add_model_new_token(self, new_token): """ Process streaming tokens with real-time think tag detection and content classification @@ -607,6 +655,11 @@ def add_message(self, agent_name, process_type, content, **kwargs): agent_id=kwargs.get("agent_id"), agent_name=kwargs.get("agent_name"), explicit_agent_id=explicit_agent_id, + metadata={ + key: kwargs[key] + for key in ("error_code", "retryable") + if key in kwargs + }, ) @contextmanager @@ -750,7 +803,8 @@ class Message: def __init__(self, message_type: ProcessType, content, tool_name: str = None, tool_arguments: dict = None, agent_id=None, agent_name: str = None, depth: int = 0, tool_call_id: str | None = None, - invocation_id: str | None = None): + invocation_id: str | None = None, attempt_id: str | None = None, + metadata: dict[str, Any] | None = None): self.message_type = message_type self.content = content self.tool_name = tool_name @@ -760,6 +814,8 @@ def __init__(self, message_type: ProcessType, content, tool_name: str = None, self.depth = depth self.tool_call_id = tool_call_id self.invocation_id = invocation_id + self.attempt_id = attempt_id + self.metadata = metadata or {} # generate json format and convert to string def to_json(self): @@ -786,4 +842,7 @@ def to_json(self): result["depth"] = self.depth if self.invocation_id is not None: result["invocation_id"] = self.invocation_id + if self.attempt_id is not None: + result["attempt_id"] = self.attempt_id + result.update(self.metadata) return json.dumps(result, ensure_ascii=False) diff --git a/test/backend/services/test_agent_model_attempt_persistence.py b/test/backend/services/test_agent_model_attempt_persistence.py new file mode 100644 index 000000000..422d60ae3 --- /dev/null +++ b/test/backend/services/test_agent_model_attempt_persistence.py @@ -0,0 +1,172 @@ +import asyncio +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from management.services.agent import run as agent_run_service +from management.services.agent.run import ( + _finalize_buffered_unit_fragments, + _is_continuation, + _rollback_model_attempt_units, +) + + +def test_cmsr_003_rollback_removes_only_matching_attempt_units(): + units = [ + {"type": "model_output_thinking", "_attempt_id": "failed", "unit_content": "leak"}, + {"type": "model_output_thinking", "_attempt_id": "sibling", "unit_content": "keep"}, + {"type": "step_count", "unit_content": "step"}, + ] + + assert _rollback_model_attempt_units(units, "failed") == 1 + assert [unit["unit_content"] for unit in units] == ["keep", "step"] + + +def test_cmsr_003_attempt_metadata_is_never_persisted(): + units = [ + { + "type": "model_output_thinking", + "_attempt_id": "successful", + "_content_fragments": ["hello", " world"], + "unit_content": "", + } + ] + + assert _finalize_buffered_unit_fragments(units) == len("hello world") + assert units == [ + { + "type": "model_output_thinking", + "unit_content": "hello world", + "content": "hello world", + } + ] + + +def test_cmsr_003_continuation_requires_matching_attempt_and_invocation(): + current = { + "type": "model_output_thinking", + "_attempt_id": "attempt-1", + "invocation_id": "subagent-1", + } + matching = {"attempt_id": "attempt-1", "invocation_id": "subagent-1"} + sibling = {"attempt_id": "attempt-2", "invocation_id": "subagent-1"} + + assert _is_continuation(current, True, "model_output_thinking", matching) + assert not _is_continuation(current, True, "model_output_thinking", sibling) + + +def _agent_request(): + return SimpleNamespace( + conversation_id=999, + history=[], + is_debug=False, + ) + + +def _agent_run_info(outcome: str = "completed"): + return SimpleNamespace( + agent_config=SimpleNamespace(pre_run_tool_events=()), + cancellation_scope=None, + stop_event=asyncio.Event(), + human_interaction=None, + attempt_outcome=outcome, + thread_future=None, + ) + + +def _configure_stream_mocks(monkeypatch, persisted_batches): + monkeypatch.setattr(agent_run_service, "save_message", MagicMock(return_value=4242)) + monkeypatch.setattr( + agent_run_service, + "persist_assistant_run_batch", + lambda **kwargs: persisted_batches.append(kwargs), + ) + monkeypatch.setattr( + agent_run_service, "_unregister_agent_run_after_execution", MagicMock() + ) + monkeypatch.setattr( + agent_run_service.streaming_channel_manager, "complete_channel", AsyncMock() + ) + monkeypatch.setattr(agent_run_service, "_cleanup_channel_later", AsyncMock()) + + async def run_managed(_lane, _spec, fn, *args, **kwargs): + return fn(*args, **kwargs) + + async def wait_for_cancel(*_args, **_kwargs): + await asyncio.Event().wait() + + monkeypatch.setattr(agent_run_service.runtime_thread_manager, "run", run_managed) + monkeypatch.setattr(agent_run_service, "_poll_runtime_cancel_signal", wait_for_cancel) + + +@pytest.mark.asyncio +async def test_cmsr_003_stream_rolls_back_failed_attempt_before_persistence(monkeypatch): + async def fake_agent_run(*_args, **_kwargs): + yield json.dumps({"type": "model_attempt_control", "content": "", "phase": "begin", "attempt_id": "failed"}) + yield json.dumps({"type": "model_output_thinking", "content": "leak", "attempt_id": "failed"}) + yield json.dumps({"type": "model_attempt_control", "content": "", "phase": "rollback", "attempt_id": "failed"}) + yield json.dumps({"type": "model_attempt_control", "content": "", "phase": "begin", "attempt_id": "success"}) + yield json.dumps({"type": "model_output_thinking", "content": "keep", "attempt_id": "success"}) + yield json.dumps({"type": "model_attempt_control", "content": "", "phase": "commit", "attempt_id": "success"}) + yield json.dumps({"type": "final_answer", "content": "done"}) + + persisted_batches = [] + _configure_stream_mocks(monkeypatch, persisted_batches) + monkeypatch.setattr(agent_run_service, "agent_run", fake_agent_run) + channel = SimpleNamespace(publish=AsyncMock()) + + chunks = [ + chunk + async for chunk in agent_run_service._stream_agent_chunks( + _agent_request(), + "user1", + "tenant1", + _agent_run_info(), + MagicMock(), + channel=channel, + ) + ] + + assert len(chunks) == 7 + assert len(persisted_batches) == 1 + units = persisted_batches[0]["message_units"] + assert [unit["unit_content"] for unit in units] == ["keep", "done"] + assert all("_attempt_id" not in unit for unit in units) + + +@pytest.mark.asyncio +async def test_cmsr_004_terminal_error_is_persisted_once_with_failed_status(monkeypatch): + async def fake_agent_run(*_args, **_kwargs): + yield json.dumps( + { + "type": "error", + "content": "The model request failed.", + "error_code": "model_unknown_error", + "retryable": False, + } + ) + + persisted_batches = [] + _configure_stream_mocks(monkeypatch, persisted_batches) + monkeypatch.setattr(agent_run_service, "agent_run", fake_agent_run) + channel = SimpleNamespace(publish=AsyncMock()) + + chunks = [ + chunk + async for chunk in agent_run_service._stream_agent_chunks( + _agent_request(), + "user1", + "tenant1", + _agent_run_info("failed"), + MagicMock(), + channel=channel, + ) + ] + + error_chunks = [chunk for chunk in chunks if '"type": "error"' in chunk] + assert len(error_chunks) == 1 + assert '"retryable": false' in error_chunks[0] + assert persisted_batches[0]["terminal_status"] == "failed" + assert [unit["type"] for unit in persisted_batches[0]["message_units"]] == ["error"] diff --git a/test/common/openai_compatible_mock_server.py b/test/common/openai_compatible_mock_server.py new file mode 100644 index 000000000..0f24c246f --- /dev/null +++ b/test/common/openai_compatible_mock_server.py @@ -0,0 +1,430 @@ +#!/usr/bin/env python3 +"""Deterministic OpenAI-compatible server for Nexent deployment tests. + +The server is deliberately self-contained and never forwards requests to a +real provider. It implements the OpenAI surfaces used by Nexent and exposes a +small control API for selecting retry/failure scenarios. +""" + +from __future__ import annotations + +import argparse +import json +import re +import socket +import threading +import time +import uuid +from dataclasses import dataclass, field +from http import HTTPStatus +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any +from urllib.parse import urlsplit + + +SUPPORTED_SCENARIOS = { + "success", + "partial_then_success", + "always_429", + "always_503", + "always_401", + "empty_stop", + "length", +} + + +@dataclass +class MockState: + scenario: str = "success" + response_text: str = "MOCK_SUCCESS" + retry_after: float = 0.0 + partial_chunk_delay: float = 0.05 + request_count: int = 0 + requests: list[dict[str, Any]] = field(default_factory=list) + lock: threading.Lock = field(default_factory=threading.Lock, repr=False) + + def configure(self, payload: dict[str, Any]) -> dict[str, Any]: + scenario = payload.get("scenario", "success") + if scenario not in SUPPORTED_SCENARIOS: + raise ValueError(f"unsupported scenario: {scenario}") + with self.lock: + self.scenario = scenario + self.response_text = str(payload.get("response_text", "MOCK_SUCCESS")) + self.retry_after = max(0.0, float(payload.get("retry_after", 0.0))) + self.partial_chunk_delay = max( + 0.0, float(payload.get("partial_chunk_delay", 0.05)) + ) + self.request_count = 0 + self.requests.clear() + return self._snapshot_unlocked() + + def record(self, payload: dict[str, Any], has_bearer_auth: bool) -> tuple[int, str]: + with self.lock: + self.request_count += 1 + request_number = self.request_count + self.requests.append( + { + "request_number": request_number, + "model": payload.get("model"), + "stream": payload.get("stream") is True, + "include_usage": ( + isinstance(payload.get("stream_options"), dict) + and payload["stream_options"].get("include_usage") is True + ), + "message_count": len(payload.get("messages", [])), + "max_tokens": payload.get("max_tokens"), + "stop_count": len(payload.get("stop", []) or []), + "has_bearer_auth": has_bearer_auth, + } + ) + return request_number, self.scenario + + def snapshot(self) -> dict[str, Any]: + with self.lock: + return self._snapshot_unlocked() + + def _snapshot_unlocked(self) -> dict[str, Any]: + return { + "scenario": self.scenario, + "response_text": self.response_text, + "retry_after": self.retry_after, + "request_count": self.request_count, + "requests": list(self.requests), + } + + +def _extract_requested_answer(payload: dict[str, Any], fallback: str) -> str: + messages = payload.get("messages") + if not isinstance(messages, list): + return fallback + text_parts: list[str] = [] + for message in reversed(messages): + if not isinstance(message, dict) or message.get("role") != "user": + continue + content = message.get("content") + if isinstance(content, str): + text_parts.append(content) + elif isinstance(content, list): + text_parts.extend( + str(item.get("text", "")) + for item in content + if isinstance(item, dict) and item.get("type") == "text" + ) + break + user_text = "\n".join(text_parts) + match = re.search( + r"reply\s+with\s+exactly\s+([A-Za-z0-9_.:-]+)", user_text, re.IGNORECASE + ) + return match.group(1) if match else fallback + + +class OpenAICompatibleMockHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + server_version = "NexentOpenAICompatibleMock/1.0" + + @property + def state(self) -> MockState: + return self.server.mock_state # type: ignore[attr-defined] + + def log_message(self, _format: str, *_args: Any) -> None: + return + + def _send_bytes( + self, status: int, body: bytes, content_type: str = "application/json" + ) -> None: + self.send_response(status) + self.send_header("Content-Type", content_type) + self.send_header("Content-Length", str(len(body))) + self.send_header("x-request-id", f"mock-{uuid.uuid4().hex}") + self.end_headers() + self.wfile.write(body) + + def _send_json(self, status: int, payload: dict[str, Any]) -> None: + self._send_bytes(status, json.dumps(payload).encode()) + + def _read_json(self) -> dict[str, Any]: + try: + length = int(self.headers.get("Content-Length", "0")) + payload = json.loads(self.rfile.read(length) or b"{}") + except (ValueError, json.JSONDecodeError) as exc: + raise ValueError("request body must be valid JSON") from exc + if not isinstance(payload, dict): + raise ValueError("request body must be a JSON object") + return payload + + def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract + path = urlsplit(self.path).path + if path == "/__stats": + self._send_json(HTTPStatus.OK, self.state.snapshot()) + return + if path in {"/v1/models", "/models"}: + self._send_json( + HTTPStatus.OK, + { + "object": "list", + "data": [ + { + "id": "nexent-mock-model", + "object": "model", + "created": 0, + "owned_by": "nexent-test", + } + ], + }, + ) + return + self._send_error(HTTPStatus.NOT_FOUND, "not_found", "Unknown endpoint") + + def do_POST(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract + path = urlsplit(self.path).path + try: + payload = self._read_json() + except ValueError as exc: + self._send_error(HTTPStatus.BAD_REQUEST, "invalid_request_error", str(exc)) + return + + if path == "/__control": + try: + snapshot = self.state.configure(payload) + except (TypeError, ValueError) as exc: + self._send_error(HTTPStatus.BAD_REQUEST, "invalid_request_error", str(exc)) + return + self._send_json(HTTPStatus.OK, snapshot) + return + + if path not in {"/v1/chat/completions", "/chat/completions"}: + self._send_error(HTTPStatus.NOT_FOUND, "not_found", "Unknown endpoint") + return + if not isinstance(payload.get("model"), str) or not isinstance( + payload.get("messages"), list + ): + self._send_error( + HTTPStatus.BAD_REQUEST, + "invalid_request_error", + "model and messages are required", + ) + return + + has_bearer_auth = self.headers.get("Authorization", "").startswith("Bearer ") + request_number, scenario = self.state.record(payload, has_bearer_auth) + if scenario == "always_429": + self._send_error( + HTTPStatus.TOO_MANY_REQUESTS, + "rate_limit_error", + "Injected rate limit", + retry_after=self.state.retry_after, + ) + return + if scenario == "always_503": + self._send_error( + HTTPStatus.SERVICE_UNAVAILABLE, + "server_error", + "Injected temporary outage", + retry_after=self.state.retry_after, + ) + return + if scenario == "always_401": + self._send_error( + HTTPStatus.UNAUTHORIZED, + "authentication_error", + "Injected authentication failure", + ) + return + if scenario == "partial_then_success" and request_number == 1: + self._send_partial_stream(payload) + return + + finish_reason = "length" if scenario == "length" else "stop" + answer = "" if scenario == "empty_stop" else _extract_requested_answer( + payload, self.state.response_text + ) + if payload.get("stream") is True: + self._send_complete_stream(payload, answer, finish_reason) + else: + self._send_non_stream_response(payload, answer, finish_reason) + + def _send_error( + self, + status: int, + error_type: str, + message: str, + *, + retry_after: float | None = None, + ) -> None: + body = json.dumps( + { + "error": { + "message": message, + "type": error_type, + "param": None, + "code": error_type, + } + } + ).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.send_header("x-request-id", f"mock-{uuid.uuid4().hex}") + if retry_after is not None: + self.send_header("Retry-After", str(retry_after)) + self.end_headers() + self.wfile.write(body) + + @staticmethod + def _completion_chunk( + *, + request_id: str, + model: str, + delta: dict[str, Any], + finish_reason: str | None, + usage: dict[str, int] | None = None, + choices: bool = True, + ) -> bytes: + payload: dict[str, Any] = { + "id": request_id, + "object": "chat.completion.chunk", + "created": int(time.time()), + "model": model, + "choices": ( + [{"index": 0, "delta": delta, "finish_reason": finish_reason}] + if choices + else [] + ), + } + if usage is not None: + payload["usage"] = usage + return f"data: {json.dumps(payload)}\n\n".encode() + + def _send_complete_stream( + self, payload: dict[str, Any], answer: str, finish_reason: str + ) -> None: + request_id = f"chatcmpl-mock-{uuid.uuid4().hex}" + model = payload["model"] + code = f'final_answer({json.dumps(answer)})' if answer else "" + chunks = [ + self._completion_chunk( + request_id=request_id, + model=model, + delta={"role": "assistant", "content": ""}, + finish_reason=None, + ), + self._completion_chunk( + request_id=request_id, + model=model, + delta={"reasoning_content": "Deterministic mock reasoning. "}, + finish_reason=None, + ), + ] + chunks.extend( + self._completion_chunk( + request_id=request_id, + model=model, + delta={"content": code[index : index + 8]}, + finish_reason=None, + ) + for index in range(0, len(code), 8) + ) + chunks.extend( + [ + self._completion_chunk( + request_id=request_id, + model=model, + delta={}, + finish_reason=finish_reason, + ), + self._completion_chunk( + request_id=request_id, + model=model, + delta={}, + finish_reason=None, + usage={ + "prompt_tokens": 100, + "completion_tokens": max(1, len(code) // 4), + "total_tokens": 100 + max(1, len(code) // 4), + }, + choices=False, + ), + b"data: [DONE]\n\n", + ] + ) + self._send_bytes(HTTPStatus.OK, b"".join(chunks), "text/event-stream") + + def _send_partial_stream(self, payload: dict[str, Any]) -> None: + request_id = f"chatcmpl-mock-{uuid.uuid4().hex}" + model = payload["model"] + chunks = [ + self._completion_chunk( + request_id=request_id, + model=model, + delta={"reasoning_content": f"CMSR_LEAK_{index:02d}"}, + finish_reason=None, + ) + for index in range(12) + ] + self.send_response(HTTPStatus.OK) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Transfer-Encoding", "chunked") + self.send_header("x-request-id", f"mock-{uuid.uuid4().hex}") + self.end_headers() + for chunk in chunks: + self.wfile.write(f"{len(chunk):X}\r\n".encode() + chunk + b"\r\n") + self.wfile.flush() + time.sleep(self.state.partial_chunk_delay) + self.close_connection = True + try: + self.connection.shutdown(socket.SHUT_RDWR) + except OSError: + pass + self.connection.close() + + def _send_non_stream_response( + self, payload: dict[str, Any], answer: str, finish_reason: str + ) -> None: + code = f'final_answer({json.dumps(answer)})' if answer else "" + self._send_json( + HTTPStatus.OK, + { + "id": f"chatcmpl-mock-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": int(time.time()), + "model": payload["model"], + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "reasoning_content": "Deterministic mock reasoning. ", + "content": code, + }, + "finish_reason": finish_reason, + } + ], + "usage": { + "prompt_tokens": 100, + "completion_tokens": max(1, len(code) // 4), + "total_tokens": 100 + max(1, len(code) // 4), + }, + }, + ) + + +class OpenAICompatibleMockServer(ThreadingHTTPServer): + daemon_threads = True + + def __init__(self, address: tuple[str, int], state: MockState | None = None): + self.mock_state = state or MockState() + super().__init__(address, OpenAICompatibleMockHandler) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--host", default="0.0.0.0") + parser.add_argument("--port", type=int, default=18080) + args = parser.parse_args() + server = OpenAICompatibleMockServer((args.host, args.port)) + print(f"openai_compatible_mock_ready={args.host}:{args.port}", flush=True) + server.serve_forever() + + +if __name__ == "__main__": + main() diff --git a/test/common/test_openai_compatible_mock_server.py b/test/common/test_openai_compatible_mock_server.py new file mode 100644 index 000000000..65feb7166 --- /dev/null +++ b/test/common/test_openai_compatible_mock_server.py @@ -0,0 +1,172 @@ +from __future__ import annotations + +import http.client +import json +import threading +import urllib.error +import urllib.request + +import pytest + +from test.common.openai_compatible_mock_server import OpenAICompatibleMockServer + + +@pytest.fixture() +def mock_server(): + server = OpenAICompatibleMockServer(("127.0.0.1", 0)) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield server, server.server_address[1] + finally: + server.shutdown() + server.server_close() + thread.join(timeout=2) + + +def _post_json(port: int, path: str, payload: dict): + request = urllib.request.Request( + f"http://127.0.0.1:{port}{path}", + data=json.dumps(payload).encode(), + headers={"Authorization": "Bearer test-key", "Content-Type": "application/json"}, + method="POST", + ) + with urllib.request.urlopen(request, timeout=3) as response: + return response.status, dict(response.headers), response.read() + + +def _stream_content(body: bytes) -> str: + content_parts = [] + for line in body.decode().splitlines(): + if not line.startswith("data: ") or line == "data: [DONE]": + continue + payload = json.loads(line[6:]) + for choice in payload.get("choices", []): + content = choice.get("delta", {}).get("content") + if content: + content_parts.append(content) + return "".join(content_parts) + + +def test_cmsr_mock_stream_matches_nexent_openai_contract(mock_server): + _server, port = mock_server + status, _headers, body = _post_json( + port, + "/v1/chat/completions", + { + "model": "nexent-mock-model", + "messages": [{"role": "user", "content": "Reply with exactly CONTRACT_OK"}], + "stream": True, + "stream_options": {"include_usage": True}, + "max_tokens": 512, + "stop": ["Observation:"], + }, + ) + + assert status == 200 + assert b'"object": "chat.completion.chunk"' in body + assert b'"reasoning_content": "Deterministic mock reasoning. "' in body + assert _stream_content(body) == 'final_answer("CONTRACT_OK")' + assert b'"finish_reason": "stop"' in body + assert b'"prompt_tokens": 100' in body + assert body.endswith(b"data: [DONE]\n\n") + + stats = _server.mock_state.snapshot() + assert stats["request_count"] == 1 + assert stats["requests"] == [ + { + "request_number": 1, + "model": "nexent-mock-model", + "stream": True, + "include_usage": True, + "message_count": 1, + "max_tokens": 512, + "stop_count": 1, + "has_bearer_auth": True, + } + ] + + +@pytest.mark.parametrize( + ("scenario", "expected_status", "expected_type"), + [ + ("always_429", 429, "rate_limit_error"), + ("always_503", 503, "server_error"), + ("always_401", 401, "authentication_error"), + ], +) +def test_cmsr_mock_standard_error_contract( + mock_server, scenario, expected_status, expected_type +): + _server, port = mock_server + _post_json(port, "/__control", {"scenario": scenario, "retry_after": 0.25}) + + with pytest.raises(urllib.error.HTTPError) as caught: + _post_json( + port, + "/v1/chat/completions", + {"model": "nexent-mock-model", "messages": [], "stream": True}, + ) + + assert caught.value.code == expected_status + error = json.loads(caught.value.read()) + assert error["error"]["type"] == expected_type + if expected_status in {429, 503}: + assert caught.value.headers["Retry-After"] == "0.25" + + +def test_cmsr_mock_partial_then_success_is_deterministic(mock_server): + server, port = mock_server + _post_json( + port, + "/__control", + { + "scenario": "partial_then_success", + "response_text": "RECOVERED", + "partial_chunk_delay": 0, + }, + ) + request_body = json.dumps( + {"model": "nexent-mock-model", "messages": [], "stream": True} + ) + + connection = http.client.HTTPConnection("127.0.0.1", port, timeout=3) + connection.request( + "POST", + "/v1/chat/completions", + body=request_body, + headers={"Content-Type": "application/json", "Authorization": "Bearer test"}, + ) + response = connection.getresponse() + with pytest.raises(http.client.IncompleteRead) as caught: + response.read() + assert b"CMSR_LEAK_00" in caught.value.partial + assert b"CMSR_LEAK_11" in caught.value.partial + connection.close() + + status, _headers, body = _post_json( + port, + "/v1/chat/completions", + {"model": "nexent-mock-model", "messages": [], "stream": True}, + ) + assert status == 200 + assert _stream_content(body) == 'final_answer("RECOVERED")' + assert server.mock_state.snapshot()["request_count"] == 2 + + +def test_cmsr_mock_non_stream_and_models_contract(mock_server): + _server, port = mock_server + with urllib.request.urlopen(f"http://127.0.0.1:{port}/v1/models", timeout=3) as response: + models = json.loads(response.read()) + assert models["data"][0]["id"] == "nexent-mock-model" + + status, _headers, body = _post_json( + port, + "/v1/chat/completions", + {"model": "nexent-mock-model", "messages": [], "stream": False}, + ) + completion = json.loads(body) + assert status == 200 + assert completion["object"] == "chat.completion" + assert completion["choices"][0]["finish_reason"] == "stop" + assert completion["usage"]["total_tokens"] > 100 diff --git a/test/sdk/core/agents/test_core_agent.py b/test/sdk/core/agents/test_core_agent.py index 38432812c..0968b2aaf 100644 --- a/test/sdk/core/agents/test_core_agent.py +++ b/test/sdk/core/agents/test_core_agent.py @@ -2715,6 +2715,42 @@ def mock_step_stream(_action_step): assert len(agent.memory.steps) == 2 assert agent.memory.steps[0].error is not None + def test_cmsr_004_terminal_model_error_stops_react_without_memory_append( + self, monkeypatch + ): + """A depleted model budget must not become another recoverable ReAct step.""" + + class FakeAgentError(Exception): + pass + + monkeypatch.setattr(core_agent_module, "AgentError", FakeAgentError) + agent = self._create_canonical_run_agent(monkeypatch) + terminal = core_agent_module.ModelInvocationTerminalError( + SimpleNamespace(value="model_timeout"), + 5, + cause=TimeoutError("provider detail"), + ) + physical_steps = 0 + + def failing_step(_action_step): + nonlocal physical_steps + physical_steps += 1 + if False: + yield None + raise terminal + + agent._step_stream = failing_step + + with pytest.raises(core_agent_module.ModelInvocationTerminalError) as exc_info: + list(agent._run_stream("test task", max_steps=10)) + + assert exc_info.value is terminal + assert physical_steps == 1 + assert agent.step_number == 1 + assert agent.memory.steps == [] + agent._finalize_step.assert_not_called() + agent._collect_step_metrics.assert_not_called() + def test_planning_run_retries_empty_direct_answer_then_verifies_valid_answer(self, monkeypatch): """Planning runs reset state, retry an empty answer, and verify the next answer.""" module = core_agent_module @@ -2881,6 +2917,22 @@ def test_handle_max_steps_reached_model_error_fallback(self): ] assert len(error_calls) >= 1 + def test_cmsr_004_max_steps_propagates_terminal_model_error(self): + agent, module = self._create_agent_for_handle_max_steps_test() + terminal = module.ModelInvocationTerminalError( + SimpleNamespace(value="model_timeout"), + 5, + cause=TimeoutError("provider detail"), + ) + agent.model = MagicMock(side_effect=terminal) + agent._finalize_step = MagicMock() + + with pytest.raises(module.ModelInvocationTerminalError) as exc_info: + agent._handle_max_steps_reached("original task") + + assert exc_info.value is terminal + agent._finalize_step.assert_not_called() + def test_handle_max_steps_reached_empty_content_uses_fallback(self, caplog, monkeypatch): """Empty max-step synthesis returns a visible fallback and records why.""" agent, _module = self._create_agent_for_handle_max_steps_test() diff --git a/test/sdk/core/agents/test_nexent_agent.py b/test/sdk/core/agents/test_nexent_agent.py index 80b369298..faa630a36 100644 --- a/test/sdk/core/agents/test_nexent_agent.py +++ b/test/sdk/core/agents/test_nexent_agent.py @@ -2028,6 +2028,32 @@ def test_agent_run_with_observer_with_exception(nexent_agent_instance, mock_core ) +def test_cmsr_004_terminal_model_error_emits_one_safe_error( + nexent_agent_instance, mock_core_agent +): + nexent_agent_instance.agent = mock_core_agent + terminal_error_type = nexent_agent.ModelInvocationTerminalError + model_error_code = terminal_error_type.safe_message.__globals__["ModelErrorCode"] + terminal = terminal_error_type( + model_error_code.SERVICE_UNAVAILABLE, + 5, + cause=RuntimeError("private provider body"), + ) + mock_core_agent.run.side_effect = terminal + + with pytest.raises(terminal_error_type) as exc_info: + nexent_agent_instance.agent_run_with_observer("test query") + + assert exc_info.value is terminal + mock_core_agent.observer.add_message.assert_called_once_with( + agent_name="test_agent", + process_type=ProcessType.ERROR, + content="The model service is temporarily unavailable. Please try again later.", + error_code="model_service_unavailable", + retryable=False, + ) + + def test_agent_run_with_observer_invalid_agent_type(nexent_agent_instance): """Test agent_run_with_observer raises TypeError when agent is not a CoreAgent.""" nexent_agent_instance.agent = "not_core_agent" diff --git a/test/sdk/core/agents/test_subagent_wrapper.py b/test/sdk/core/agents/test_subagent_wrapper.py index 1cc88287e..ba336f996 100644 --- a/test/sdk/core/agents/test_subagent_wrapper.py +++ b/test/sdk/core/agents/test_subagent_wrapper.py @@ -7,6 +7,7 @@ import pytest from nexent.core.agents.subagent_wrapper import SubAgentToolWrapper, _default_task_extractor +from nexent.core.model_errors import ModelErrorCode, ModelInvocationTerminalError class InnerAgent: @@ -146,3 +147,23 @@ def test_call_still_balances_observer_events_when_inner_raises(observer: Mock) - start_kwargs = observer.add_subagent_start.call_args.kwargs end_kwargs = observer.add_subagent_end.call_args.kwargs assert start_kwargs["invocation_id"] == end_kwargs["invocation_id"] + + +def test_cmsr_004_managed_subagent_preserves_terminal_model_error(observer: Mock) -> None: + terminal = ModelInvocationTerminalError( + ModelErrorCode.SERVICE_UNAVAILABLE, + 5, + cause=RuntimeError("provider detail"), + ) + wrapper = SubAgentToolWrapper( + Mock(side_effect=terminal), + observer, + agent_id="agent-1", + agent_name="Research", + ) + + with pytest.raises(ModelInvocationTerminalError) as exc_info: + wrapper(task="x") + + assert exc_info.value is terminal + assert observer.add_subagent_end.call_count == 1 diff --git a/test/sdk/core/models/test_model_silent_retry.py b/test/sdk/core/models/test_model_silent_retry.py new file mode 100644 index 000000000..d69b1b719 --- /dev/null +++ b/test/sdk/core/models/test_model_silent_retry.py @@ -0,0 +1,37 @@ +import pytest + +from nexent.core.models.retry import ( + DEFAULT_MODEL_RETRY, + ModelErrorCode, + classify_model_error, +) + + +class HttpError(RuntimeError): + def __init__(self, status_code: int, message: str = "provider error"): + super().__init__(message) + self.status_code = status_code + + +def test_cmsr_001_default_budget_is_five_total_attempts(): + assert DEFAULT_MODEL_RETRY.max_attempts == 5 + + +@pytest.mark.parametrize( + ("error", "retryable", "error_code"), + [ + (HttpError(429), True, ModelErrorCode.RATE_LIMIT_EXHAUSTED), + (HttpError(503), True, ModelErrorCode.SERVICE_UNAVAILABLE), + (TimeoutError("read timeout"), True, ModelErrorCode.TIMEOUT), + (ConnectionError("connection reset"), True, ModelErrorCode.CONNECTION_ERROR), + (HttpError(401), False, ModelErrorCode.AUTHENTICATION_ERROR), + (HttpError(404), False, ModelErrorCode.NOT_FOUND), + (HttpError(422), False, ModelErrorCode.INVALID_REQUEST), + (RuntimeError("unclassified provider bug"), False, ModelErrorCode.UNKNOWN_ERROR), + ], +) +def test_cmsr_002_error_classification_is_typed(error, retryable, error_code): + classification = classify_model_error(error) + + assert classification.retryable is retryable + assert classification.error_code is error_code diff --git a/test/sdk/core/models/test_openai_llm.py b/test/sdk/core/models/test_openai_llm.py index be4840333..8aa9cc247 100644 --- a/test/sdk/core/models/test_openai_llm.py +++ b/test/sdk/core/models/test_openai_llm.py @@ -2,6 +2,7 @@ import types import importlib.util from pathlib import Path +from types import SimpleNamespace # Ensure SDK package is importable by adding sdk/ to sys.path (do not fallback to stubs) sys.path.insert(0, str(Path(__file__).resolve().parents[4] / "sdk")) @@ -391,8 +392,9 @@ def close(self): stream = TimeoutStream(chunks_before_timeout) model.client.chat.completions.create = lambda **kwargs: stream - with pytest.raises(openai_llm_module.httpx.ReadTimeout): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError) as exc_info: model([{"role": "user", "content": "secret-prompt"}]) + assert exc_info.value.error_code is openai_llm_module.ModelErrorCode.TIMEOUT timeout_records = [ record for record in caplog.records @@ -1011,15 +1013,13 @@ def test_provider_context_overflow_stops_after_two_recovery_dispatches(openai_mo ) with patch.object(openai_model_instance, "_prepare_completion_kwargs", return_value={}): - with pytest.raises( - openai_llm_module.ProviderContextOverflowRetryExhausted, - match="persisted after two recovery dispatches", - ): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError) as exc_info: openai_model_instance.__call__( messages, context_rebuild=lambda: messages, _overflow_recovery_ordinal=2, ) + assert exc_info.value.error_code is openai_llm_module.ModelErrorCode.CONTEXT_OVERFLOW def test_provider_context_overflow_without_rebuild_is_retry_unsafe(openai_model_instance): @@ -1029,11 +1029,9 @@ def test_provider_context_overflow_without_rebuild_is_retry_unsafe(openai_model_ ) with patch.object(openai_model_instance, "_prepare_completion_kwargs", return_value={}): - with pytest.raises( - openai_llm_module.ProviderContextOverflowRetryUnsafe, - match="cannot be safely rebuilt", - ): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError) as exc_info: openai_model_instance.__call__(messages, context_rebuild=None) + assert exc_info.value.error_code is openai_llm_module.ModelErrorCode.CONTEXT_OVERFLOW def test_provider_context_overflow_does_not_recover_unrelated_error(openai_model_instance): @@ -1353,10 +1351,11 @@ def test_call_rejects_reasoning_only_response_and_records_diagnostics( ] with pytest.raises( - openai_llm_module.EmptyModelResponseError, + openai_llm_module.ModelInvocationTerminalError, match="finish_reason=length", - ): + ) as exc_info: openai_model_instance.__call__(messages) + assert exc_info.value.error_code is openai_llm_module.ModelErrorCode.EMPTY_RESPONSE_EXHAUSTED diagnostics = openai_model_instance.last_response_diagnostics assert diagnostics["finish_reason"] == "length" @@ -1447,8 +1446,25 @@ def test_init_with_ssl_verify_true(): assert kwargs["timeout"].read == 60.0 -def test_ut_sdk_tlm_035_uses_openai_http_implementation_timeout(): - """Use the Timeout class owned by the HTTP implementation behind OpenAI.""" +def test_cmsr_001_init_disables_hidden_openai_transport_retries(): + captured = {} + + def fake_base_init(self, *args, **kwargs): + captured.update(kwargs) + self.client = SimpleNamespace() + + with patch.object( + openai_llm_module.OpenAIServerModel, + "__init__", + fake_base_init, + ): + ImportedOpenAIModel(observer=MagicMock()) + + assert captured["client_kwargs"]["max_retries"] == 0 + + +def test_ut_sdk_tlm_035_falls_back_to_public_httpx_timeout_for_test_double(): + """Use public httpx when the injected OpenAI client has no HTTP base.""" class SDKTimeout: def __init__(self, *, connect, read, write, pool): @@ -1457,11 +1473,7 @@ def __init__(self, *, connect, read, write, pool): self.write = write self.pool = pool - sdk_httpx = types.SimpleNamespace(Timeout=SDKTimeout) - openai_base_client = types.ModuleType("openai._base_client") - openai_base_client.httpx2 = sdk_httpx - - with patch.dict(sys.modules, {"openai._base_client": openai_base_client}), \ + with patch.object(openai_llm_module.httpx, "Timeout", SDKTimeout), \ patch("openai.DefaultHttpxClient") as mock_httpx_client: ImportedOpenAIModel(observer=MagicMock(), ssl_verify=True) @@ -1472,6 +1484,29 @@ def __init__(self, *, connect, read, write, pool): ) +def test_cmsr_compatible_timeout_uses_default_clients_http_implementation(): + """OpenAI's httpx2 client must receive an httpx2 timeout, not httpx.Timeout.""" + + compatible_timeout = MagicMock() + timeout_type = MagicMock(return_value=compatible_timeout) + http_module = SimpleNamespace(Timeout=timeout_type) + compatible_client_base = type("Client", (), {}) + compatible_client_base.__module__ = "httpx2._client" + default_client = type("DefaultClient", (compatible_client_base,), {}) + + with patch.object(openai_llm_module.importlib, "import_module", return_value=http_module): + result = openai_llm_module._build_compatible_http_timeout( + default_client, + connect=10.0, + read=60.0, + write=30.0, + pool=10.0, + ) + + assert result is compatible_timeout + timeout_type.assert_called_once_with(connect=10.0, read=60.0, write=30.0, pool=10.0) + + # --------------------------------------------------------------------------- # Tests for monitoring and token_tracker integration # --------------------------------------------------------------------------- @@ -1718,10 +1753,11 @@ def test_call_api_returns_string_raises_value_error(openai_model_instance): messages = [{"role": "user", "content": [{"text": "Hello"}]}] with patch.object(openai_model_instance, "_prepare_completion_kwargs", return_value={}): + openai_model_instance.retry_config = _retry_model_config() # Mock the client to return a string instead of a stream openai_model_instance.client.chat.completions.create.return_value = "error: rate limit exceeded" - with pytest.raises(ValueError, match="LLM API returned error string: error: rate limit exceeded"): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError, match="LLM API returned error string: error: rate limit exceeded"): openai_model_instance.__call__(messages) @@ -1730,10 +1766,11 @@ def test_call_api_returns_dict_with_error_raises_value_error(openai_model_instan messages = [{"role": "user", "content": [{"text": "Hello"}]}] with patch.object(openai_model_instance, "_prepare_completion_kwargs", return_value={}): + openai_model_instance.retry_config = _retry_model_config() # Mock the client to return a dict error response openai_model_instance.client.chat.completions.create.return_value = {"error": "rate limit exceeded"} - with pytest.raises(ValueError, match="LLM API returned error: rate limit exceeded"): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError, match="LLM API returned error: rate limit exceeded"): openai_model_instance.__call__(messages) @@ -1742,10 +1779,11 @@ def test_call_api_returns_dict_with_message_raises_value_error(openai_model_inst messages = [{"role": "user", "content": [{"text": "Hello"}]}] with patch.object(openai_model_instance, "_prepare_completion_kwargs", return_value={}): + openai_model_instance.retry_config = _retry_model_config() # Mock the client to return a dict with 'message' field openai_model_instance.client.chat.completions.create.return_value = {"message": "invalid api key"} - with pytest.raises(ValueError, match="LLM API returned error: invalid api key"): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError, match="LLM API returned error: invalid api key"): openai_model_instance.__call__(messages) @@ -1754,10 +1792,11 @@ def test_call_api_returns_plain_dict_raises_value_error(openai_model_instance): messages = [{"role": "user", "content": [{"text": "Hello"}]}] with patch.object(openai_model_instance, "_prepare_completion_kwargs", return_value={}): + openai_model_instance.retry_config = _retry_model_config() # Mock the client to return a plain dict openai_model_instance.client.chat.completions.create.return_value = {"status": "fail"} - with pytest.raises(ValueError, match="LLM API returned error:"): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError, match="LLM API returned error:"): openai_model_instance.__call__(messages) @@ -2405,34 +2444,89 @@ def fake_create(stream=True, **kwargs): openai_model_instance.retry_config = _retry_model_config(max_attempts=max_attempts) openai_model_instance.client.chat.completions.create.side_effect = fake_create - with pytest.raises(_StatusErr) as exc_info: + with pytest.raises(openai_llm_module.ModelInvocationTerminalError) as exc_info: openai_model_instance.__call__([{"role": "user", "content": "hello"}]) - assert exc_info.value.status_code == 503 + assert exc_info.value.error_code is openai_llm_module.ModelErrorCode.SERVICE_UNAVAILABLE + assert exc_info.value.attempts == max_attempts assert calls["n"] == max_attempts assert openai_model_instance.last_retry_count == max_attempts - 1 -def test_non_retryable_fails_immediately(openai_model_instance): - """A 401 must NOT be retried; the call fails on the first attempt.""" +def test_cmsr_001_default_retry_budget_makes_exactly_five_physical_calls( + openai_model_instance, +): calls = {"n": 0} def fake_create(stream=True, **kwargs): calls["n"] += 1 - raise _StatusErr(401, "Unauthorized") + raise _StatusErr(503, "Service Unavailable") + + openai_model_instance.retry_config = openai_llm_module.ModelRetryConfig( + backoff_base_seconds=0, + max_backoff_seconds=0, + jitter=False, + ) + openai_model_instance.client.chat.completions.create.side_effect = fake_create + + with pytest.raises(openai_llm_module.ModelInvocationTerminalError) as exc_info: + openai_model_instance.__call__([{"role": "user", "content": "hello"}]) + + assert calls["n"] == 5 + assert exc_info.value.attempts == 5 + + +@pytest.mark.parametrize( + ("failure", "expected_code"), + [ + (_StatusErr(401, "Unauthorized"), openai_llm_module.ModelErrorCode.AUTHENTICATION_ERROR), + (_StatusErr(400, "Bad request"), openai_llm_module.ModelErrorCode.INVALID_REQUEST), + (RuntimeError("unclassified provider bug"), openai_llm_module.ModelErrorCode.UNKNOWN_ERROR), + ], +) +def test_non_retryable_fails_immediately(openai_model_instance, failure, expected_code): + """Authentication, request and unknown failures must fail on the first call.""" + calls = {"n": 0} + + def fake_create(stream=True, **kwargs): + calls["n"] += 1 + raise failure openai_model_instance.retry_config = _retry_model_config() openai_model_instance.client.chat.completions.create.side_effect = fake_create - with pytest.raises(_StatusErr) as exc_info: + with pytest.raises(openai_llm_module.ModelInvocationTerminalError) as exc_info: openai_model_instance.__call__([{"role": "user", "content": "hello"}]) - assert exc_info.value.status_code == 401 + assert exc_info.value.error_code is expected_code assert calls["n"] == 1 -def test_reasoning_only_stop_response_retries_once_then_propagates(openai_model_instance): - """A reasoning-only stop response gets one transparent retry.""" +def test_cmsr_002_retry_after_controls_retry_wait(openai_model_instance): + calls = {"n": 0} + rate_limit = _StatusErr(429, "Rate limit") + rate_limit.response = SimpleNamespace(headers={"Retry-After": "3.5"}) + + def fake_create(stream=True, **kwargs): + calls["n"] += 1 + if calls["n"] == 1: + raise rate_limit + return [_make_content_chunk("ok")] + + wait_event = MagicMock() + wait_event.is_set.return_value = False + openai_model_instance.stop_event = wait_event + openai_model_instance.retry_config = _retry_model_config(backoff_base=1.0) + openai_model_instance.client.chat.completions.create.side_effect = fake_create + + result = openai_model_instance.__call__([{"role": "user", "content": "hello"}]) + + assert result is not None + wait_event.wait.assert_called_once_with(3.5) + + +def test_reasoning_only_stop_response_exhausts_shared_attempt_budget(openai_model_instance): + """A reasoning-only stop response uses the configured shared attempt budget.""" calls = {"n": 0} def fake_create(stream=True, **kwargs): @@ -2446,10 +2540,11 @@ def fake_create(stream=True, **kwargs): openai_model_instance.retry_config = _retry_model_config() openai_model_instance.client.chat.completions.create.side_effect = fake_create - with pytest.raises(openai_llm_module.EmptyModelResponseError): + with pytest.raises(openai_llm_module.ModelInvocationTerminalError) as exc_info: openai_model_instance.__call__([{"role": "user", "content": "hello"}]) - assert calls["n"] == 2 + assert exc_info.value.error_code is openai_llm_module.ModelErrorCode.EMPTY_RESPONSE_EXHAUSTED + assert calls["n"] == openai_model_instance.retry_config.max_attempts def test_reasoning_only_stop_response_recovers_on_retry(openai_model_instance): diff --git a/test/sdk/core/utils/test_observer_model_attempts.py b/test/sdk/core/utils/test_observer_model_attempts.py new file mode 100644 index 000000000..33a1437d7 --- /dev/null +++ b/test/sdk/core/utils/test_observer_model_attempts.py @@ -0,0 +1,66 @@ +import json + +from nexent.core.utils.observer import MessageObserver, ProcessType + + +def _events(observer: MessageObserver) -> list[dict]: + return [json.loads(item) for item in observer.get_cached_message()] + + +def test_cmsr_003_model_attempt_events_stamp_chunks_and_reset_parser_state(): + observer = MessageObserver(lang="en") + + observer.begin_model_attempt("attempt-one", 1) + observer.add_model_reasoning_content("partial reasoning") + observer.add_model_new_token("partial code") + observer.rollback_model_attempt("attempt-one", 1) + + assert not observer.token_buffer + assert not observer.think_buffer + assert observer.current_mode is ProcessType.MODEL_OUTPUT_THINKING + assert observer.in_think_mode is False + + events = _events(observer) + assert events[0] == { + "type": "model_attempt_control", + "content": "", + "phase": "begin", + "attempt_id": "attempt-one", + "attempt": 1, + } + assert events[1]["attempt_id"] == "attempt-one" + assert events[-1]["phase"] == "rollback" + + +def test_cmsr_003_new_attempt_never_inherits_failed_attempt_identity(): + observer = MessageObserver(lang="en") + observer.begin_model_attempt("attempt-one", 1) + observer.rollback_model_attempt("attempt-one", 1) + observer.begin_model_attempt("attempt-two", 2) + observer.add_model_reasoning_content("clean") + observer.commit_model_attempt("attempt-two", 2) + + events = _events(observer) + model_event = next(event for event in events if event["type"] == "model_output_deep_thinking") + assert model_event["attempt_id"] == "attempt-two" + assert events[-1]["phase"] == "commit" + + +def test_cmsr_004_terminal_error_keeps_string_content_and_stable_metadata(): + observer = MessageObserver(lang="en") + observer.add_message( + "agent", + ProcessType.ERROR, + "The model request failed.", + error_code="model_unknown_error", + retryable=False, + ) + + assert _events(observer) == [ + { + "type": "error", + "content": "The model request failed.", + "error_code": "model_unknown_error", + "retryable": False, + } + ] From a5f7a833b3323cb8b6cce8e550f9acf3bc806691 Mon Sep 17 00:00:00 2001 From: gs-aion Date: Sun, 20 Sep 2026 17:52:25 +0800 Subject: [PATCH 7/9] Fix: AIDP knowledge base bug fix (#3967) --- .../ext_components/aidp/apps/aidp_mgmt_app.py | 355 +++++- .../aidp/services/aidp_access_service.py | 113 +- .../aidp/services/aidp_service.py | 678 ++++++++++- frontend/const/knowledgeBase.ts | 5 + .../aidp/components/AidpCreateKbModal.tsx | 1009 +++++++---------- .../aidp/components/AidpDocumentList.tsx | 570 +++++++--- .../AidpKnowledgeBaseModalParts.tsx | 162 +++ .../components/AidpKnowledgeConfiguration.tsx | 343 ++++-- .../aidp/components/AidpKnowledgeList.tsx | 502 +++++--- .../aidp/components/AidpUpdateKbModal.tsx | 201 ++-- .../aidp/hooks/useAidpGroupOptions.ts | 30 + .../aidp/services/aidpKnowledgeService.ts | 166 ++- .../aidp/services/aidpUploadUtils.ts | 15 + frontend/lib/aidpDocumentStatus.ts | 128 +++ frontend/public/locales/en/common.json | 52 + frontend/public/locales/zh/common.json | 52 + frontend/services/api.ts | 4 + frontend/tests/aidpDocumentStatus.test.ts | 140 +++ .../mock_servers/aidp_mgmt_mock_server.py | 414 ++++++- .../aidp/test_aidp_access_service.py | 22 + .../ext_components/aidp/test_aidp_mgmt_app.py | 482 +++++++- test/ext_components/aidp/test_aidp_service.py | 883 ++++++++++++++- 22 files changed, 5089 insertions(+), 1237 deletions(-) create mode 100644 frontend/ext_components/aidp/components/AidpKnowledgeBaseModalParts.tsx create mode 100644 frontend/ext_components/aidp/hooks/useAidpGroupOptions.ts create mode 100644 frontend/ext_components/aidp/services/aidpUploadUtils.ts create mode 100644 frontend/lib/aidpDocumentStatus.ts create mode 100644 frontend/tests/aidpDocumentStatus.test.ts diff --git a/backend/ext_components/aidp/apps/aidp_mgmt_app.py b/backend/ext_components/aidp/apps/aidp_mgmt_app.py index 1b4bccefc..5fa9fe224 100644 --- a/backend/ext_components/aidp/apps/aidp_mgmt_app.py +++ b/backend/ext_components/aidp/apps/aidp_mgmt_app.py @@ -18,28 +18,27 @@ import time from http import HTTPStatus from typing import Annotated, List, Optional +from uuid import UUID -from fastapi import APIRouter, File, Path, Query, Request, UploadFile -from fastapi.responses import JSONResponse +from fastapi import APIRouter, File, HTTPException, Path, Query, Request, UploadFile +from fastapi.responses import JSONResponse, StreamingResponse +from nexent.core.concurrency import run_blocking from pydantic import BaseModel, Field from sqlalchemy.exc import IntegrityError -from nexent.core.concurrency import run_blocking +from starlette.background import BackgroundTask from consts.const import AIDP_API_KEY, AIDP_SERVER_URL from consts.error_code import ErrorCode from consts.exceptions import AppException, UnauthorizedError from database.user_tenant_db import get_user_role_by_tenant from ext_components.aidp.consts.aidp_exceptions import ( - AidpKbConflictError, - AidpKbNotFoundError, - AidpKbPermissionDeniedError, - AidpKbSyncError, AidpGroupValidationError, + AidpKbConflictError, ) from ext_components.aidp.database import aidp_permission_db from ext_components.aidp.services import aidp_permission_service as perms -from ext_components.aidp.services.aidp_kb_update_service import save_kb_settings from ext_components.aidp.services.aidp_access_service import ( + get_cached_aidp_channels, get_cached_aidp_doc_count, get_cached_aidp_kb_detail, invalidate_aidp_catalog_cache, @@ -47,25 +46,32 @@ invalidate_aidp_kb_detail_cache, resolve_current_aidp_access, ) +from ext_components.aidp.services.aidp_kb_update_service import save_kb_settings +from ext_components.aidp.services.aidp_permission_service import ( + EDIT, + PRIVATE, + READ_ONLY, + _validate_group_ids_strict, # noqa: F401 - retained as a module-level compatibility symbol +) from ext_components.aidp.services.aidp_service import ( _timestamp_to_iso, count_aidp_docs_impl, create_aidp_kb_impl, delete_aidp_kb_impl, get_aidp_kb_impl, + list_aidp_channels_impl, + list_aidp_doc_history_impl, list_aidp_docs_impl, list_aidp_models_impl, + remove_aidp_docs_impl, + stream_aidp_doc_impl, + select_aidp_channel, update_aidp_kb_impl, upload_aidp_docs_impl, ) -from ext_components.aidp.services.aidp_permission_service import ( - EDIT, - PRIVATE, - READ_ONLY, - _validate_group_ids_strict, -) from utils import auth_utils as auth_utils_module + aidp_mgmt_router = APIRouter(prefix="/aidp-mgmt") logger = logging.getLogger("aidp_mgmt_app") @@ -74,6 +80,14 @@ AIDP_OTHER_FILE_MAX_SIZE_BYTES = 1024 * 1024 * 1024 AIDP_SMALL_FILE_EXTENSIONS = {"txt", "xls", "xlsx", "csv"} +# AIDP document statuses (mirrors the file-history vocabulary): UPLOADING, +# PROCESSING and EXTRACTING are the stages a file walks through, COMPLETED and +# FAILED are the two terminal outcomes. Only the terminal ones end the wait, so +# every other reported status is counted as work in progress: `processing_count` +# is what keeps the frontend polling, and a build that reports a stage we do not +# know yet must not stop it early either. +_TERMINAL_DOC_STATUSES = ("COMPLETED", "FAILED") + def _upload_failure(file_name: str, reason_zh: str, reason_en: str) -> dict: return { @@ -187,6 +201,17 @@ class SetPermissionRequest(BaseModel): ) +class RemoveAidpDocumentsRequest(BaseModel): + + file_uuids: List[UUID] = Field(..., min_length=1, description="AIDP file UUIDs") + + +class DownloadAidpDocumentRequest(BaseModel): + """AIDP file selected for download.""" + + file_uuid: UUID = Field(..., description="AIDP file UUID") + + # --------------------------------------------------------------------------- # Auth helpers # --------------------------------------------------------------------------- @@ -243,11 +268,6 @@ def _raise_aidp_conflict(exc: IntegrityError) -> None: ) -# HTTPException is imported lazily to keep FastAPI's exception handler in -# control of the response body. -from fastapi import HTTPException # noqa: E402 (placed here to avoid editing mid-file) - - def _credentials() -> tuple[str, str]: return AIDP_SERVER_URL, AIDP_API_KEY @@ -263,8 +283,16 @@ def _is_user_role(user_id: str, tenant_id: str) -> bool: return (role or "USER").upper() == "USER" -def _current_accessible_rows(user_id: str, tenant_id: str) -> list[dict]: - """Return the current AIDP catalog intersected with local user access.""" +def _current_accessible_rows( + user_id: str, + tenant_id: str, + keyword: str | None = None, +) -> list[dict]: + """Return the current AIDP catalog intersected with local user access. + + ``keyword`` is forwarded to AIDP so the remote catalog is already narrowed + before the permission intersection runs. + """ server_url, api_key = _credentials() snapshot = resolve_current_aidp_access( server_url=server_url, @@ -272,6 +300,7 @@ def _current_accessible_rows(user_id: str, tenant_id: str) -> list[dict]: user_id=user_id, tenant_id=tenant_id, aidp_tenant_id="aidp", + keyword=keyword, ) return snapshot.accessible_rows @@ -316,6 +345,179 @@ def _load_cached_doc_count(server_url: str, api_key: str, kb_id: str) -> int: ) +# Knowledge bases whose all-status history has already been reported as +# unavailable. Documents are polled every few seconds while they process, so the +# fallback is reported once per KB instead of once per poll. +_HISTORY_FALLBACK_REPORTED: set[str] = set() + + +def _log_history_fallback(kds_id: str, reason: str) -> None: + """Report (once per KB) that the document list fell back to ingested files. + + The status column can only be filled from the history payload, so a silent + fallback shows up as a column of dashes. Logging the concrete reason keeps + the cause traceable: an AIDP build without the endpoint, a channel list whose + fields we cannot read, and a failing history request all land here. + """ + if kds_id in _HISTORY_FALLBACK_REPORTED: + logger.debug( + "AIDP file history still unavailable for KB %s (%s)", + kds_id, + reason, + ) + return + _HISTORY_FALLBACK_REPORTED.add(kds_id) + logger.warning( + "AIDP all-status file history unavailable for KB %s (%s); the document list falls " + "back to ingested files only, so files under processing stay invisible and the " + "status column stays empty", + kds_id, + reason, + ) + + +def _resolve_doc_history_channel( + server_url: str, + api_key: str, + kds_id: str, +) -> dict | None: + """Resolve the ingestion channel whose directory feeds ``kds_id``. + + Returns ``None`` when AIDP exposes no channel carrying both ``fs_id`` and a + source directory, in which case the caller keeps using the completed-files + listing. The channel payload shape is reported when resolution fails, so an + unreadable field name is visible instead of silently degrading. + """ + channels = get_cached_aidp_channels( + server_url=server_url, + api_key=api_key, + kds_id=kds_id, + loader=lambda: list_aidp_channels_impl(server_url, api_key, kds_id), + ) + channel = select_aidp_channel(channels, kds_id) + if channel is None: + _log_history_fallback( + kds_id, + "no channel exposes fs_id + a source dir (channels=%d, sample keys=%s)" + % (len(channels), sorted(channels[0].keys()) if channels else []), + ) + return channel + + +async def _load_doc_history( + server_url: str, + api_key: str, + kds_id: str, +) -> dict | None: + """Fetch the all-status document history for ``kds_id`` when available. + + Any failure degrades to the legacy completed-files listing rather than + failing the request: the channel/history endpoints are newer than the + document list and may be missing on older AIDP builds, and a knowledge base + with no matching channel must still render its ingested files. Every + degraded case is logged with the offending KB so the cause stays traceable. + """ + try: + channel = await run_blocking( + "aidp-doc-history-channel", + _resolve_doc_history_channel, + server_url, + api_key, + kds_id, + lane="control-io", + owner="config", + ) + if not channel: + # The reason is reported by ``_resolve_doc_history_channel``. + return None + history = await run_blocking( + "aidp-doc-history", + list_aidp_doc_history_impl, + server_url, + api_key, + channel["fs_id"], + channel["src_dir"], + kds_id, + lane="control-io", + owner="config", + ) + items = history.get("value") if isinstance(history, dict) else None + if isinstance(items, list) and not items: + # A resolved channel directory holding no file would blank the table + # and hide the KB's ingested files — the directory may simply not be + # where this KB's uploads live. The KB-scoped listing is always safe + # to show, so an empty history is treated as unusable rather than + # authoritative (an empty KB still renders an empty list either way). + _log_history_fallback( + kds_id, + "history returned no files for the resolved directory " + f"(fs_id={channel['fs_id']}, dir_path={channel['src_dir']})", + ) + return None + return history + except AppException as exc: + _log_history_fallback(kds_id, f"history request failed: {exc}") + return None + except Exception as exc: # noqa: BLE001 - history is an optional enhancement + _log_history_fallback(kds_id, f"unexpected history error: {exc!r}") + return None + + +def _history_sort_key(item: dict) -> tuple: + """Sort key placing the most recently uploaded file first. + + Numeric upload timestamps sort above entries that only expose an ISO + ``created_at`` string, and entries with neither sink to the bottom — the + ordering is only ever used to bring fresh uploads to the top of page 1. + """ + raw = item.get("first_upload_time") + if raw is None: + raw = item.get("create_time") + try: + return (1, float(raw)) + except (TypeError, ValueError): + created_at = item.get("created_at") + return (0, created_at) if isinstance(created_at, str) else (0, "") + + +def _paginate_history_documents(result: dict, page: int, page_size: int) -> dict: + """Slice an all-status history payload into one page. + + The history API returns the whole channel directory in one response, so the + total is exact and the document Count endpoint is not needed. Newest files + come first so an upload shows up at the top of page 1 as soon as it is + accepted, instead of only after ingestion completes. + + ``processing_count`` covers the WHOLE directory, not just the returned page: + the frontend keeps polling while it is non-zero, so a file still being + processed on another page also keeps the status column live. + """ + raw_items = result.get("value") + items = ( + [item for item in raw_items if isinstance(item, dict)] + if isinstance(raw_items, list) + else [] + ) + ordered = sorted(items, key=_history_sort_key, reverse=True) + start = (page - 1) * page_size + end = start + page_size + # Count every non-terminal status, so a file that is uploading or extracting + # keeps the frontend polling exactly like one that is being chunked. + processing_count = sum( + 1 + for item in ordered + if isinstance(item.get("status"), str) + and item["status"].strip().upper() not in _TERMINAL_DOC_STATUSES + ) + return { + "value": ordered[start:end], + "total_count": len(ordered), + "has_more": end < len(ordered), + "total_reliable": True, + "processing_count": processing_count, + } + + # --------------------------------------------------------------------------- # Handlers # --------------------------------------------------------------------------- @@ -326,21 +528,38 @@ async def list_knowledge_bases( request: Request, page: Annotated[int, Query(ge=1, description="Page number starting from 1")] = 1, page_size: Annotated[int, Query(ge=1, le=100, description="Page size from 1 to 100")] = 10, + keyword: Annotated[ + str | None, + Query(max_length=200, description="Optional name filter forwarded to AIDP"), + ] = None, ) -> JSONResponse: - """List KBs the caller can access. + """List KBs the caller can access, optionally filtered by name. Resolution order: - 1. Fetch every KB visible to the currently configured AIDP credentials. + 1. Fetch the AIDP catalog visible to the configured credentials, narrowed + server-side by ``keyword`` when one is supplied. 2. Intersect that catalog with the caller's effective Nexent permissions. 3. Paginate the intersection, then fetch details for the visible page. + + A non-blank ``keyword`` therefore narrows both the fetched set and the + reported ``total_count``: both describe the filtered, permission-scoped set. """ user_id, tenant_id = await _auth(request) server_url, api_key = _credentials() + normalized_keyword = (keyword or "").strip() or None started_at = time.perf_counter() + # ``run_blocking`` (not a bare to_thread) keeps AIDP catalog reads on the + # managed control-io lane, and the keyword is still forwarded so the remote + # catalog is narrowed server-side before the permission intersection. rows = await run_blocking( - "aidp-accessible-rows", _current_accessible_rows, user_id, tenant_id, - lane="control-io", owner="config", + "aidp-accessible-rows", + _current_accessible_rows, + user_id, + tenant_id, + normalized_keyword, + lane="control-io", + owner="config", ) access_resolve_ms = (time.perf_counter() - started_at) * 1000 total_count = len(rows) @@ -720,6 +939,28 @@ async def list_documents( server_url, api_key = _credentials() started_at = time.perf_counter() + + # Preferred source: the all-status history, so files appear in the list + # while they are still being chunked/embedded (and when they failed). The + # payload covers the whole channel directory, so this branch paginates + # in-process and reports an exact total. + history_result = await _load_doc_history(server_url, api_key, kds_id) + if history_result is not None: + result = _paginate_history_documents(history_result, page, page_size) + logger.info( + "AIDP document list timing: total_ms=%.1f kb_id=%s page=%d page_size=%d " + "page_count=%d total_count=%d total_reliable=True source=history", + (time.perf_counter() - started_at) * 1000, + kds_id, + page, + page_size, + len(result["value"]), + result["total_count"], + ) + return JSONResponse(status_code=HTTPStatus.OK, content=result) + + # Fallback: the completed-files listing (historical behaviour), used when + # AIDP has no channel/history support for this knowledge base. list_result, count_result = await asyncio.gather( run_blocking( "aidp-list-documents", @@ -769,11 +1010,14 @@ async def list_documents( result["total_count"] = int(total_count) result["has_more"] = has_more + # The completed-files listing carries no processing statuses, so nothing + # keeps the frontend polling for this knowledge base. + result["processing_count"] = 0 if not count_reliable: result["total_reliable"] = False logger.info( "AIDP document list timing: total_ms=%.1f kb_id=%s page=%d page_size=%d " - "page_count=%d total_count=%d total_reliable=%s", + "page_count=%d total_count=%d total_reliable=%s source=completed", (time.perf_counter() - started_at) * 1000, kds_id, page, @@ -785,6 +1029,65 @@ async def list_documents( return JSONResponse(status_code=HTTPStatus.OK, content=result) +@aidp_mgmt_router.post("/knowledge-bases/{kds_id}/documents/remove") +async def remove_documents( + request: Request, + kds_id: Annotated[str, Path(description="Knowledge base ID")], + body: RemoveAidpDocumentsRequest, +) -> JSONResponse: + """Remove AIDP documents.""" + user_id, tenant_id = await _auth(request) + perms.require_permission(kds_id, user_id, tenant_id, required="EDIT") + + server_url, api_key = _credentials() + result = await run_blocking( + "aidp-remove-documents", + remove_aidp_docs_impl, + server_url, + api_key, + kds_id, + [str(file_uuid) for file_uuid in body.file_uuids], + lane="control-io", + owner="config", + ) + + success_list = result["success_list"] + if success_list: + invalidate_aidp_kb_detail_cache(server_url, api_key, kds_id) + invalidate_aidp_doc_count_cache(server_url, api_key, kds_id) + return JSONResponse(status_code=HTTPStatus.OK, content=result) + + +@aidp_mgmt_router.post("/knowledge-bases/{kds_id}/documents/download") +async def download_document( + request: Request, + kds_id: Annotated[str, Path(description="Knowledge base ID")], + body: DownloadAidpDocumentRequest, +) -> StreamingResponse: + """Proxy an AIDP document as a binary attachment.""" + user_id, tenant_id = await _auth(request) + perms.require_permission(kds_id, user_id, tenant_id, required="READ") + + server_url, api_key = _credentials() + aidp_response = await stream_aidp_doc_impl( + server_url, + api_key, + kds_id, + str(body.file_uuid), + ) + response_headers = { + "Content-Disposition": aidp_response.headers["Content-Disposition"], + "X-File-Size": aidp_response.headers["X-File-Size"], + } + + return StreamingResponse( + aidp_response.aiter_bytes(), + media_type=aidp_response.headers["Content-Type"], + headers=response_headers, + background=BackgroundTask(aidp_response.aclose), + ) + + @aidp_mgmt_router.patch("/aidp-permissions/{kds_id}") async def set_permission( request: Request, diff --git a/backend/ext_components/aidp/services/aidp_access_service.py b/backend/ext_components/aidp/services/aidp_access_service.py index 8ebc23e90..80b1a9b49 100644 --- a/backend/ext_components/aidp/services/aidp_access_service.py +++ b/backend/ext_components/aidp/services/aidp_access_service.py @@ -23,15 +23,26 @@ _DETAIL_CACHE_MAX_ENTRIES = 256 _DOC_COUNT_CACHE_TTL_SECONDS = 30.0 _DOC_COUNT_CACHE_MAX_ENTRIES = 256 +_CHANNELS_CACHE_TTL_SECONDS = 60.0 +_CHANNELS_CACHE_MAX_ENTRIES = 32 _catalog_cache: OrderedDict[tuple[str, str], tuple[float, list[dict]]] = OrderedDict() _detail_cache: OrderedDict[tuple[str, str, str], tuple[float, dict]] = OrderedDict() _doc_count_cache: OrderedDict[tuple[str, str, str], tuple[float, int]] = OrderedDict() +_channels_cache: OrderedDict[tuple[str, str], tuple[float, list[dict]]] = OrderedDict() _catalog_inflight: dict[tuple[str, str], Future[Any]] = {} _detail_inflight: dict[tuple[str, str, str], Future[Any]] = {} _doc_count_inflight: dict[tuple[str, str, str], Future[Any]] = {} +_channels_inflight: dict[tuple[str, str], Future[Any]] = {} _catalog_versions: dict[tuple[str, str], int] = {} _detail_versions: dict[tuple[str, str, str], int] = {} _doc_count_versions: dict[tuple[str, str, str], int] = {} +_channels_versions: dict[tuple[str, str], int] = {} +# Keyword-filtered catalogs live in their own namespace. Sharing ``_catalog_cache`` +# would let a search result overwrite the full catalog under the same key and +# silently hide KBs from every non-search caller for the rest of the TTL. +_search_catalog_cache: OrderedDict[tuple[str, str, str], tuple[float, list[dict]]] = OrderedDict() +_search_catalog_inflight: dict[tuple[str, str, str], Future[Any]] = {} +_search_catalog_versions: dict[tuple[str, str, str], int] = {} _cache_lock = threading.RLock() _T = TypeVar("_T") @@ -136,13 +147,39 @@ def _get_remote_catalog( api_key: str, aidp_tenant_id: str, force_refresh: bool, + keyword: str | None = None, ) -> list[dict]: - key = _cache_key(server_url, aidp_tenant_id) + """Return the remote KB catalog, optionally filtered server-side by keyword. + + Keyword-filtered results are cached in their own namespace. Both variants + come from the same endpoint but describe different sets, so sharing a cache + slot would let a search response mask the full catalog (or vice versa) for + the remainder of the TTL. + """ + normalized_keyword = (keyword or "").strip() + catalog_key = _cache_key(server_url, aidp_tenant_id) + + if normalized_keyword: + return _get_or_load_cached( + cache=_search_catalog_cache, + inflight=_search_catalog_inflight, + versions=_search_catalog_versions, + key=(*catalog_key, normalized_keyword.lower()), + ttl_seconds=_CATALOG_CACHE_TTL_SECONDS, + max_entries=_CATALOG_CACHE_MAX_ENTRIES, + loader=lambda: _extract_remote_items( + fetch_all_aidp_knowledge_bases_impl( + server_url, api_key, keyword=normalized_keyword + ) + ), + force_refresh=force_refresh, + ) + return _get_or_load_cached( cache=_catalog_cache, inflight=_catalog_inflight, versions=_catalog_versions, - key=key, + key=catalog_key, ttl_seconds=_CATALOG_CACHE_TTL_SECONDS, max_entries=_CATALOG_CACHE_MAX_ENTRIES, loader=lambda: _extract_remote_items( @@ -196,6 +233,43 @@ def get_cached_aidp_doc_count( ) +def get_cached_aidp_channels( + server_url: str, + api_key: str, + kds_id: str, + loader: Callable[[], Any], + aidp_tenant_id: str = "aidp", + force_refresh: bool = False, +) -> list[dict]: + """Return one knowledge base's AIDP ingestion channels with short caching. + + Channels are catalogued per knowledge base, so the KB id is part of the cache + key: without it a channel resolved for one KB would be reused for another and + the history lookup would read the wrong directory. + + Every document-list request resolves a channel, and the UI polls that list + while files are still being processed, so the result is cached to keep + polling at one upstream request per interval. A channel describes how a KB + ingests files and is not affected by uploads, so a plain TTL is enough and + there is nothing to invalidate on write. + """ + + def _load() -> list[dict]: + return _extract_remote_items(loader()) + + key = (*_cache_key(server_url, aidp_tenant_id), str(kds_id)) + return _get_or_load_cached( + cache=_channels_cache, + inflight=_channels_inflight, + versions=_channels_versions, + key=key, + ttl_seconds=_CHANNELS_CACHE_TTL_SECONDS, + max_entries=_CHANNELS_CACHE_MAX_ENTRIES, + loader=_load, + force_refresh=force_refresh, + ) + + def resolve_current_aidp_access( server_url: str, api_key: str, @@ -203,8 +277,14 @@ def resolve_current_aidp_access( tenant_id: str, aidp_tenant_id: str = "aidp", force_refresh: bool = False, + keyword: str | None = None, ) -> AidpAccessSnapshot: - """Return the current AIDP catalog intersected with local user access.""" + """Return the current AIDP catalog intersected with local user access. + + ``keyword`` narrows the remote catalog server-side before the permission + intersection runs, so the returned snapshot only contains matching KBs. + Passing ``None``/blank keeps the full catalog. + """ started_at = time.perf_counter() remote_started_at = time.perf_counter() remote_items = _get_remote_catalog( @@ -212,6 +292,7 @@ def resolve_current_aidp_access( api_key=api_key, aidp_tenant_id=aidp_tenant_id, force_refresh=force_refresh, + keyword=keyword, ) remote_ms = (time.perf_counter() - remote_started_at) * 1000 permission_started_at = time.perf_counter() @@ -267,17 +348,40 @@ def invalidate_aidp_catalog_cache( api_key: str | None = None, aidp_tenant_id: str = "aidp", ) -> None: - """Invalidate one credential-scoped catalog, or every catalog when omitted.""" + """Invalidate one credential-scoped catalog, or every catalog when omitted. + + Keyword-filtered catalogs are dropped alongside the full catalog: they are + projections of the same remote data, so leaving them behind would let a + stale search result outlive the mutation that invalidated the listing. + """ with _cache_lock: if server_url is None or api_key is None: _catalog_cache.clear() for key in set(_catalog_versions) | set(_catalog_inflight): _catalog_versions[key] = _catalog_versions.get(key, 0) + 1 + _search_catalog_cache.clear() + for search_key in set(_search_catalog_versions) | set(_search_catalog_inflight): + _search_catalog_versions[search_key] = ( + _search_catalog_versions.get(search_key, 0) + 1 + ) return key = _cache_key(server_url, aidp_tenant_id) _catalog_cache.pop(key, None) _catalog_versions[key] = _catalog_versions.get(key, 0) + 1 + # Search keys are ``(*catalog_key, keyword)``; match them by prefix. + for search_key in [ + candidate + for candidate in set(_search_catalog_cache) + | set(_search_catalog_versions) + | set(_search_catalog_inflight) + if candidate[:2] == key + ]: + _search_catalog_cache.pop(search_key, None) + _search_catalog_versions[search_key] = ( + _search_catalog_versions.get(search_key, 0) + 1 + ) + def invalidate_aidp_kb_detail_cache( server_url: str | None = None, @@ -329,6 +433,7 @@ def invalidate_aidp_doc_count_cache( __all__ = [ "AidpAccessSnapshot", + "get_cached_aidp_channels", "get_cached_aidp_doc_count", "get_cached_aidp_kb_detail", "invalidate_aidp_catalog_cache", diff --git a/backend/ext_components/aidp/services/aidp_service.py b/backend/ext_components/aidp/services/aidp_service.py index ac5cb4d37..50400aee4 100644 --- a/backend/ext_components/aidp/services/aidp_service.py +++ b/backend/ext_components/aidp/services/aidp_service.py @@ -5,15 +5,16 @@ import logging import time from datetime import datetime, timezone -from typing import Any, Callable, Dict, List -from urllib.parse import urljoin +from typing import Any, Callable, Dict, List, NoReturn +from urllib.parse import quote, urljoin import httpx +from nexent.utils.http_client_manager import http_client_manager from consts.const import AIDP_TENANT_ID from consts.error_code import ErrorCode from consts.exceptions import AppException -from nexent.utils.http_client_manager import http_client_manager + logger = logging.getLogger("aidp_service") @@ -66,6 +67,39 @@ def _extract_upstream_error(response: httpx.Response) -> str | None: return reason[:_MAX_UPSTREAM_ERROR_REASON_LENGTH] +def _raise_aidp_http_error(error: httpx.HTTPStatusError, operation: str) -> NoReturn: + """Map an AIDP HTTP error to the common application exception format.""" + response = error.response + upstream_reason = _extract_upstream_error(response) + logger.exception( + "AIDP %s HTTP error: status_code=%s upstream_reason=%s", + operation, + response.status_code, + upstream_reason or "unavailable", + ) + details = { + "upstream_status": response.status_code, + "upstream_reason": upstream_reason, + } + error_code = { + 401: ErrorCode.AIDP_AUTH_ERROR, + 403: ErrorCode.AIDP_AUTH_ERROR, + 429: ErrorCode.AIDP_RATE_LIMIT, + }.get(response.status_code, ErrorCode.AIDP_SERVICE_ERROR) + fallback_message = { + ErrorCode.AIDP_AUTH_ERROR: f"AIDP authentication failed: {str(error)}", + ErrorCode.AIDP_RATE_LIMIT: f"AIDP rate limit exceeded: {str(error)}", + }.get( + error_code, + f"AIDP API HTTP error {response.status_code}: {str(error)}", + ) + raise AppException( + error_code, + upstream_reason or fallback_message, + details=details, + ) + + def _extract_upload_failures(response: httpx.Response) -> List[Dict[str, str]]: """Extract per-file upload failures from AIDP's structured error body.""" try: @@ -113,6 +147,48 @@ def _get_list_path(tenant_id: str | None = None) -> str: return f"/KnowledgeBase/Tenants/{_resolve_tenant_id(tenant_id)}/KnowledgeBases" +def _get_channels_path(kds_id: str, tenant_id: str | None = None) -> str: + """Build the knowledge-base scoped ingestion-pipeline (channel) API path. + + AIDP exposes the channel catalog per knowledge base + (``.../KnowledgeBases/{kds_id}/Channels``), mirroring the knowledge-file + endpoints. The channel carries the ``fs_id`` + source directory the history + request is addressed with, which is why the KB id is part of the path. + """ + return ( + f"/KnowledgeBase/Tenants/{_resolve_tenant_id(tenant_id)}" + f"/KnowledgeBases/{kds_id}/Channels" + ) + + +def _get_doc_history_path(kds_id: str, tenant_id: str | None = None) -> str: + """Build the knowledge-base scoped knowledge-file history API path. + + Like the channel catalog, the history endpoint belongs to one knowledge base + (``.../KnowledgeBases/{kds_id}/KnowledgeFiles/History``); the body addresses + the request to a channel directory inside that knowledge base. + """ + return ( + f"/KnowledgeBase/Tenants/{_resolve_tenant_id(tenant_id)}" + f"/KnowledgeBases/{kds_id}/KnowledgeFiles/History" + ) + + +def _build_list_query(page: int, page_size: int, keyword: str | None = None) -> str: + """Build the query string for a knowledge-base list request. + + AIDP applies ``keyword`` server-side (matching on the KB name), so the + value is forwarded verbatim. The parameter is omitted entirely when no + keyword is supplied, which keeps plain listing requests byte-identical to + the pre-search behaviour (same URL, so upstream caches still hit). + """ + query = f"?page={page}&page_size={page_size}" + normalized_keyword = (keyword or "").strip() + if normalized_keyword: + query += f"&keyword={quote(normalized_keyword, safe='')}" + return query + + def _timestamp_to_iso(value: Any) -> str | None: """Convert a numeric Unix timestamp (seconds or milliseconds) to ISO-8601 UTC. @@ -140,6 +216,31 @@ def _timestamp_to_iso(value: Any) -> str | None: return datetime.fromtimestamp(ts, tz=timezone.utc).isoformat().replace("+00:00", "Z") +# AIDP wraps list payloads in ``value`` (the shape every established endpoint in +# this adapter returns). The remaining keys cover builds that put the items +# behind a transport envelope, and a bare array is accepted as-is, so a renamed +# wrapper cannot silently turn a populated list into "nothing returned". +_AIDP_LIST_KEYS = ("value", "data", "result", "records", "items", "list") + + +def _extract_list_payload(payload: Any) -> list | None: + """Return the item list carried by an AIDP list response, if there is one. + + ``None`` means the response does not carry a list at all, which callers + report as an unexpected response format instead of degrading into an empty + result that looks like "no data upstream". + """ + if isinstance(payload, list): + return payload + if not isinstance(payload, dict): + return None + for key in _AIDP_LIST_KEYS: + candidate = payload.get(key) + if isinstance(candidate, list): + return candidate + return None + + def _normalize_aidp_doc(raw: Dict[str, Any]) -> Dict[str, Any]: """Map an AIDP document item to the shape the frontend expects. @@ -157,6 +258,86 @@ def _normalize_aidp_doc(raw: Dict[str, Any]) -> Dict[str, Any]: return out +def _normalize_doc_status(value: Any) -> str | None: + """Normalize an AIDP file status to its canonical upper-case token. + + AIDP reports ``COMPLETED`` / ``PROCESSING`` / ``FAILED``; normalizing here + means the frontend can compare against one form regardless of upstream + capitalization. ``None`` is returned for missing/blank values so callers can + tell "no status reported" apart from a real status. Non-string values (for + example a numeric ``state``) are ignored instead of being stringified. + """ + if isinstance(value, str) and value.strip(): + return value.strip().upper() + return None + + +# Deployments spell the per-file processing status differently, so it is read +# through aliases like the channel fields further down (``_CHANNEL_*_KEYS``). +# Only the forms actually observed are listed; adding one here is enough to +# support another deployment. +_HISTORY_STATUS_KEYS = ( + "status", + "file_status", + "fileStatus", + "file_state", + "doc_status", +) + + +def _extract_doc_status(raw: Dict[str, Any]) -> str | None: + """Return the first recognized file status carried by ``raw``. + + The canonical ``status`` wins, then the observed aliases in order, so a + deployment that renames the field keeps working without a code change. + """ + for key in _HISTORY_STATUS_KEYS: + status = _normalize_doc_status(raw.get(key)) + if status is not None: + return status + return None + + +def _normalize_history_doc(raw: Dict[str, Any]) -> Dict[str, Any]: + """Map an AIDP knowledge-file history item to the frontend document shape. + + The history payload carries the same identity fields as the document list + (``file_ino_no`` / ``file_name`` / ``file_size`` / ``file_type`` / + ``first_upload_time``) plus a processing status. Whatever field the upstream + used, the value is re-emitted under the canonical ``status`` key so the + frontend only has to read one field. A blank or missing upstream status is + dropped rather than forwarded as whitespace, so callers can treat "no status + reported" as a single state. + """ + out = _normalize_aidp_doc(raw) + status = _extract_doc_status(raw) + if status is None: + out.pop("status", None) + else: + out["status"] = status + return out + + +def _warn_when_history_status_unreadable(items: List[Any], fs_id: str) -> None: + """Log the payload shape when no history item carries a usable status. + + The processing status is the reason the history endpoint is used at all: a + payload whose status field we cannot read renders as a column of dashes, so + the keys of a sample item are logged to make the real field name + discoverable from the service log instead of leaving the gap invisible. + """ + item_dicts = [item for item in items if isinstance(item, dict)] + if not item_dicts or any("status" in item for item in item_dicts): + return + logger.warning( + "AIDP file history for fs_id=%s returned %d item(s) with no readable status; " + "sample keys=%s", + fs_id, + len(item_dicts), + sorted(item_dicts[0].keys()), + ) + + def _validate_params(server_url: str, api_key: str) -> str: """Validate parameters and return normalized base URL.""" if not server_url or not isinstance(server_url, str): @@ -183,6 +364,7 @@ def _validate_params(server_url: str, api_key: str) -> str: _AIDP_RETRY_BACKOFF_FACTOR = 0.5 _AIDP_RETRYABLE_STATUS_CODES = {408, 429, 500, 502, 503, 504} _AIDP_READ_TIMEOUT_SECONDS = 30.0 +_AIDP_DOWNLOAD_TIMEOUT_SECONDS = 120.0 def _request_with_retry( @@ -260,8 +442,12 @@ def fetch_aidp_knowledge_bases_impl( api_key: str, page: int = 1, page_size: int = 10, + keyword: str | None = None, ) -> Dict[str, Any]: - """Fetch a single page from AIDP API (simple passthrough).""" + """Fetch a single page from AIDP API (simple passthrough). + + ``keyword`` is forwarded to AIDP as an optional server-side filter. + """ normalized_url = _validate_params(server_url, api_key) headers = { @@ -269,7 +455,7 @@ def fetch_aidp_knowledge_bases_impl( "Content-Type": "application/json", } - list_path = f"{_get_list_path()}?page={page}&page_size={page_size}" + list_path = f"{_get_list_path()}{_build_list_query(page, page_size, keyword)}" list_url = urljoin(f"{normalized_url}/", list_path) logger.info("Fetching AIDP knowledge bases from %s", list_url) @@ -343,9 +529,25 @@ def _normalize_response(raw: Dict[str, Any]) -> Dict[str, Any]: } +def _kb_matches_keyword(item: Dict[str, Any], keyword: str) -> bool: + """Return whether a KB row carries ``keyword`` in its name or description. + + Only used as the fallback for AIDP builds that ignore the server-side + filter, so plain case-insensitive substring semantics are enough — the + upstream match set is always trusted when the upstream actually narrowed. + """ + needle = keyword.lower() + for field in ("kds_name", "name", "description"): + value = item.get(field) + if isinstance(value, str) and needle in value.lower(): + return True + return False + + def fetch_all_aidp_knowledge_bases_impl( server_url: str, api_key: str, + keyword: str | None = None, ) -> Dict[str, Any]: """Fetch every AIDP knowledge-base page using the dedicated Count API. @@ -354,8 +556,20 @@ def fetch_all_aidp_knowledge_bases_impl( the number of pages, and every list request uses the configured tenant. Duplicate resources are removed by ``kds_id`` while preserving their first-seen order. + + ``keyword`` is forwarded to every list call so AIDP filters server-side. + The Count API has no keyword support, so the page count is deliberately + derived from the unfiltered total: AIDP returns the matching subset on the + earliest pages and empty pages afterwards, which this loop tolerates. The + caller therefore still receives every match. + + A build that ignores ``keyword`` answers with the whole catalog instead. That + is detected after the loop (the collected set is not narrower than the + reported total) and the keyword is then applied locally, so a search cannot + silently return the unfiltered list. """ normalized_url = _validate_params(server_url, api_key) + normalized_keyword = (keyword or "").strip() page_size = 100 started_at = time.perf_counter() count_started_at = time.perf_counter() @@ -395,7 +609,10 @@ def fetch_all_aidp_knowledge_bases_impl( list_started_at = time.perf_counter() for current_page in range(1, total_pages + 1): - page_path = f"{_get_list_path()}?page={current_page}&page_size={page_size}" + page_path = ( + f"{_get_list_path()}" + f"{_build_list_query(current_page, page_size, keyword)}" + ) current_url = urljoin(f"{normalized_url}/", page_path) logger.info( @@ -446,6 +663,26 @@ def fetch_all_aidp_knowledge_bases_impl( all_items.append(item) accumulated_count = len(all_items) + result_count = accumulated_count + if normalized_keyword and accumulated_count >= total_count: + # The upstream answered a keyword request with the entire catalog, so + # this AIDP build ignores the parameter. Narrow locally instead of + # reporting an unfiltered list as a search result: every KB has + # already been fetched above, so the match set stays complete. + all_items = [ + item + for item in all_items + if _kb_matches_keyword(item, normalized_keyword) + ] + result_count = len(all_items) + logger.info( + "AIDP returned the full catalog (%d of %d reported KBs) for keyword %r; " + "applied the local name/description filter -> %d match(es)", + accumulated_count, + total_count, + normalized_keyword, + result_count, + ) list_ms = (time.perf_counter() - list_started_at) * 1000 logger.info( "AIDP KB catalog timing: total_ms=%.1f count_ms=%.1f list_ms=%.1f " @@ -460,7 +697,7 @@ def fetch_all_aidp_knowledge_bases_impl( return { "value": all_items, - "total_count": accumulated_count, + "total_count": result_count, "next_link": None, } except httpx.RequestError as e: @@ -1012,6 +1249,107 @@ def upload_aidp_docs_impl( ) +def remove_aidp_docs_impl( + server_url: str, + api_key: str, + kds_id: str, + file_uuids: List[str], +) -> Dict[str, Any]: + """Remove one or more documents from an AIDP knowledge base.""" + normalized_url = _validate_params(server_url, api_key) + + headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + } + remove_path = f"{_get_list_path()}/{kds_id}/KnowledgeFiles/Remove" + remove_url = urljoin(f"{normalized_url}/", remove_path) + logger.info("Removing %d AIDP documents from %s", len(file_uuids), remove_url) + + try: + client = http_client_manager.get_sync_client( + base_url=normalized_url, + timeout=_AIDP_READ_TIMEOUT_SECONDS, + verify_ssl=False, + ) + response = _request_with_retry( + lambda: client.post( + remove_url, + headers=headers, + json={"file_uuids": file_uuids}, + ), + context=f"remove-docs:{kds_id}", + ) + response.raise_for_status() + result = response.json() + return result + except httpx.RequestError as e: + logger.exception("AIDP document removal request failed: %s", e) + raise AppException( + ErrorCode.AIDP_CONNECTION_ERROR, + f"AIDP API request failed: {str(e)}", + ) + except httpx.HTTPStatusError as e: + _raise_aidp_http_error(e, "document removal") + except ValueError as e: + logger.exception("Failed to parse AIDP document removal response: %s", e) + raise AppException( + ErrorCode.AIDP_RESPONSE_ERROR, + f"Failed to parse AIDP API response: {str(e)}", + ) + + +async def stream_aidp_doc_impl( + server_url: str, + api_key: str, + kds_id: str, + file_uuid: str, +) -> httpx.Response: + """Open a streaming response for one AIDP document.""" + normalized_url = _validate_params(server_url, api_key) + + headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + } + download_path = f"{_get_list_path()}/{kds_id}/KnowledgeFiles/Download" + download_url = urljoin(f"{normalized_url}/", download_path) + logger.info("Downloading AIDP document %s from %s", file_uuid, download_url) + + response: httpx.Response | None = None + try: + client = http_client_manager.get_async_client( + base_url=normalized_url, + timeout=_AIDP_DOWNLOAD_TIMEOUT_SECONDS, + verify_ssl=False, + ) + response = await client.send( + client.build_request( + "POST", + download_url, + headers=headers, + json={"file_uuid": file_uuid}, + ), + stream=True, + ) + if response.status_code >= 400: + await response.aread() + response.raise_for_status() + return response + except httpx.RequestError as e: + if response is not None: + await response.aclose() + logger.exception("AIDP document download request failed: %s", e) + raise AppException( + ErrorCode.AIDP_CONNECTION_ERROR, + f"AIDP API request failed: {str(e)}", + ) + except httpx.HTTPStatusError as e: + if response is not None: + await response.aclose() + _raise_aidp_http_error(e, "document download") + + def count_aidp_docs_impl(server_url: str, api_key: str, kds_id: str) -> int: """Get total document count in a KB via AIDP POST .../Count endpoint. @@ -1172,6 +1510,332 @@ def list_aidp_docs_impl( ) +# ==================== Knowledge-file history (all-status listing) ==================== + +# Channel entries are read through alias lists because AIDP deployments spell +# these fields differently. Only the forms actually observed are listed; adding +# a new spelling here is enough to support another deployment. +_CHANNEL_FS_ID_KEYS = ( + "fs_id", + "fsId", + "fsid", + "file_system_id", + "fileSystemId", + "file_sys_id", + "filesystem_id", +) +_CHANNEL_SRC_DIR_KEYS = ( + "src_dir", + "srcDir", + "source_dir", + "sourceDir", + "source_path", + "sourcePath", + "dir_path", + "dirPath", + "input_dir", + "inputDir", +) +_CHANNEL_KB_ID_KEYS = ( + "kds_id", + "kdsId", + "knowledge_base_id", + "knowledgeBaseId", + "kb_id", + "kbId", + "knowledge_base_ids", + "knowledgeBaseIds", +) + + +def _first_non_empty_string(item: Dict[str, Any], keys: tuple) -> str | None: + """Return the first non-empty string among ``keys``, or ``None``.""" + for key in keys: + value = item.get(key) + if isinstance(value, str) and value.strip(): + return value.strip() + return None + + +def _lookup_channel_field(item: Dict[str, Any], keys: tuple) -> str | None: + """Return the first non-empty string among ``keys``, nested dicts included. + + Deployments nest the addressing fields differently (for example under a + ``config`` or ``file_system`` object next to the channel's own keys), so one + level of nested dictionaries is inspected before giving up. Without this a + channel holding both values could look unusable purely because of nesting, + which silently drops the whole status feature back to the ingested-files + listing. + """ + found = _first_non_empty_string(item, keys) + if found is not None: + return found + for value in item.values(): + if isinstance(value, dict): + found = _first_non_empty_string(value, keys) + if found is not None: + return found + return None + + +def _references_kb_id(value: Any, kds_id: str) -> bool: + """Return whether a channel field references ``kds_id``. + + Accepts a scalar id, a list of ids, or a directory string such as + ``/knowledge/`` — a channel is bound to a knowledge base either by + an explicit id field or by pointing its source directory at that KB. + """ + if isinstance(value, str): + return value == kds_id or kds_id in value.split("/") + if isinstance(value, (list, tuple, set)): + return any(_references_kb_id(entry, kds_id) for entry in value) + return False + + +def _channel_references_kb_id(item: Dict[str, Any], kds_id: str) -> bool: + """Return whether any channel field (nested ones included) references the KB. + + Fields are matched by name first, then every remaining scalar/list value is + checked so a build that binds a channel to its knowledge base through an + undocumented field (a pipeline name, a path) still resolves. + """ + for key, value in item.items(): + if key in _CHANNEL_KB_ID_KEYS and _references_kb_id(value, kds_id): + return True + if isinstance(value, dict) and _channel_references_kb_id(value, kds_id): + return True + if isinstance(value, (str, list, tuple, set)) and _references_kb_id( + value, kds_id + ): + return True + return False + + +def select_aidp_channel( + channels: List[Dict[str, Any]], + kds_id: str | None = None, +) -> Dict[str, str] | None: + """Pick the ingestion channel whose files feed ``kds_id``. + + Selection order: + 1. a channel that references ``kds_id`` (id field or source directory), + 2. otherwise the first channel carrying both ``fs_id`` and a source dir. + + Step 2 covers deployments that keep a single channel per tenant, where the + history lookup is intentionally directory-wide. Returns ``None`` when no + channel carries the pair the history API needs, so the caller can fall back + to the completed-files listing. + """ + usable: List[Dict[str, str]] = [] + for channel in channels: + if not isinstance(channel, dict): + continue + fs_id = _lookup_channel_field(channel, _CHANNEL_FS_ID_KEYS) + src_dir = _lookup_channel_field(channel, _CHANNEL_SRC_DIR_KEYS) + if not fs_id or not src_dir: + continue + usable.append({"fs_id": fs_id, "src_dir": src_dir}) + if kds_id and _channel_references_kb_id(channel, kds_id): + return usable[-1] + + return usable[0] if usable else None + + +def list_aidp_channels_impl( + server_url: str, + api_key: str, + kds_id: str, + tenant_id: str | None = None, +) -> Dict[str, Any]: + """List the ingestion pipelines feeding one knowledge base via AIDP API. + + Endpoint: ``GET /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{kds_id}/Channels`` + Response: ``{"value": [{"fs_id": ..., "src_dir": ...}, ...]}`` + """ + normalized_url = _validate_params(server_url, api_key) + if not isinstance(kds_id, str) or not kds_id.strip(): + raise AppException( + ErrorCode.AIDP_CONFIG_INVALID, + "AIDP channel listing requires a non-empty knowledge base id", + ) + + headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + } + + channels_path = _get_channels_path(kds_id, tenant_id) + channels_url = urljoin(f"{normalized_url}/", channels_path) + logger.info("Listing AIDP channels from %s", channels_url) + + try: + client = http_client_manager.get_sync_client( + base_url=normalized_url, + timeout=_AIDP_READ_TIMEOUT_SECONDS, + verify_ssl=False, + ) + response = _request_with_retry( + lambda: client.get(channels_url, headers=headers), + context="list-channels", + ) + response.raise_for_status() + result = response.json() + items = _extract_list_payload(result) + if items is None: + raise AppException( + ErrorCode.AIDP_RESPONSE_ERROR, + "Unexpected AIDP channel list response format", + ) + # Downstream readers always read ``value``, so the extracted items are + # re-emitted under that canonical key whichever envelope this build used. + payload = result if isinstance(result, dict) else {} + payload["value"] = [item for item in items if isinstance(item, dict)] + return payload + except httpx.RequestError as e: + logger.exception("AIDP request failed: %s", e) + raise AppException( + ErrorCode.AIDP_CONNECTION_ERROR, + f"AIDP API request failed: {str(e)}", + ) + except httpx.HTTPStatusError as e: + logger.exception( + "AIDP API HTTP error: %s, status_code: %s", + e, + e.response.status_code, + ) + if e.response.status_code in (401, 403): + raise AppException( + ErrorCode.AIDP_AUTH_ERROR, + f"AIDP authentication failed: {str(e)}", + ) + if e.response.status_code == 429: + raise AppException( + ErrorCode.AIDP_RATE_LIMIT, + f"AIDP rate limit exceeded: {str(e)}", + ) + raise AppException( + ErrorCode.AIDP_SERVICE_ERROR, + f"AIDP API HTTP error {e.response.status_code}: {str(e)}", + ) + except ValueError as e: + logger.exception("Failed to parse AIDP API response: %s", e) + raise AppException( + ErrorCode.AIDP_RESPONSE_ERROR, + f"Failed to parse AIDP API response: {str(e)}", + ) + + +def list_aidp_doc_history_impl( + server_url: str, + api_key: str, + fs_id: str, + dir_path: str, + kds_id: str, + tenant_id: str | None = None, +) -> Dict[str, Any]: + """List every file in a channel directory regardless of processing status. + + Endpoint: ``POST /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{kds_id}/KnowledgeFiles/History`` + Body: ``{"fs_id": , "dir_path": }`` + Response: ``{"value": [, ...]}`` + + Unlike ``list_aidp_docs_impl`` this returns files that are still being + chunked/embedded (``PROCESSING``) or that failed (``FAILED``), which is what + lets the UI show an upload immediately instead of only after ingestion. + """ + normalized_url = _validate_params(server_url, api_key) + + if not isinstance(kds_id, str) or not kds_id.strip(): + raise AppException( + ErrorCode.AIDP_CONFIG_INVALID, + "AIDP file history requires a non-empty knowledge base id", + ) + normalized_fs_id = fs_id.strip() if isinstance(fs_id, str) else "" + normalized_dir_path = dir_path.strip() if isinstance(dir_path, str) else "" + if not normalized_fs_id or not normalized_dir_path: + raise AppException( + ErrorCode.AIDP_CONFIG_INVALID, + "AIDP file history requires a non-empty fs_id and dir_path", + ) + + headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + } + + history_path = _get_doc_history_path(kds_id, tenant_id) + history_url = urljoin(f"{normalized_url}/", history_path) + logger.info( + "Listing AIDP knowledge-file history from %s (fs_id=%s, dir_path=%s)", + history_url, + normalized_fs_id, + normalized_dir_path, + ) + + try: + client = http_client_manager.get_sync_client( + base_url=normalized_url, + timeout=_AIDP_READ_TIMEOUT_SECONDS, + verify_ssl=False, + ) + response = _request_with_retry( + lambda: client.post( + history_url, + headers=headers, + json={"fs_id": normalized_fs_id, "dir_path": normalized_dir_path}, + ), + context=f"list-doc-history:{normalized_fs_id}", + ) + response.raise_for_status() + result = response.json() + items = _extract_list_payload(result) + if items is None: + raise AppException( + ErrorCode.AIDP_RESPONSE_ERROR, + "Unexpected AIDP file history response format", + ) + normalized_items = [ + _normalize_history_doc(item) if isinstance(item, dict) else item + for item in items + ] + _warn_when_history_status_unreadable(normalized_items, normalized_fs_id) + payload = result if isinstance(result, dict) else {} + payload["value"] = normalized_items + return payload + except httpx.RequestError as e: + logger.exception("AIDP request failed: %s", e) + raise AppException( + ErrorCode.AIDP_CONNECTION_ERROR, + f"AIDP API request failed: {str(e)}", + ) + except httpx.HTTPStatusError as e: + logger.exception( + "AIDP API HTTP error: %s, status_code: %s", + e, + e.response.status_code, + ) + if e.response.status_code in (401, 403): + raise AppException( + ErrorCode.AIDP_AUTH_ERROR, + f"AIDP authentication failed: {str(e)}", + ) + if e.response.status_code == 429: + raise AppException( + ErrorCode.AIDP_RATE_LIMIT, + f"AIDP rate limit exceeded: {str(e)}", + ) + raise AppException( + ErrorCode.AIDP_SERVICE_ERROR, + f"AIDP API HTTP error {e.response.status_code}: {str(e)}", + ) + except ValueError as e: + logger.exception("Failed to parse AIDP API response: %s", e) + raise AppException( + ErrorCode.AIDP_RESPONSE_ERROR, + f"Failed to parse AIDP API response: {str(e)}", + ) + + # AIDP ModelService endpoint for listing applicable models. def _get_models_path(tenant_id: str | None = None) -> str: """Build the tenant-scoped model service API path.""" diff --git a/frontend/const/knowledgeBase.ts b/frontend/const/knowledgeBase.ts index 0569f4660..e6a7bd853 100644 --- a/frontend/const/knowledgeBase.ts +++ b/frontend/const/knowledgeBase.ts @@ -4,6 +4,11 @@ export const AIDP_KNOWLEDGE_BASE_NAME_PATTERN = /^[\u4e00-\u9fa5a-zA-Z][\u4e00-\u9fa5a-zA-Z0-9_]{0,255}$/; +// Debounce for the AIDP knowledge-base search box. Keystrokes are coalesced +// until the user pauses for this long, so typing "report" issues one request +// instead of six. +export const KB_SEARCH_DEBOUNCE_MS = 300; + // Document status constants export const DOCUMENT_STATUS = { WAIT_FOR_PROCESSING: "WAIT_FOR_PROCESSING", diff --git a/frontend/ext_components/aidp/components/AidpCreateKbModal.tsx b/frontend/ext_components/aidp/components/AidpCreateKbModal.tsx index 5a14363ad..5e4eb93ed 100644 --- a/frontend/ext_components/aidp/components/AidpCreateKbModal.tsx +++ b/frontend/ext_components/aidp/components/AidpCreateKbModal.tsx @@ -1,69 +1,55 @@ "use client"; -import React, { useState, useMemo, useEffect, useRef } from "react"; +import React, { useEffect, useMemo, useRef, useState } from "react"; import { useTranslation } from "react-i18next"; import { useQuery } from "@tanstack/react-query"; +import type { TFunction } from "i18next"; import { - Modal, + Collapse, Form, - Input, InputNumber, - Steps, - Upload, - Button, - message, + Modal, + Select, Space, - Divider, - Collapse, Switch, Tooltip, - Select, + Upload, + message, } from "antd"; -import { InboxOutlined, QuestionCircleOutlined } from "@ant-design/icons"; +import { + InboxOutlined, + QuestionCircleOutlined, + SettingOutlined, +} from "@ant-design/icons"; import type { AidpKnowledgeBaseItem } from "@/types/agentConfig"; -import type { AidpModelItem } from "@/ext_components/aidp/services/aidpKnowledgeService"; +import type { + AidpModelItem, + AidpUploadResponse, +} from "@/ext_components/aidp/services/aidpKnowledgeService"; import aidpKnowledgeService from "@/ext_components/aidp/services/aidpKnowledgeService"; -import { USER_ROLES } from "@/const/auth"; - -/** - * Antd's Upload component (Dragger) requires ``originFileObj`` to satisfy - * the ``RcFile`` shape (``File`` + ``uid`` + ``lastModifiedDate``). We - * store raw ``File`` objects in component state, so we cast at the - * render boundary. The structural cast is sufficient because antd does - * not read the extra fields — it only requires them to exist for type - * compatibility. - */ -type RcFileLike = File & { uid: string; lastModifiedDate: Date }; -import { - AIDP_ACCEPT_STRING, - AIDP_KNOWLEDGE_BASE_NAME_PATTERN, -} from "@/const/knowledgeBase"; +import { AIDP_ACCEPT_STRING } from "@/const/knowledgeBase"; import { partitionAidpFiles, validateAidpFiles, } from "@/services/uploadService"; -import { useGroupList } from "@/hooks/group/useGroupList"; -import { useAuthorizationContext } from "@/components/providers/AuthorizationProvider"; +import { getAidpUploadFailureDetails } from "@/ext_components/aidp/services/aidpUploadUtils"; +import { collectUploadedFileIds } from "@/lib/aidpDocumentStatus"; +import { useAidpGroupOptions } from "../hooks/useAidpGroupOptions"; +import { + AIDP_MODAL_STYLES, + AidpKnowledgeBaseBasicFields, + AidpKnowledgeBaseModalFooter, + AidpKnowledgeBaseModalHeader, + AidpKnowledgeBasePermissionFields, +} from "./AidpKnowledgeBaseModalParts"; const { Dragger } = Upload; -// Preferred VLM model name when present in AIDP's available list. -// Falls back to the first model in the list if this specific one is absent. const PREFERRED_VLM_MODEL = "Qwen3-VL-8B-Instruct"; +type AidpPermission = "EDIT" | "READ_ONLY" | "PRIVATE"; -/** - * Default AIDP knowledge base configuration. - * Aligned with sdk/nexent/core/knowledge_base/config.py (build_create_payload defaults). - * - * Required fields per AIDP schema: - * chunk_token_num (> 0), chunk_overlap_num (>= 0) - * Reference fills the rest (is_personal, topk, similarity, smartsplit, caption_enable). - * ``vlm_model`` is no longer a hardcoded constant — it is resolved at runtime - * from the list of models AIDP advertises as applicable to the KnowledgeBase - * application (see ``useQuery(["aidp-models"])``). - */ const AIDP_CREATE_DEFAULTS = { chunk_token_num: 1024, chunk_overlap_num: 128, @@ -72,15 +58,116 @@ const AIDP_CREATE_DEFAULTS = { topk: 10, similarity: 0.0, smartsplit: 1, - // caption_enable: int 0/1. caption_enable: 0, }; +const validateAidpCreateFiles = ( + files: File[], + t: TFunction, + notify: typeof message +) => { + if (files.length === 0) return true; + const validation = validateAidpFiles(files); + if (validation.valid.length === files.length) return true; + partitionAidpFiles(files, t, notify); + return false; +}; + +const getAidpCreatePermissionValues = ( + isUser: boolean, + values: { ingroup_permission?: string; group_ids?: unknown } +) => { + const configuredPermission = values.ingroup_permission; + let permission: AidpPermission = "READ_ONLY"; + if (isUser) { + permission = "PRIVATE"; + } else if ( + configuredPermission === "EDIT" || + configuredPermission === "READ_ONLY" || + configuredPermission === "PRIVATE" + ) { + permission = configuredPermission; + } + const groupIds = + isUser || permission === "PRIVATE" + ? [] + : Array.isArray(values.group_ids) + ? values.group_ids + : []; + return { permission, groupIds }; +}; + +const showAidpCreateUploadResult = ( + result: AidpUploadResponse, + language: string, + t: TFunction +) => { + const failureDetails = getAidpUploadFailureDetails( + result.failed_list, + language, + t("aidpKnowledge.uploadFailed") + ); + const failureLines = failureDetails.map((detail, index) => ( +
{detail}
+ )); + const allFailedContent = + failureLines.length > 0 ? ( + failureLines + ) : ( +
{t("aidpKnowledge.uploadFailed")}
+ ); + + if (result.summary.failed > 0 && result.summary.success === 0) { + message.warning( +
+
{t("aidpKnowledge.createKbSuccess")}
+ {allFailedContent} +
+ ); + return; + } + if (result.summary.failed > 0) { + message.info( +
+
{t("aidpKnowledge.createKbSuccess")}
+
+ {t("aidpKnowledge.uploadPartial", { + success: result.summary.success, + failed: result.summary.failed, + })} +
+ {failureLines} +
+ ); + return; + } + message.success( + `${t("aidpKnowledge.createKbSuccess")} | ${t( + "aidpKnowledge.uploadSuccess", + { count: result.summary.success } + )}` + ); +}; + +const getAidpCreateErrorReason = ( + error: unknown, + knowledgeBaseCreated: boolean, + t: TFunction +) => { + if (error instanceof Error && error.message.trim()) return error.message; + return knowledgeBaseCreated + ? t("aidpKnowledge.uploadFailed") + : t("aidpKnowledge.createKbFailed"); +}; + interface AidpCreateKbModalProps { open: boolean; existingKbs: AidpKnowledgeBaseItem[]; onCancel: () => void; - onSuccess: (knowledgeBase: AidpKnowledgeBaseItem) => void; + onSuccess: ( + knowledgeBase: AidpKnowledgeBaseItem, + uploadedFileIds?: string[] + ) => void; } const AidpCreateKbModal: React.FC = ({ @@ -91,140 +178,79 @@ const AidpCreateKbModal: React.FC = ({ }) => { const { t, i18n } = useTranslation(); const [form] = Form.useForm(); - const [current, setCurrent] = useState(0); const [loading, setLoading] = useState(false); const [fileList, setFileList] = useState([]); const fileListRef = useRef([]); + const pendingFilesRef = useRef([]); + const rafIdRef = useRef(null); useEffect(() => { fileListRef.current = fileList; }, [fileList]); - // Antd fires beforeUpload once per file in a multi-select batch. - // The `newFiles` array may-or-may-not be the same reference across the N - // calls (behavior differs between and and antd versions), - // so we cannot rely on reference-equality for a single-call-per-batch guard. - // Instead, we collect each file in beforeUpload and schedule a single - // requestAnimationFrame flush that runs validate/add once per batch. - // This guarantees partitionAidpFiles + toast execute exactly ONCE per - // user selection, regardless of antd's internal dispatch count. - const pendingFilesRef = useRef([]); - const rafIdRef = useRef(null); + const { isUser, canConfigureGroupPermissions, groupOptions } = + useAidpGroupOptions(); - // Load the tenant's groups so the user can pick which groups may access - // the new KB. When no tenant context is available we fall back to an - // empty list and disable the group picker. - // NOTE: ``useAuthorizationContext`` is required — it is the only context - // that exposes ``user: User | null`` (with ``tenantId``). The similarly- - // named ``useAuthenticationContext`` only carries ``session`` and the - // plain ``useAuthentication`` hook doesn't carry ``user`` at all. - const { user } = useAuthorizationContext(); - const isUser = user?.role === USER_ROLES.USER; - const canConfigureGroupPermissions = !!user && !isUser; - const tenantId = user?.tenantId ?? null; - const { data: groupListData } = useGroupList( - canConfigureGroupPermissions ? tenantId : null - ); - const groupOptions = useMemo( - () => - (groupListData?.groups ?? []).map((g) => ({ - value: g.group_id, - label: g.group_name, - })), - [groupListData] - ); - const [formValues, setFormValues] = useState<{ - name: string; - description?: string; - vlm_model?: string; - chunk_token_num: number; - chunk_overlap_num: number; - caption_enable: number; - ingroup_permission: "EDIT" | "READ_ONLY" | "PRIVATE"; - group_ids: number[]; - }>({ - name: "", - chunk_token_num: AIDP_CREATE_DEFAULTS.chunk_token_num, - chunk_overlap_num: AIDP_CREATE_DEFAULTS.chunk_overlap_num, - caption_enable: AIDP_CREATE_DEFAULTS.caption_enable, - ingroup_permission: isUser ? "PRIVATE" : "READ_ONLY", - group_ids: [], - }); - - // Drive the vlm_model dropdown's visibility off the live Switch value. - // useWatch gives us a re-render whenever caption_enable toggles, without - // forcing the user to manually sync the form value to local state. const captionEnabled = Form.useWatch("caption_enable", form); - // Track in-group permission live so the group_ids picker can be disabled - // at PRIVATE without calling Form.useWatch inside a conditional sub-render - // (which would violate the Rules of Hooks). const ingroupPermission = Form.useWatch("ingroup_permission", form); useEffect(() => { if (!open) return; form.setFieldsValue({ + chunk_token_num: AIDP_CREATE_DEFAULTS.chunk_token_num, + chunk_overlap_num: AIDP_CREATE_DEFAULTS.chunk_overlap_num, + caption_enable: AIDP_CREATE_DEFAULTS.caption_enable === 1, ingroup_permission: isUser ? "PRIVATE" : "READ_ONLY", group_ids: [], }); - setFormValues((previous) => ({ - ...previous, - ingroup_permission: isUser ? "PRIVATE" : previous.ingroup_permission, - group_ids: isUser ? [] : previous.group_ids, - })); }, [form, isUser, open]); - // Fetch applicable VLM models from AIDP. Only run when modal is open to - // avoid hitting the (relatively slow) admin endpoint unnecessarily. const { data: vlmModelsData, isLoading: vlmModelsLoading } = useQuery({ queryKey: ["aidp-models", "llm", "KnowledgeBase"], queryFn: () => aidpKnowledgeService.listModels("llm", "KnowledgeBase"), enabled: open, - staleTime: 5 * 60 * 1000, // 5 min + staleTime: 5 * 60 * 1000, }); const vlmModelOptions = useMemo(() => { const models: AidpModelItem[] = vlmModelsData?.models ?? []; return models - .map((m) => m.model_name) - .filter( - (name): name is string => typeof name === "string" && name.length > 0 - ); + .map((model) => model.model_name) + .filter((name): name is string => Boolean(name)); }, [vlmModelsData]); - // Resolve the default VLM model: prefer the hardcoded - // PREFERRED_VLM_MODEL if present, otherwise the first in the list. - // Falls back to PREFERRED_VLM_MODEL (sent to AIDP as-is) when the - // models endpoint returns empty, matching the previous behavior. const defaultVlmModel = useMemo(() => { if (vlmModelOptions.length === 0) return PREFERRED_VLM_MODEL; - if (vlmModelOptions.includes(PREFERRED_VLM_MODEL)) - return PREFERRED_VLM_MODEL; - return vlmModelOptions[0]; + return vlmModelOptions.includes(PREFERRED_VLM_MODEL) + ? PREFERRED_VLM_MODEL + : vlmModelOptions[0]; }, [vlmModelOptions]); - // Pre-populate vlm_model on the form whenever the default is resolved, - // so the user sees a meaningful default on first open. useEffect(() => { - if (!open) return; - if (!defaultVlmModel) return; - const current = form.getFieldValue("vlm_model"); - if (!current || !vlmModelOptions.includes(current)) { + if (!open || !defaultVlmModel) return; + const currentModel = form.getFieldValue("vlm_model"); + if (!currentModel || !vlmModelOptions.includes(currentModel)) { form.setFieldValue("vlm_model", defaultVlmModel); } - }, [open, defaultVlmModel, vlmModelOptions, form]); + }, [defaultVlmModel, form, open, vlmModelOptions]); - // Duplicate name check against existing KBs const existingNames = useMemo( () => new Set( (existingKbs || []) .map((kb) => kb.kds_name?.toLowerCase().trim()) - .filter((n): n is string => !!n) + .filter((name): name is string => Boolean(name)) ), [existingKbs] ); - const handleNext = async () => { + const handleSubmit = async () => { + let knowledgeBaseCreated = false; + // Files AIDP accepted while creating the KB. Reported to the parent so it + // can keep refreshing the document list until they finish processing. + let uploadedFileIds: string[] = []; + let createdKnowledgeBase: AidpKnowledgeBaseItem | null = null; + try { const values = await form.validateFields(); const name = values.name.trim(); @@ -234,169 +260,63 @@ const AidpCreateKbModal: React.FC = ({ return; } - // Save form values before fields unmount - setFormValues({ + if (!validateAidpCreateFiles(fileList, t, message)) return; + + setLoading(true); + const { permission, groupIds } = getAidpCreatePermissionValues( + isUser, + values + ); + const captionEnable = values.caption_enable ? 1 : 0; + + const created = await aidpKnowledgeService.createKb({ name, - description: values.description?.trim() || undefined, - vlm_model: values.caption_enable - ? values.vlm_model || defaultVlmModel || undefined - : "", + description: values.description?.trim() || "", chunk_token_num: values.chunk_token_num ?? AIDP_CREATE_DEFAULTS.chunk_token_num, chunk_overlap_num: values.chunk_overlap_num ?? AIDP_CREATE_DEFAULTS.chunk_overlap_num, - caption_enable: values.caption_enable ? 1 : 0, - // The permission select is disabled at PRIVATE so users cannot pick - // group_ids while PRIVATE; we always coerce to [] for safety. - ingroup_permission: isUser - ? "PRIVATE" - : (values.ingroup_permission ?? "READ_ONLY"), - group_ids: - isUser || (values.ingroup_permission ?? "READ_ONLY") === "PRIVATE" - ? [] - : Array.isArray(values.group_ids) - ? values.group_ids - : [], - }); - setCurrent(1); - } catch { - // form validation error, do nothing - } - }; - - const handleBack = () => { - // Restore formValues into the Form when remounting Step 0, - // since antd Form clears field values when the Form is unmounted. - form.setFieldsValue(formValues); - setCurrent(0); - }; - - const handleSubmit = async (skipUpload: boolean) => { - let knowledgeBaseCreated = false; - let createdKdsId = ""; - let createdKnowledgeBase: AidpKnowledgeBaseItem | null = null; - try { - if (!formValues.name?.trim()) { - message.error(t("aidpKnowledge.kbNameRequired")); - setCurrent(0); - return; - } - setLoading(true); - - const permission = isUser ? "PRIVATE" : formValues.ingroup_permission; - const groupIds = isUser ? [] : formValues.group_ids; - - // Defense-in-depth: re-validate every file in case beforeUpload was bypassed - if (!skipUpload && fileList.length > 0) { - const validation = validateAidpFiles(fileList); - if (validation.valid.length !== fileList.length) { - setLoading(false); - partitionAidpFiles(fileList, t, message); - return; - } - } - - // Step 1: Create KB - // Aligned with sdk/nexent/core/knowledge_base/mapper.py#build_create_payload - const created = await aidpKnowledgeService.createKb({ - name: formValues.name.trim(), - description: formValues.description || "", - chunk_token_num: formValues.chunk_token_num, - chunk_overlap_num: formValues.chunk_overlap_num, embedding_model: AIDP_CREATE_DEFAULTS.embedding_model, - vlm_model: - formValues.caption_enable === 1 - ? formValues.vlm_model || defaultVlmModel || "" - : "", + vlm_model: captionEnable + ? values.vlm_model || defaultVlmModel || "" + : "", is_personal: AIDP_CREATE_DEFAULTS.is_personal, topk: AIDP_CREATE_DEFAULTS.topk, similarity: AIDP_CREATE_DEFAULTS.similarity, smartsplit: AIDP_CREATE_DEFAULTS.smartsplit, - caption_enable: formValues.caption_enable, - // v7.1: forward in-group permission + groups to the backend so the - // knowledge-base permission row is created in lockstep with the KB. + caption_enable: captionEnable, ingroup_permission: permission, group_ids: groupIds, }); knowledgeBaseCreated = true; - createdKdsId = String(created.kds_id || ""); createdKnowledgeBase = { ...created, - kds_id: createdKdsId, - kds_name: created.kds_name || formValues.name.trim(), - description: created.description ?? formValues.description ?? "", + kds_id: String(created.kds_id || ""), + kds_name: created.kds_name || name, + description: created.description ?? values.description?.trim() ?? "", permission: "EDIT", ingroup_permission: permission, - group_ids: permission === "PRIVATE" ? [] : groupIds, + group_ids: groupIds, resource_status: "ACTIVE", - is_multimodal: formValues.caption_enable === 1, + is_multimodal: captionEnable === 1, }; - // Step 2: Upload files (if any and not skipped) - if (!skipUpload && fileList.length > 0 && created.kds_id) { + if (fileList.length > 0 && created.kds_id) { const result = await aidpKnowledgeService.uploadDocs( created.kds_id, fileList ); - - const failureDetails = result.failed_list.map((item) => { - const reason = i18n.language.startsWith("zh") - ? item.reason_zh || item.reason_en - : item.reason_en || item.reason_zh; - return `${item.file_name}: ${reason || t("aidpKnowledge.uploadFailed")}`; - }); - const failureLines = failureDetails.map((detail, index) => ( -
{detail}
- )); - - if (result.summary.failed > 0 && result.summary.success === 0) { - message.warning( -
-
{t("aidpKnowledge.createKbSuccess")}
- {failureLines.length > 0 ? ( - failureLines - ) : ( -
{t("aidpKnowledge.uploadFailed")}
- )} -
- ); - } else if (result.summary.failed > 0) { - message.info( -
-
{t("aidpKnowledge.createKbSuccess")}
-
- {t("aidpKnowledge.uploadPartial", { - success: result.summary.success, - failed: result.summary.failed, - })} -
- {failureLines} -
- ); - } else { - message.success( - t("aidpKnowledge.createKbSuccess") + - " | " + - t("aidpKnowledge.uploadSuccess", { - count: result.summary.success, - }) - ); - } + showAidpCreateUploadResult(result, i18n.language, t); + uploadedFileIds = collectUploadedFileIds(result.success_list); } else { message.success(t("aidpKnowledge.createKbSuccess")); } handleReset(); - if (createdKnowledgeBase) { - onSuccess(createdKnowledgeBase); - } + if (createdKnowledgeBase) + onSuccess(createdKnowledgeBase, uploadedFileIds); } catch (error) { - const reason = - error instanceof Error && error.message.trim() - ? error.message - : knowledgeBaseCreated - ? t("aidpKnowledge.uploadFailed") - : t("aidpKnowledge.createKbFailed"); + const reason = getAidpCreateErrorReason(error, knowledgeBaseCreated, t); message.error( knowledgeBaseCreated ? `${t("aidpKnowledge.createKbSuccess")} | ${reason}` @@ -404,9 +324,8 @@ const AidpCreateKbModal: React.FC = ({ ); if (knowledgeBaseCreated) { handleReset(); - if (createdKnowledgeBase) { - onSuccess(createdKnowledgeBase); - } + if (createdKnowledgeBase) + onSuccess(createdKnowledgeBase, uploadedFileIds); } } finally { setLoading(false); @@ -415,21 +334,12 @@ const AidpCreateKbModal: React.FC = ({ const handleReset = () => { form.resetFields(); - setCurrent(0); setFileList([]); pendingFilesRef.current = []; if (rafIdRef.current !== null) { cancelAnimationFrame(rafIdRef.current); rafIdRef.current = null; } - setFormValues({ - name: "", - chunk_token_num: AIDP_CREATE_DEFAULTS.chunk_token_num, - chunk_overlap_num: AIDP_CREATE_DEFAULTS.chunk_overlap_num, - caption_enable: AIDP_CREATE_DEFAULTS.caption_enable, - ingroup_permission: isUser ? "PRIVATE" : "READ_ONLY", - group_ids: [], - }); }; const handleCancel = () => { @@ -437,330 +347,253 @@ const AidpCreateKbModal: React.FC = ({ onCancel(); }; - // ---- Render steps ---- - - const renderStep0 = () => ( - <> -
- - - - - - - - {canConfigureGroupPermissions && ( - <> - {/* USER creation is always a personal PRIVATE KB. */} - - - - - )} - - - {t("aidpKnowledge.createCaptionEnable")} - - - - - } - > - - - - {/* VLM model picker is only relevant when multimodal captioning is - enabled. Hide the dropdown entirely when the Switch is off so - users aren't shown an irrelevant choice, and so the backend - receives an empty ``vlm_model`` (see handleSubmit). */} - {captionEnabled && ( - - {t("aidpKnowledge.createVlmModel")} - - - - - } - > - ({ + label: name, + value: name, + }))} + filterOption={(input, option) => + (option?.label as string) + ?.toLowerCase() + .includes(input.toLowerCase()) ?? false + } + /> + + )} + + ), + }, + ]} + /> +
+ ); }; diff --git a/frontend/ext_components/aidp/components/AidpDocumentList.tsx b/frontend/ext_components/aidp/components/AidpDocumentList.tsx index 3394e4450..8e8e7a037 100644 --- a/frontend/ext_components/aidp/components/AidpDocumentList.tsx +++ b/frontend/ext_components/aidp/components/AidpDocumentList.tsx @@ -1,9 +1,9 @@ -import React, { useState, useCallback, useRef } from "react"; +import React, { useCallback, useRef, useState } from "react"; import { useTranslation } from "react-i18next"; -import { Button, Pagination, Upload, message, Tooltip } from "antd"; +import { Button, Modal, Pagination, Tag, Upload, message, Tooltip } from "antd"; import { - UploadOutlined, + FileTextOutlined, InboxOutlined, ReloadOutlined, } from "@ant-design/icons"; @@ -12,23 +12,144 @@ import type { AidpKnowledgeBaseItem } from "@/types/agentConfig"; import type { AidpDocumentItem } from "@/ext_components/aidp/services/aidpKnowledgeService"; import aidpKnowledgeService from "@/ext_components/aidp/services/aidpKnowledgeService"; import { AIDP_ACCEPT_STRING } from "@/const/knowledgeBase"; +import log from "@/lib/logger"; +import { + AIDP_DOC_IN_PROGRESS_STATUSES, + AIDP_DOCUMENT_STATUS, + collectUploadedFileIds, + normalizeAidpDocStatus, +} from "@/lib/aidpDocumentStatus"; import { partitionAidpFiles } from "@/services/uploadService"; +import { getAidpUploadFailureDetails } from "@/ext_components/aidp/services/aidpUploadUtils"; const { Dragger } = Upload; +// AIDP rejects a re-uploaded file with a per-file reason, but that reason is +// shaped exactly like any other upload failure — the user cannot tell a +// duplicate apart from a genuine error. The wording AIDP actually returns is +// "文件已存在,请重命名或删除已有文件" / "File already exists. Please rename or +// delete the existing file.", which never contains the literal word +// "duplicate". Match on those phrases in BOTH languages (the backend returns +// reason_zh and reason_en together, independently of the UI language) plus the +// generic duplicate markers, case-insensitively so upstream capitalisation +// changes cannot silently break the detection. +const DUPLICATE_UPLOAD_REASON_MARKERS = [ + "already exists", + "duplicate", + "已存在", + "重复", +]; + +const isDuplicateUploadReason = ( + ...reasons: Array +): boolean => { + const haystack = reasons + .filter((reason): reason is string => Boolean(reason)) + .join(" ") + .toLowerCase(); + if (!haystack) return false; + return DUPLICATE_UPLOAD_REASON_MARKERS.some((marker) => + haystack.includes(marker) + ); +}; + +/** + * Labels for the in-progress statuses, which all render as a blue tag. + * + * Uploading and extracting are the states a user watches right after an upload, + * so they get the same treatment as processing instead of falling through to the + * verbatim fallback below. + */ +const IN_PROGRESS_STATUS_LABELS: Record = { + [AIDP_DOCUMENT_STATUS.UPLOADING]: "aidpKnowledge.docStatusUploading", + [AIDP_DOCUMENT_STATUS.PROCESSING]: "aidpKnowledge.docStatusProcessing", + [AIDP_DOCUMENT_STATUS.EXTRACTING]: "aidpKnowledge.docStatusExtracting", +}; + +/** + * Table cell showing the ingestion status of a document. + * + * Uploading, processing and extracting are blue because the file is still on its + * way in (being uploaded, chunked/embedded, or having its content extracted) and + * the list keeps refreshing itself until the status resolves; `COMPLETED` green + * and `FAILED` red are the two terminal outcomes. A missing status means the + * backend fell back to the completed-files listing, which only reports ingested + * files — those render as a dash like any other empty cell. An unrecognised + * status is shown verbatim rather than hidden, so a new AIDP status is visible + * instead of silently blank. + */ +const DocumentStatusCell: React.FC<{ status?: string }> = ({ status }) => { + const { t } = useTranslation(); + const normalized = normalizeAidpDocStatus(status); + + if (!normalized) { + return -; + } + + if (AIDP_DOC_IN_PROGRESS_STATUSES.includes(normalized)) { + const labelKey = IN_PROGRESS_STATUS_LABELS[normalized]; + return ( + + {labelKey ? t(labelKey) : status} + + ); + } + + if (normalized === AIDP_DOCUMENT_STATUS.COMPLETED) { + return ( + + {t("aidpKnowledge.docStatusCompleted")} + + ); + } + + if (normalized === AIDP_DOCUMENT_STATUS.FAILED) { + return ( + + {t("aidpKnowledge.docStatusFailed")} + + ); + } + + return ( + + {status} + + ); +}; + +const resolveDownloadFilename = (response: Response, fallback: string) => { + const contentDisposition = response.headers.get("content-disposition") || ""; + const encodedName = /filename\*=UTF-8''([^;]+)/i.exec( + contentDisposition + )?.[1]; + if (encodedName) { + try { + return decodeURIComponent(encodedName); + } catch { + // Use the regular filename or document name when decoding fails. + } + } + const plainName = /filename="?([^";]+)"?/i.exec(contentDisposition)?.[1]; + return plainName || response.headers.get("x-file-name") || fallback; +}; + interface AidpDocumentListProps { activeKb: AidpKnowledgeBaseItem | null; documents: AidpDocumentItem[]; totalDocs: number; - /** True when `totalDocs` came from the AIDP Count API; when false the - * total is a fallback estimate and "共 N 条" should be suppressed. */ + /** True when `totalDocs` came from the AIDP Count API. */ totalReliable: boolean; hasMore: boolean; isLoading: boolean; currentPage: number; pageSize: number; onPageChange: (page: number) => void; - onDocsUploaded: () => void; + /** Called after documents change, with the ids AIDP returned for any newly + * accepted uploads (an empty list for deletions/refreshes). The parent uses + * them to keep refreshing the list until each uploaded file reports a + * terminal processing status. */ + onDocsUploaded: (uploadedFileIds: string[]) => void; onRefresh: () => void; } @@ -47,6 +168,10 @@ const AidpDocumentList: React.FC = ({ }) => { const { t, i18n } = useTranslation(); const [uploading, setUploading] = useState(false); + const [deleting, setDeleting] = useState(false); + const [downloadingFileUuid, setDownloadingFileUuid] = useState( + null + ); // Antd fires beforeUpload once per file in a multi-select batch. // The `fileList` array may-or-may-not be the same reference across the N // calls (behavior differs between and and antd versions), @@ -56,10 +181,85 @@ const AidpDocumentList: React.FC = ({ const pendingFilesRef = useRef([]); const rafIdRef = useRef(null); + const isUnavailable = + activeKb?.resource_status === "UNAVAILABLE" || + activeKb?.resource_status === "ORPHANED"; + const canDeleteDocuments = + !!activeKb && !isUnavailable && activeKb.permission === "EDIT"; + const canDownloadDocuments = + !!activeKb && + !isUnavailable && + (activeKb.permission === "EDIT" || activeKb.permission === "READ_ONLY"); + + const handleDownload = useCallback( + async (document: AidpDocumentItem) => { + if (!activeKb || !document.file_uuid) return; + setDownloadingFileUuid(document.file_uuid); + try { + const response = await aidpKnowledgeService.downloadDoc( + activeKb.kds_id, + document.file_uuid + ); + const blob = await response.blob(); + const downloadUrl = URL.createObjectURL(blob); + const link = window.document.createElement("a"); + link.href = downloadUrl; + link.download = resolveDownloadFilename(response, document.file_name); + link.click(); + URL.revokeObjectURL(downloadUrl); + message.success(t("aidpKnowledge.downloadSuccess")); + } catch (error) { + log.error("Failed to download AIDP document:", error); + message.error(t("aidpKnowledge.downloadFailed")); + } finally { + setDownloadingFileUuid(null); + } + }, + [activeKb, t] + ); + + const handleDelete = useCallback( + (document: AidpDocumentItem) => { + if (!activeKb || !document.file_uuid) return; + Modal.confirm({ + title: t("aidpKnowledge.confirmDeleteDocTitle"), + content: t("aidpKnowledge.confirmDeleteDocContent"), + okText: t("common.confirm"), + cancelText: t("common.cancel"), + okButtonProps: { danger: true }, + centered: true, + onOk: async () => { + setDeleting(true); + try { + const result = await aidpKnowledgeService.removeDoc( + activeKb.kds_id, + document.file_uuid + ); + if (result.summary.success > 0) { + message.success(t("aidpKnowledge.deleteDocSuccess")); + } else { + message.error(t("aidpKnowledge.deleteDocFailed")); + } + if (result.summary.success > 0) { + // Deletions are not uploads: nothing to wait for, so the parent + // simply refreshes the list. + onDocsUploaded([]); + } + } catch (error) { + log.error("Failed to delete AIDP document:", error); + message.error(t("aidpKnowledge.deleteDocFailed")); + } finally { + setDeleting(false); + } + }, + }); + }, + [activeKb, onDocsUploaded, t] + ); + const handleUpload = useCallback( async (fileList: File[]) => { - if (!activeKb) return; - if (fileList.length === 0) return; + if (!activeKb || fileList.length === 0) return; setUploading(true); try { @@ -68,11 +268,21 @@ const AidpDocumentList: React.FC = ({ fileList ); + // A rejected duplicate is not an error the user can debug, so give it a + // dedicated message instead of echoing AIDP's "please rename or delete" + // instruction, which is not actionable in this dialog. Every other + // failure keeps the shared reason formatter. const failureDetails = result.failed_list.map((item) => { - const reason = i18n.language.startsWith("zh") - ? item.reason_zh || item.reason_en - : item.reason_en || item.reason_zh; - return `${item.file_name}: ${reason || t("aidpKnowledge.uploadFailed")}`; + if (isDuplicateUploadReason(item.reason_zh, item.reason_en)) { + return t("aidpKnowledge.uploadDuplicateFile", { + fileName: item.file_name, + }); + } + return getAidpUploadFailureDetails( + [item], + i18n.language, + t("aidpKnowledge.uploadFailed") + )[0]; }); const failureLines = failureDetails.map((detail, index) => (
{detail}
@@ -98,12 +308,12 @@ const AidpDocumentList: React.FC = ({ {failureLines} ); - onDocsUploaded(); + onDocsUploaded(collectUploadedFileIds(result.success_list)); } else { message.success( t("aidpKnowledge.uploadSuccess", { count: result.summary.success }) ); - onDocsUploaded(); + onDocsUploaded(collectUploadedFileIds(result.success_list)); } } catch (error) { const reason = @@ -118,7 +328,6 @@ const AidpDocumentList: React.FC = ({ [activeKb, i18n.language, onDocsUploaded, t] ); - // Format file size for display const formatSize = (bytes?: number): string => { if (!bytes || bytes === 0) return "-"; if (bytes < 1024) return `${bytes} B`; @@ -128,206 +337,231 @@ const AidpDocumentList: React.FC = ({ return `${(bytes / (1024 * 1024 * 1024)).toFixed(1)} GB`; }; + const effectiveTotal = totalReliable + ? totalDocs + : hasMore + ? currentPage * pageSize + 1 + : currentPage * pageSize; + + const canUpload = + !!activeKb && !isUnavailable && activeKb.permission === "EDIT"; + + const renderUploadArea = () => { + if (!canUpload) { + const reasonKey = !activeKb + ? "aidpKnowledge.uploadNoKb" + : isUnavailable + ? "aidpKnowledge.uploadKbUnavailable" + : "aidpKnowledge.uploadReadOnly"; + return ( +
+

{t(reasonKey)}

+
+ ); + } + + return ( + { + pendingFilesRef.current.push(_file); + if (rafIdRef.current === null) { + rafIdRef.current = requestAnimationFrame(() => { + const batch = pendingFilesRef.current; + pendingFilesRef.current = []; + rafIdRef.current = null; + + const { valid } = partitionAidpFiles(batch, t, message); + if (valid.length > 0) void handleUpload(valid); + }); + } + return false; + }} + disabled={uploading} + className="!rounded-xl !border-blue-200 !bg-blue-50/30" + > +

+ +

+

+ {uploading + ? t("aidpKnowledge.uploading") + : t("aidpKnowledge.uploadHint")} +

+
+
{t("aidpKnowledge.uploadHintCount")}
+
{t("aidpKnowledge.uploadHintSize")}
+
+ {t("aidpKnowledge.uploadHintFormats")} +
+
+
+ ); + }; + return ( -
- {/* Header */} -
-
-
-

- {activeKb?.kds_name || ""} -

- +
+
+
+
+ +
+
+ {/* An empty `title` leaves the Tooltip inert, so it is safe + while the knowledge base detail is still loading. */} + +

+ {activeKb?.kds_name || ""} +

+
+ {t("aidpKnowledge.tagDocs", { count: totalDocs })}
- -
+ +
- {/* Document table */} -
+
{isLoading ? ( -
+
-
-

+

+

{t("aidpKnowledge.loadingDocs")}

) : documents.length > 0 ? ( -
+
- + - - - + - + - + {documents.map((doc) => ( - - + - - - + ))}
+ {t("aidpKnowledge.docFileName")} + {t("aidpKnowledge.docType")} + + {t("aidpKnowledge.docStatus")} + {t("aidpKnowledge.docSize")} + {t("aidpKnowledge.docCreatedAt")} + {t("aidpKnowledge.docActions")} +
-
- {doc.file_name} -
-
- {doc.file_ino_no} -
+
+ {/* Long file names are clamped by CSS on a + width-bounded inner element (`max-width` on a + table cell is ignored by browsers, so the clamp + must live on the div itself) and the full value + is surfaced through an antd Tooltip on hover, + matching the local knowledge-base list. */} + +
+ {doc.file_name} +
+
+ +
+ {doc.file_ino_no} +
+
+ {doc.file_type || "-"} + + {formatSize(doc.file_size)} + {doc.created_at ? new Date(doc.created_at).toLocaleString() : "-"} +
+ {canDownloadDocuments && ( + + )} + {canDeleteDocuments && ( + + )} +
+
) : ( -
+
{t("aidpKnowledge.noDocuments")}
)}
- {/* Server-side pagination. - AIDP exposes a dedicated Count API for documents which the backend - now calls alongside the list request. When Count succeeds, - `totalReliable` is true and we display the full pagination (page - numbers + "共 N 条"). When Count fails (e.g. the endpoint is not - available on a particular AIDP instance), `totalReliable` is false - and we fall back to simple prev/next mode without a total, using - `has_more` to decide whether the next-page button should enable. */} - {documents.length > 0 && - (() => { - // When total is unreliable we still need antd to know when to - // enable "next": set total just past the current page if there is - // a next page, otherwise clamp to the current page end. - const effectiveTotal = totalReliable - ? totalDocs - : hasMore - ? currentPage * pageSize + 1 - : currentPage * pageSize; - return ( -
- t("aidpKnowledge.showTotal", { count: total }) - : undefined - } - size="small" - /> -
- ); - })()} - - {/* Upload area — gated by ``activeKb.permission`` and ``resource_status``. - - Per v7.1 §7.1, READ_ONLY callers may view existing documents but - must not be able to upload. UNAVAILABLE / ORPHANED KBs are - read-only regardless of permission because the AIDP backend cannot - service the request. The container is replaced with a hint instead - of disabling the Dragger so the visual structure stays consistent - and screen-reader users get an explicit reason. */} -
- {(() => { - const isUnavailable = - activeKb?.resource_status === "UNAVAILABLE" || - activeKb?.resource_status === "ORPHANED"; - const canUpload = - !!activeKb && !isUnavailable && activeKb.permission === "EDIT"; - if (!canUpload) { - const reasonKey = !activeKb - ? "aidpKnowledge.uploadNoKb" - : isUnavailable - ? "aidpKnowledge.uploadKbUnavailable" - : "aidpKnowledge.uploadReadOnly"; - return ( -
-

{t(reasonKey)}

-
- ); - } - return ( - { - // Queue the file and defer validation + upload until the - // synchronous batch of beforeUpload calls finishes. Each batch - // flushes in a single frame so toasts and handleUpload run once. - pendingFilesRef.current.push(_file); - if (rafIdRef.current === null) { - rafIdRef.current = requestAnimationFrame(() => { - const batch = pendingFilesRef.current; - pendingFilesRef.current = []; - rafIdRef.current = null; - - const { valid } = partitionAidpFiles(batch, t, message); - if (valid.length > 0) { - handleUpload(valid); - } - }); - } - return false; - }} - disabled={uploading} - > -

- -

-

- {uploading - ? t("aidpKnowledge.uploading") - : t("aidpKnowledge.uploadHint")} -

-
-
{t("aidpKnowledge.uploadHintCount")}
-
{t("aidpKnowledge.uploadHintSize")}
-
- {t("aidpKnowledge.uploadHintFormats")} -
-
-
- ); - })()} + {documents.length > 0 && ( +
+ t("aidpKnowledge.showTotal", { count }) + : undefined + } + size="small" + /> +
+ )} + +
+ {renderUploadArea()}
); diff --git a/frontend/ext_components/aidp/components/AidpKnowledgeBaseModalParts.tsx b/frontend/ext_components/aidp/components/AidpKnowledgeBaseModalParts.tsx new file mode 100644 index 000000000..85b2d77a8 --- /dev/null +++ b/frontend/ext_components/aidp/components/AidpKnowledgeBaseModalParts.tsx @@ -0,0 +1,162 @@ +import React from "react"; +import { Button, Form, Input, Select } from "antd"; +import type { TFunction } from "i18next"; + +import { AIDP_KNOWLEDGE_BASE_NAME_PATTERN } from "@/const/knowledgeBase"; +import type { AidpGroupOption } from "../hooks/useAidpGroupOptions"; + +export const AIDP_MODAL_STYLES = { + container: { overflow: "hidden", borderRadius: 16, padding: 0 }, + body: { padding: 0 }, + footer: { + margin: 0, + padding: "12px 20px 16px", + borderTop: "1px solid #f0f0f0", + }, +}; + +interface AidpModalHeaderProps { + title: string; + subtitle: string; +} + +export const AidpKnowledgeBaseModalHeader: React.FC = ({ + title, + subtitle, +}) => ( +
+

+ {title} +

+

{subtitle}

+
+); + +interface AidpModalFooterProps { + onCancel: () => void; + onSubmit: () => void; + loading: boolean; + cancelText: string; + submitText: string; +} + +export const AidpKnowledgeBaseModalFooter: React.FC = ({ + onCancel, + onSubmit, + loading, + cancelText, + submitText, +}) => ( +
+ + +
+); + +interface AidpKnowledgeBaseBasicFieldsProps { + t: TFunction; +} + +export const AidpKnowledgeBaseBasicFields: React.FC< + AidpKnowledgeBaseBasicFieldsProps +> = ({ t }) => ( + <> + + + + + + + +); + +interface AidpPermissionFieldsProps { + t: TFunction; + groupOptions: AidpGroupOption[]; + ingroupPermission?: string; + showSearch?: boolean; +} + +const hasRequiredGroupIds = (permission: string, value: unknown) => + permission === "PRIVATE" || (Array.isArray(value) && value.length > 0); + +export const AidpKnowledgeBasePermissionFields: React.FC< + AidpPermissionFieldsProps +> = ({ t, groupOptions, ingroupPermission, showSearch = false }) => ( + <> + + + + +); diff --git a/frontend/ext_components/aidp/components/AidpKnowledgeConfiguration.tsx b/frontend/ext_components/aidp/components/AidpKnowledgeConfiguration.tsx index 9dac6c97c..d9463c498 100644 --- a/frontend/ext_components/aidp/components/AidpKnowledgeConfiguration.tsx +++ b/frontend/ext_components/aidp/components/AidpKnowledgeConfiguration.tsx @@ -1,6 +1,6 @@ "use client"; -import React, { useState, useEffect, useCallback } from "react"; +import React, { useState, useEffect, useCallback, useRef } from "react"; import { useTranslation } from "react-i18next"; import { App, Row, Col, Modal } from "antd"; @@ -9,8 +9,13 @@ import { InfoCircleFilled } from "@ant-design/icons"; import { SETUP_PAGE_CONTAINER, TWO_COLUMN_LAYOUT, - STANDARD_CARD, } from "@/const/layoutConstants"; +import { KB_SEARCH_DEBOUNCE_MS } from "@/const/knowledgeBase"; +import { + AIDP_DOC_STATUS_POLL_MS, + AIDP_DOC_UPLOAD_WATCH_TIMEOUT_MS, + findPendingUploadIds, +} from "@/lib/aidpDocumentStatus"; import type { AidpKnowledgeBaseItem } from "@/types/agentConfig"; import aidpKnowledgeService, { type AidpKbDetail, @@ -40,12 +45,17 @@ const AidpKnowledgeConfiguration: React.FC = () => { // which may not contain the currently active KB. `selectedKb` is the item // itself — set on selection, kept stable across list refetches. const [activeKbId, setActiveKbId] = useState(null); - const [selectedKb, setSelectedKb] = useState(null); - const [activeKbDetail, setActiveKbDetail] = useState(null); + const [selectedKb, setSelectedKb] = useState( + null + ); + const [, setActiveKbDetail] = useState(null); const [documents, setDocuments] = useState([]); const [totalDocs, setTotalDocs] = useState(0); const [docHasMore, setDocHasMore] = useState(false); const [docTotalReliable, setDocTotalReliable] = useState(true); + // Files still being processed across the whole knowledge base (not just the + // visible page). Drives the status polling below. + const [docProcessingCount, setDocProcessingCount] = useState(0); const [loadingDocs, setLoadingDocs] = useState(false); // ---- Pagination state ---- @@ -57,37 +67,60 @@ const AidpKnowledgeConfiguration: React.FC = () => { // ---- Modal state ---- const [createModalOpen, setCreateModalOpen] = useState(false); const [updateModalOpen, setUpdateModalOpen] = useState(false); - const [editingKb, setEditingKb] = useState(null); + const [editingKb, setEditingKb] = useState( + null + ); + + // ---- Keyword search state ---- + // `kbKeyword` is the raw input value and keeps the text field responsive; + // `debouncedKbKeyword` is what actually drives requests. Splitting them means + // typing never waits on the network, and a pause settles on one request. + const [kbKeyword, setKbKeyword] = useState(""); + const [debouncedKbKeyword, setDebouncedKbKeyword] = useState(""); + + useEffect(() => { + const timer = setTimeout(() => { + setDebouncedKbKeyword(kbKeyword.trim()); + }, KB_SEARCH_DEBOUNCE_MS); + return () => clearTimeout(timer); + }, [kbKeyword]); // ---- Fetch KB list (server-side pagination: each page fetches page_size items + Count total) ---- - const fetchKbs = useCallback(async (page: number = 1) => { - setLoadingKbs(true); - try { - const result = await aidpKnowledgeService.listKbs( - page, - KB_PAGE_SIZE, - ); - setKbs(result.value); - setKbTotal(result.total_count ?? result.value.length); - setKbHasMore(result.has_more ?? false); - setKbTotalReliable(result.total_reliable !== false); - setKbPage(page); - } catch (error) { - log.error("Failed to fetch AIDP knowledge bases:", error); - appMessage.error(t("aidpKnowledge.fetchKbsFailed")); - setKbs([]); - setKbTotal(0); - setKbHasMore(false); - setKbTotalReliable(false); - } finally { - setLoadingKbs(false); - } - }, [appMessage, t]); + const fetchKbs = useCallback( + async (page: number = 1, keyword: string = "") => { + setLoadingKbs(true); + try { + const result = await aidpKnowledgeService.listKbs( + page, + KB_PAGE_SIZE, + keyword + ); + setKbs(result.value); + setKbTotal(result.total_count ?? result.value.length); + setKbHasMore(result.has_more ?? false); + setKbTotalReliable(result.total_reliable !== false); + setKbPage(page); + } catch (error) { + log.error("Failed to fetch AIDP knowledge bases:", error); + appMessage.error(t("aidpKnowledge.fetchKbsFailed")); + setKbs([]); + setKbTotal(0); + setKbHasMore(false); + setKbTotalReliable(false); + } finally { + setLoadingKbs(false); + } + }, + [appMessage, t] + ); - // Auto-fetch on mount + // Fetch on mount, and again whenever the debounced keyword settles. + // Every keyword change restarts at page 1 on purpose: the previous page + // number is meaningless against a different result set and would otherwise + // render an empty list whenever the filtered set is shorter than that page. useEffect(() => { - fetchKbs(); - }, [fetchKbs]); + fetchKbs(1, debouncedKbKeyword); + }, [fetchKbs, debouncedKbKeyword]); // ---- Cleanup legacy localStorage credentials on mount ---- // v7.1: AIDP credentials moved backend-side; frontends that pre-date the @@ -105,46 +138,129 @@ const AidpKnowledgeConfiguration: React.FC = () => { }, []); // ---- Fetch documents for active KB (server-side pagination) ---- + // `silent` refreshes are used by the status poller: the table keeps rendering + // the previous page instead of flashing the loading placeholder, and a failed + // poll stays out of the way (logged only) so a short upstream hiccup cannot + // spam the user with a toast every interval. The manual refresh button always + // runs a non-silent fetch, so errors stay visible when the user asks for them. const fetchDocs = useCallback( - async (kbId: string, page: number = 1) => { - setLoadingDocs(true); + async (kbId: string, page: number = 1, options?: { silent?: boolean }) => { + const silent = options?.silent === true; + if (!silent) setLoadingDocs(true); try { const result = await aidpKnowledgeService.listDocs( kbId, page, - DOC_PAGE_SIZE, + DOC_PAGE_SIZE ); const count = result.total_count ?? result.value.length; setDocuments(result.value); setTotalDocs(count); setDocHasMore(result.has_more ?? false); setDocTotalReliable(result.total_reliable !== false); + setDocProcessingCount(result.processing_count ?? 0); setDocPage(page); + + // Settle the upload watch: a just-uploaded file counts as done only + // once it is listed with a terminal status. Files AIDP has not listed + // yet stay pending on purpose, so the refresh keeps running instead of + // stopping while the list is still missing the upload. + if (pendingUploadIdsRef.current.length > 0) { + const stillPending = findPendingUploadIds( + pendingUploadIdsRef.current, + result.value + ); + pendingUploadIdsRef.current = stillPending; + if (stillPending.length === 0) setUploadWatchActive(false); + } } catch (error) { log.error("Failed to fetch AIDP documents:", error); - appMessage.error(t("aidpKnowledge.fetchDocsFailed")); - setDocuments([]); - setTotalDocs(0); - setDocHasMore(false); - setDocTotalReliable(false); + if (!silent) { + appMessage.error(t("aidpKnowledge.fetchDocsFailed")); + setDocuments([]); + setTotalDocs(0); + setDocHasMore(false); + setDocTotalReliable(false); + setDocProcessingCount(0); + } } finally { - setLoadingDocs(false); + if (!silent) setLoadingDocs(false); } }, [appMessage, t] ); + // ---- Upload watch - state - --- + // Files the user just uploaded and that are not settled yet. Kept in a ref + // so `fetchDocs` can settle them without becoming a new function on every + // watch update (which would restart the polling interval). + const pendingUploadIdsRef = useRef([]); + const uploadWatchStartedAtRef = useRef(0); + const [uploadWatchActive, setUploadWatchActive] = useState(false); + + /** Start refreshing until every just-uploaded file reaches a terminal state. */ + const startUploadWatch = useCallback((uploadedFileIds: string[]) => { + if (uploadedFileIds.length === 0) return; + pendingUploadIdsRef.current = uploadedFileIds; + uploadWatchStartedAtRef.current = Date.now(); + setUploadWatchActive(true); + }, []); + + /** Stop watching uploads (no upload in flight to wait for). */ + const stopUploadWatch = useCallback(() => { + pendingUploadIdsRef.current = []; + setUploadWatchActive(false); + }, []); + + // ---- Poll the document list while work is outstanding ---- + // Two independent reasons to poll: + // * a file just uploaded by the user has not settled yet (`uploadWatchActive`), + // which covers the window where AIDP has accepted the upload but does not + // list it yet - polling on processing_count alone would never start there; + // * the knowledge base still reports files being ingested + // (`docProcessingCount`, counted across the whole KB, so a processing file + // on another page keeps the status column live). + // Polling stops once neither holds: every file is COMPLETED or FAILED. The + // upload watch additionally gives up after AIDP_DOC_UPLOAD_WATCH_TIMEOUT_MS so + // a silently dropped upload cannot keep the list refreshing forever. + const shouldPollDocs = uploadWatchActive || docProcessingCount > 0; + useEffect(() => { + if (!activeKbId || !shouldPollDocs) return; + const tick = () => { + if ( + uploadWatchActive && + Date.now() - uploadWatchStartedAtRef.current > + AIDP_DOC_UPLOAD_WATCH_TIMEOUT_MS + ) { + stopUploadWatch(); + return; + } + void fetchDocs(activeKbId, docPage, { silent: true }); + }; + const timer = window.setInterval(tick, AIDP_DOC_STATUS_POLL_MS); + return () => window.clearInterval(timer); + }, [ + activeKbId, + docPage, + shouldPollDocs, + uploadWatchActive, + fetchDocs, + stopUploadWatch, + ]); + // ---- Handle KB selection ---- const handleSelectKb = useCallback( (kb: AidpKnowledgeBaseItem) => { + stopUploadWatch(); setActiveKbId(kb.kds_id); setSelectedKb(kb); setDocPage(1); setDocHasMore(false); setDocTotalReliable(true); + setDocProcessingCount(0); fetchDocs(kb.kds_id, 1); }, - [fetchDocs] + [fetchDocs, stopUploadWatch] ); // ---- Handle KB deletion ---- @@ -166,6 +282,7 @@ const AidpKnowledgeConfiguration: React.FC = () => { // If the deleted KB was active, clear selection if (activeKbId === kb.kds_id) { + stopUploadWatch(); setActiveKbId(null); setSelectedKb(null); setActiveKbDetail(null); @@ -173,18 +290,27 @@ const AidpKnowledgeConfiguration: React.FC = () => { setTotalDocs(0); setDocHasMore(false); setDocTotalReliable(true); + setDocProcessingCount(0); setDocPage(1); } - // Refresh list - fetchKbs(kbPage); - } catch (error) { + // Refresh list, keeping the active search filter applied + fetchKbs(kbPage, debouncedKbKeyword); + } catch { appMessage.error(t("aidpKnowledge.deleteKbFailed")); } }, }); }, - [activeKbId, appMessage, t, fetchKbs, kbPage] + [ + activeKbId, + appMessage, + t, + fetchKbs, + kbPage, + debouncedKbKeyword, + stopUploadWatch, + ] ); // ---- Edit KB ---- @@ -212,9 +338,11 @@ const AidpKnowledgeConfiguration: React.FC = () => { // ---- After create success ---- // The create response already contains the resource. Insert it locally - // instead of scanning up to 50 expensive server-side pages. + // instead of scanning up to 50 expensive server-side pages. Files uploaded + // together with the KB are watched like any other upload so their processing + // status appears without a manual refresh. const handleCreateKbSuccess = useCallback( - (newKb: AidpKnowledgeBaseItem) => { + (newKb: AidpKnowledgeBaseItem, uploadedFileIds: string[] = []) => { setCreateModalOpen(false); setKbs((current) => [newKb, ...current.filter((kb) => kb.kds_id !== newKb.kds_id)].slice( @@ -228,43 +356,66 @@ const AidpKnowledgeConfiguration: React.FC = () => { return nextTotal; }); setKbTotalReliable(true); + stopUploadWatch(); + startUploadWatch(uploadedFileIds); setActiveKbId(newKb.kds_id); setSelectedKb(newKb); setActiveKbDetail(newKb); setDocPage(1); setDocHasMore(false); setDocTotalReliable(true); + setDocProcessingCount(0); void fetchDocs(newKb.kds_id, 1); }, - [fetchDocs, kbPage] + [fetchDocs, kbPage, startUploadWatch, stopUploadWatch] ); + // ---- Refresh the active KB metadata (counts / name) ---- + const refreshActiveKbDetail = useCallback(() => { + if (!activeKbId) return; + void aidpKnowledgeService + .getKb(activeKbId) + .then((detail) => { + const refreshed = { + ...selectedKb, + ...detail, + kds_id: activeKbId, + kds_name: detail.kds_name || selectedKb?.kds_name || activeKbId, + } as AidpKnowledgeBaseItem; + setSelectedKb(refreshed); + setActiveKbDetail(detail); + setKbs((current) => + current.map((kb) => (kb.kds_id === activeKbId ? refreshed : kb)) + ); + }) + .catch((error) => + log.error("Failed to refresh active AIDP KB detail:", error) + ); + }, [activeKbId, selectedKb]); + // ---- After documents uploaded ---- - const handleDocsUploaded = useCallback(() => { - if (activeKbId) { + // Refresh immediately so the uploaded files show up without the user having + // to press refresh, then watch them until AIDP reports a terminal status. + const handleDocsUploaded = useCallback( + (uploadedFileIds: string[]) => { + startUploadWatch(uploadedFileIds); + if (!activeKbId) return; // Reset doc pagination to page 1 so data and pagination UI stay in sync setDocPage(1); void fetchDocs(activeKbId, 1); - void aidpKnowledgeService - .getKb(activeKbId) - .then((detail) => { - const refreshed = { - ...selectedKb, - ...detail, - kds_id: activeKbId, - kds_name: detail.kds_name || selectedKb?.kds_name || activeKbId, - } as AidpKnowledgeBaseItem; - setSelectedKb(refreshed); - setActiveKbDetail(detail); - setKbs((current) => - current.map((kb) => (kb.kds_id === activeKbId ? refreshed : kb)) - ); - }) - .catch((error) => - log.error("Failed to refresh active AIDP KB detail:", error) - ); - } - }, [activeKbId, fetchDocs, selectedKb]); + refreshActiveKbDetail(); + }, + [activeKbId, fetchDocs, refreshActiveKbDetail, startUploadWatch] + ); + + // ---- Manual refresh (refresh button) ---- + // Refreshes the visible page and the KB metadata but never starts an upload + // watch: the user is not waiting for a file they just added. + const handleRefreshDocs = useCallback(() => { + if (!activeKbId) return; + void fetchDocs(activeKbId, docPage); + refreshActiveKbDetail(); + }, [activeKbId, docPage, fetchDocs, refreshActiveKbDetail]); // Active KB item is stored in `selectedKb` state (not derived from `kbs`), // because the KB list is server-paginated and refetching it after upload @@ -273,18 +424,17 @@ const AidpKnowledgeConfiguration: React.FC = () => { return (
- {/* Two-column layout — content-sized cards with a single - scroll container; no card stretches to viewport height. */} -
- +
+ {/* Left column: KB list */} { hasMore={kbHasMore} currentPage={kbPage} pageSize={KB_PAGE_SIZE} - onPageChange={(page) => fetchKbs(page)} + keyword={kbKeyword} + onKeywordChange={setKbKeyword} + onPageChange={(page) => fetchKbs(page, debouncedKbKeyword)} onSelect={handleSelectKb} - onRefresh={() => fetchKbs(kbPage)} + onRefresh={() => fetchKbs(kbPage, debouncedKbKeyword)} onCreateNew={() => setCreateModalOpen(true)} onEdit={handleEditKb} onDelete={handleDeleteKb} @@ -311,6 +463,7 @@ const AidpKnowledgeConfiguration: React.FC = () => { {/* Right column: Document list or empty state */} { isLoading={loadingDocs} currentPage={docPage} pageSize={DOC_PAGE_SIZE} - onPageChange={(page) => fetchDocs(activeKbId!, page)} + onPageChange={(page) => { + // The upload watch only makes sense on the page the upload + // landed on; the status poller still covers other pages. + stopUploadWatch(); + void fetchDocs(activeKbId!, page); + }} onDocsUploaded={handleDocsUploaded} - onRefresh={handleDocsUploaded} + onRefresh={handleRefreshDocs} /> ) : ( -
-
-
-
- -
-

- {t("aidpKnowledge.selectKbTitle")} -

-

- {t("aidpKnowledge.selectKbHint")} -

+
+
+
+
+

+ {t("aidpKnowledge.selectKbTitle")} +

+

+ {t("aidpKnowledge.selectKbHint")} +

)} diff --git a/frontend/ext_components/aidp/components/AidpKnowledgeList.tsx b/frontend/ext_components/aidp/components/AidpKnowledgeList.tsx index 2911d32a0..b21d0bd9f 100644 --- a/frontend/ext_components/aidp/components/AidpKnowledgeList.tsx +++ b/frontend/ext_components/aidp/components/AidpKnowledgeList.tsx @@ -1,12 +1,19 @@ import React, { useMemo } from "react"; import { useTranslation } from "react-i18next"; -import { Button, Pagination, Tag, Tooltip } from "antd"; +import { Button, Input, Pagination, Tooltip } from "antd"; import { - PlusOutlined, - ReloadOutlined, -} from "@ant-design/icons"; -import { SquarePen, Trash2 } from "lucide-react"; + BookOpen, + CircleOff, + Eye, + FolderOpen, + Glasses, + PencilRuler, + Search, + SquarePen, + Trash2, +} from "lucide-react"; +import { PlusOutlined, ReloadOutlined } from "@ant-design/icons"; import type { AidpKnowledgeBaseItem } from "@/types/agentConfig"; import { useGroupList } from "@/hooks/group/useGroupList"; @@ -18,12 +25,18 @@ interface AidpKnowledgeListProps { activeKbId: string | null; isLoading: boolean; total: number; - /** True when `total` came from AIDP Count API (reliable). When false we - * show a simple prev/next pagination without "共 N 条". */ + /** True when `total` came from AIDP Count API. */ totalReliable: boolean; hasMore: boolean; currentPage: number; pageSize: number; + /** Raw search box value. The parent debounces it before querying, so this is + * intentionally the un-debounced text the user is currently typing. The + * filtering itself happens server-side: AIDP narrows the catalog by this + * keyword, so the page below renders exactly what the API returned and the + * reported total always matches the rendered cards. */ + keyword: string; + onKeywordChange: (value: string) => void; onPageChange: (page: number) => void; onSelect: (kb: AidpKnowledgeBaseItem) => void; onRefresh: () => void; @@ -32,6 +45,20 @@ interface AidpKnowledgeListProps { onDelete: (kb: AidpKnowledgeBaseItem) => void; } +const permissionIcon = (permission?: string) => { + const props = { size: 13, className: "text-gray-500" }; + switch (permission) { + case "EDIT": + return ; + case "READ_ONLY": + return ; + case "PRIVATE": + return ; + default: + return ; + } +}; + const AidpKnowledgeList: React.FC = ({ kbs, activeKbId, @@ -41,6 +68,8 @@ const AidpKnowledgeList: React.FC = ({ hasMore, currentPage, pageSize, + keyword, + onKeywordChange, onPageChange, onSelect, onRefresh, @@ -50,221 +79,322 @@ const AidpKnowledgeList: React.FC = ({ }) => { const { t } = useTranslation(); - // Load groups for the current tenant so we can render group_ids as names. - // ``useAuthorizationContext`` is the right hook here (the similarly-named - // ``useAuthenticationContext`` carries only ``session`` — no ``user`` object). const { user } = useAuthorizationContext(); const tenantId = user?.tenantId ?? null; const { data: groupListData } = useGroupList(tenantId); const groupById = useMemo(() => { const map = new Map(); - (groupListData?.groups ?? []).forEach((g) => { - map.set(g.group_id, g.group_name); + (groupListData?.groups ?? []).forEach((group) => { + map.set(group.group_id, group.group_name); }); return map; }, [groupListData]); - // Convert group ids to names, skipping any ids that don't resolve - // (e.g. the group was deleted or the list is not yet loaded). Aligned - // with the local-knowledge-base list renderer. - const getGroupNames = (groupIds: number[] | undefined): string[] => { - if (!Array.isArray(groupIds) || groupIds.length === 0) return []; - return groupIds + // The keyword filter lives on the server: the parent forwards it to AIDP, and + // filtering the returned page again here would hide results whenever AIDP's + // matching is not a plain substring of the name/description. Only the order is + // local, so the most recently updated KB stays on top of the page. + const displayedKbs = useMemo( + () => + [...kbs].sort((a, b) => { + const aTime = Date.parse(a.updated_at || a.created_at || "") || 0; + const bTime = Date.parse(b.updated_at || b.created_at || "") || 0; + return bTime - aTime; + }), + [kbs] + ); + + const getGroupNames = (groupIds?: number[]) => + (groupIds ?? []) .map((id) => groupById.get(id)) - .filter((name): name is string => typeof name === "string" && name.length > 0); + .filter((name): name is string => Boolean(name)); + + const permissionLabel = (permission?: string) => + t(`knowledgeBase.ingroup.permission.${permission || "DEFAULT"}`); + + const renderStatusTag = (kb: AidpKnowledgeBaseItem) => { + const isUnavailable = + kb.resource_status === "UNAVAILABLE" || kb.resource_status === "ORPHANED"; + if (isUnavailable) { + return ( + + {t("aidpKnowledge.kbUnavailable")} + + ); + } + if (kb.permission === "READ_ONLY") { + return ( + + {t("aidpKnowledge.kbReadOnly")} + + ); + } + return null; }; - // Sort alphabetically by name - const displayedKbs = useMemo(() => { - return [...kbs].sort((a, b) => - (a.kds_name || "").localeCompare(b.kds_name || "") + const renderKnowledgeCard = (kb: AidpKnowledgeBaseItem) => { + const isActive = activeKbId === kb.kds_id; + const isUnavailable = + kb.resource_status === "UNAVAILABLE" || kb.resource_status === "ORPHANED"; + const canModify = kb.permission === "EDIT" && !isUnavailable; + const groupNames = getGroupNames(kb.group_ids); + const documentCount = kb.document_count ?? 0; + const chunkCount = kb.chunk_count ?? 0; + const permission = kb.ingroup_permission || "PRIVATE"; + + return ( +
onSelect(kb)} + onKeyDown={(event) => { + if (event.target !== event.currentTarget) return; + if (event.key === "Enter" || event.key === " ") { + event.preventDefault(); + onSelect(kb); + } + }} + > +
+
+
+ +
+
+

+ {kb.kds_name} +

+ + AIDP + +
+
+ +
+ {canModify && ( + +
+
+ +

+ {kb.description?.trim() || t("knowledgeBase.description.default")} +

+ +
+ + {t("knowledgeBase.tag.documents", { count: documentCount })} + + + {t("knowledgeBase.tag.chunks", { count: chunkCount })} + + {kb.embedding_model && kb.embedding_model !== "default" && ( + + {kb.embedding_model} + + )} + {kb.is_multimodal && ( + + multimodal + + )} + {renderStatusTag(kb)} + + {permission === "PRIVATE" ? ( + + {permissionIcon(permission)} + {permissionLabel(permission)} + + ) : ( + groupNames.slice(0, 2).map((groupName) => ( + + {groupName} + + )) + )} + +
+ +
+ + {t("knowledgeBase.tag.updatedAt", { + date: + kb.updated_at || kb.created_at + ? new Date( + kb.updated_at || kb.created_at || "" + ).toLocaleDateString() + : t("aidpKnowledge.createdAtUnknown"), + })} + + + {permissionLabel(permission)} + +
+
); - }, [kbs]); + }; + + let effectiveTotal = currentPage * pageSize; + if (totalReliable) { + effectiveTotal = total; + } else if (hasMore) { + effectiveTotal += 1; + } return ( -
- {/* Header */} -
-
-

- {t("aidpKnowledge.kbListTitle")} -

-
+
+
+
+
+
+ +
+
+

+ {t("knowledgeBase.page.title")} +

+

+ {t("knowledgeBase.page.description")} +

+
+
+ +
-
- {/* List */} -
- {displayedKbs.length > 0 ? ( -
- {displayedKbs.map((kb) => { - const isActive = activeKbId === kb.kds_id; - const isUnavailable = - kb.resource_status === "UNAVAILABLE" || - kb.resource_status === "ORPHANED"; - // Only EDIT-level callers may modify the KB or its files. - const canModify = kb.permission === "EDIT" && !isUnavailable; +
+

+ {t("knowledgeBase.page.all")} + + {t("knowledgeBase.page.count", { count: total })} + +

+ } + value={keyword} + onChange={(event) => onKeywordChange(event.target.value)} + className="h-10 min-w-[240px] max-w-[420px] flex-1 !rounded-lg" + allowClear + /> +
+
- return ( -
onSelect(kb)} - > -
-
-
-

- {kb.kds_name} -

- {isUnavailable && ( - - {t("aidpKnowledge.kbUnavailable")} - - )} - {kb.permission === "READ_ONLY" && !isUnavailable && ( - - {t("aidpKnowledge.kbReadOnly")} - - )} -
- {kb.description && ( -

- {kb.description} -

- )} -
- {kb.ingroup_permission === "PRIVATE" && ( - - {t("knowledgeBase.ingroup.permission.PRIVATE")} - - )} - {/* Authorized user-group tags. Aligned with the local - knowledge base list: only render group names when - ``ingroup_permission !== "PRIVATE"``, each group - gets its own blue tag, and when there are no - groups to show we render nothing (no "not - authorized" fallback). Gated by the ``group:read`` - permission so users without group visibility see - the KB card cleanly without the tag area. */} - - {kb.ingroup_permission !== "PRIVATE" && - getGroupNames(kb.group_ids).map((groupName, idx) => ( - - {groupName} - - ))} - - {kb.created_at ? ( - - {t("aidpKnowledge.createdAt", { - date: new Date(kb.created_at).toLocaleDateString(), - })} - - ) : ( - {t("aidpKnowledge.createdAtUnknown")} - )} -
-
-
- {canModify && ( - -
-
-
- ); - })} +
+ {isLoading && kbs.length === 0 ? ( +
+ Loading...
) : ( -
- {t("aidpKnowledge.listEmpty")} +
+ + {displayedKbs.map(renderKnowledgeCard)}
)} -
- {/* Server-side pagination. - AIDP exposes a dedicated Count API for KBs which the backend calls - alongside the list request. When Count succeeds, `totalReliable` - is true and we display the full pagination (page numbers + - "共 N 条"). When Count fails (e.g. endpoint unavailable), we fall - back to simple prev/next mode using `has_more`. */} - {kbs.length > 0 && (() => { - const effectiveTotal = totalReliable - ? total - : (hasMore - ? currentPage * pageSize + 1 - : currentPage * pageSize); - return ( -
- t("aidpKnowledge.showTotal", { count: total }) - : undefined - } - size="small" - /> + {!isLoading && displayedKbs.length === 0 && ( +
+ {keyword.trim() + ? t("knowledgeBase.list.noResults") + : t("aidpKnowledge.listEmpty")}
- ); - })()} + )} +
+ + {kbs.length > 0 && ( +
+ t("aidpKnowledge.showTotal", { count }) + : undefined + } + size="small" + /> +
+ )}
); }; diff --git a/frontend/ext_components/aidp/components/AidpUpdateKbModal.tsx b/frontend/ext_components/aidp/components/AidpUpdateKbModal.tsx index cb9a0317e..415290f06 100644 --- a/frontend/ext_components/aidp/components/AidpUpdateKbModal.tsx +++ b/frontend/ext_components/aidp/components/AidpUpdateKbModal.tsx @@ -1,16 +1,21 @@ "use client"; -import React, { useEffect, useMemo } from "react"; +import React, { useEffect } from "react"; import { useTranslation } from "react-i18next"; -import { Modal, Form, Input, Select, message } from "antd"; +import { Collapse, Modal, Form, message } from "antd"; +import { SettingOutlined } from "@ant-design/icons"; import type { AidpKnowledgeBaseItem } from "@/types/agentConfig"; import aidpKnowledgeService from "@/ext_components/aidp/services/aidpKnowledgeService"; -import { AIDP_KNOWLEDGE_BASE_NAME_PATTERN } from "@/const/knowledgeBase"; -import { useGroupList } from "@/hooks/group/useGroupList"; -import { useAuthorizationContext } from "@/components/providers/AuthorizationProvider"; -import { USER_ROLES } from "@/const/auth"; +import { useAidpGroupOptions } from "../hooks/useAidpGroupOptions"; +import { + AIDP_MODAL_STYLES, + AidpKnowledgeBaseBasicFields, + AidpKnowledgeBaseModalFooter, + AidpKnowledgeBaseModalHeader, + AidpKnowledgeBasePermissionFields, +} from "./AidpKnowledgeBaseModalParts"; interface AidpUpdateKbModalProps { open: boolean; @@ -29,24 +34,9 @@ const AidpUpdateKbModal: React.FC = ({ const [form] = Form.useForm(); const [loading, setLoading] = React.useState(false); - // Mirror the create-modal wiring: the authorization context exposes - // ``user.tenantId``, which we feed into ``useGroupList`` to enumerate - // the tenant's groups for the access-group picker below. - const { user } = useAuthorizationContext(); - const isUser = user?.role === USER_ROLES.USER; - const canConfigureGroupPermissions = !!user && !isUser; - const tenantId = user?.tenantId ?? null; - const { data: groupListData } = useGroupList( - canConfigureGroupPermissions ? tenantId : null - ); - const groupOptions = useMemo( - () => - (groupListData?.groups ?? []).map((g) => ({ - value: g.group_id, - label: g.group_name, - })), - [groupListData] - ); + const { isUser, canConfigureGroupPermissions, groupOptions } = + useAidpGroupOptions(); + const [advancedOpen, setAdvancedOpen] = React.useState(false); const ingroupPermission = Form.useWatch("ingroup_permission", form); @@ -54,20 +44,21 @@ const AidpUpdateKbModal: React.FC = ({ // that predate the column — normalize to an empty array so the Select // (mode="multiple") receives a value shape it accepts. useEffect(() => { - if (open && knowledgeBase) { - form.setFieldsValue({ - name: knowledgeBase.kds_name, - description: knowledgeBase.description || "", - ingroup_permission: isUser - ? "PRIVATE" - : knowledgeBase.ingroup_permission || "READ_ONLY", - group_ids: isUser - ? [] - : Array.isArray(knowledgeBase.group_ids) - ? knowledgeBase.group_ids - : [], - }); - } + if (!open) return; + setAdvancedOpen(false); + if (!knowledgeBase) return; + form.setFieldsValue({ + name: knowledgeBase.kds_name, + description: knowledgeBase.description || "", + ingroup_permission: isUser + ? "PRIVATE" + : knowledgeBase.ingroup_permission || "READ_ONLY", + group_ids: isUser + ? [] + : Array.isArray(knowledgeBase.group_ids) + ? knowledgeBase.group_ids + : [], + }); }, [open, knowledgeBase, form, isUser]); const handleOk = async () => { @@ -164,96 +155,74 @@ const AidpUpdateKbModal: React.FC = ({ return ( + } > -
- + + - - - - - - {canConfigureGroupPermissions && ( - <> - + {canConfigureGroupPermissions && ( + + setAdvancedOpen( + Array.isArray(keys) + ? keys.includes("advanced") + : keys === "advanced" + ) + } + items={[ { - required: true, - message: t("aidpKnowledge.createIngroupPermissionRequired"), + key: "advanced", + label: ( + + + {t("aidpKnowledge.createAdvancedOptions")} + + ), + children: ( +
+ +
+ ), }, ]} - > - -
- - )} -
+ /> + )} + +
); }; diff --git a/frontend/ext_components/aidp/hooks/useAidpGroupOptions.ts b/frontend/ext_components/aidp/hooks/useAidpGroupOptions.ts new file mode 100644 index 000000000..ff0a42e36 --- /dev/null +++ b/frontend/ext_components/aidp/hooks/useAidpGroupOptions.ts @@ -0,0 +1,30 @@ +import { useMemo } from "react"; + +import { USER_ROLES } from "@/const/auth"; +import { useAuthorizationContext } from "@/components/providers/AuthorizationProvider"; +import { useGroupList } from "@/hooks/group/useGroupList"; + +export interface AidpGroupOption { + value: number; + label: string; +} + +export const useAidpGroupOptions = () => { + const { user } = useAuthorizationContext(); + const isUser = user?.role === USER_ROLES.USER; + const canConfigureGroupPermissions = Boolean(user) && !isUser; + const tenantId = user?.tenantId ?? null; + const { data: groupListData } = useGroupList( + canConfigureGroupPermissions ? tenantId : null + ); + const groupOptions = useMemo( + () => + (groupListData?.groups ?? []).map((group) => ({ + value: group.group_id, + label: group.group_name, + })), + [groupListData] + ); + + return { isUser, canConfigureGroupPermissions, groupOptions }; +}; diff --git a/frontend/ext_components/aidp/services/aidpKnowledgeService.ts b/frontend/ext_components/aidp/services/aidpKnowledgeService.ts index 405924233..80fb4830e 100644 --- a/frontend/ext_components/aidp/services/aidpKnowledgeService.ts +++ b/frontend/ext_components/aidp/services/aidpKnowledgeService.ts @@ -1,7 +1,7 @@ /** * AIDP Knowledge Base Management Service * - * Wraps the 8 AIDP management backend endpoints. + * Wraps the AIDP management backend endpoints. * Credentials (server_url, api_key) are read by the backend from environment variables. */ @@ -29,19 +29,25 @@ export interface AidpKbDetail { ingroup_permission?: "EDIT" | "READ_ONLY" | "PRIVATE"; group_ids?: number[]; resource_status?: - | "ACTIVE" - | "CREATING" - | "DELETE_PENDING" - | "ORPHANED" - | "UNAVAILABLE"; + "ACTIVE" | "CREATING" | "DELETE_PENDING" | "ORPHANED" | "UNAVAILABLE"; } export interface AidpDocumentItem { - file_ino_no: string; + file_uuid: string; + file_ino_no: number; file_name: string; file_size?: number; file_type?: string; created_at?: string; + /** + * Processing status reported by the AIDP file history endpoint: + * `PROCESSING` | `COMPLETED` | `FAILED` (upper-cased by the backend). + * Absent when the backend falls back to the completed-files listing, which + * only ever reports ingested files. + */ + status?: string; + /** Channel directory the file was ingested from. */ + dir_path?: string; } export interface AidpDocumentListResponse { @@ -52,9 +58,16 @@ export interface AidpDocumentListResponse { * fallback estimate when Count fails (false). When false the frontend * should treat the total as approximate and avoid displaying "共 N 条". */ total_reliable?: boolean; + /** + * Number of files still being processed across the WHOLE knowledge base + * (not just the returned page). The list polls while this is greater than + * zero and stops once every file has reached a terminal status. + */ + processing_count?: number; } export interface AidpUploadSuccessItem { + file_uuid: string; file_name: string; file_type: string; file_size: number; @@ -78,6 +91,62 @@ export interface AidpUploadResponse { failed_list: AidpUploadFailedItem[]; } +export interface AidpDocumentOperationItem { + file_uuid: string; +} + +export interface AidpDocumentRemoveResponse { + summary: { + total: number; + success: number; + failed: number; + }; + success_list: AidpDocumentOperationItem[]; + failed_list: AidpDocumentOperationItem[]; +} + +type AidpOperationSummary = { + total: number; + success: number; + failed: number; +}; + +type AidpOperationResponse = { + summary: AidpOperationSummary; + success_list: TSuccess[]; + failed_list: TFailure[]; +}; + +const normalizeAidpOperationResponse = ( + result: Partial> +): AidpOperationResponse => { + const successList: TSuccess[] = Array.isArray(result.success_list) + ? result.success_list + : []; + const failedList: TFailure[] = Array.isArray(result.failed_list) + ? result.failed_list + : []; + + return { + summary: { + total: + typeof result.summary?.total === "number" + ? result.summary.total + : successList.length + failedList.length, + success: + typeof result.summary?.success === "number" + ? result.summary.success + : successList.length, + failed: + typeof result.summary?.failed === "number" + ? result.summary.failed + : failedList.length, + }, + success_list: successList, + failed_list: failedList, + }; +}; + export interface AidpModelItem { /** Display / identifier used for the model (sent to AIDP as ``vlm_model``). */ model_name: string; @@ -174,15 +243,21 @@ function buildUrl( class AidpKnowledgeService { /** - * List knowledge bases (paginated). + * List knowledge bases (paginated), optionally filtered by name. + * + * `keyword` is forwarded to the backend, which passes it on to AIDP for + * server-side filtering. A blank keyword is omitted from the query string + * entirely so an unfiltered call produces the same request as before. */ async listKbs( page: number = 1, - pageSize: number = 10 + pageSize: number = 10, + keyword?: string ): Promise { const url = buildUrl(API_ENDPOINTS.aidpMgmt.knowledgeBases, { page, page_size: pageSize, + keyword: keyword?.trim() || undefined, }); const response = await fetchWithErrorHandling(url, { @@ -346,31 +421,10 @@ class AidpKnowledgeService { } const result = (await response.json()) as Partial; - const successList = Array.isArray(result.success_list) - ? result.success_list - : []; - const failedList = Array.isArray(result.failed_list) - ? result.failed_list - : []; - - return { - summary: { - total: - typeof result.summary?.total === "number" - ? result.summary.total - : successList.length + failedList.length, - success: - typeof result.summary?.success === "number" - ? result.summary.success - : successList.length, - failed: - typeof result.summary?.failed === "number" - ? result.summary.failed - : failedList.length, - }, - success_list: successList, - failed_list: failedList, - }; + return normalizeAidpOperationResponse< + AidpUploadSuccessItem, + AidpUploadFailedItem + >(result); } /** @@ -451,8 +505,52 @@ class AidpKnowledgeService { typeof result.total_reliable === "boolean" ? result.total_reliable : typeof result.total_count === "number", + processing_count: + typeof result.processing_count === "number" + ? result.processing_count + : undefined, }; } + + /** + * Remove one document from an AIDP knowledge base. + * The AIDP API accepts an array, so the single-document UI sends one item. + */ + async removeDoc( + id: string, + fileUuid: string + ): Promise { + const url = buildUrl(API_ENDPOINTS.aidpMgmt.removeKbDocuments(id), {}); + const response = await fetchWithErrorHandling(url, { + method: "POST", + headers: { + ...getAuthHeaders(), + "Content-Type": "application/json", + }, + body: JSON.stringify({ file_uuids: [fileUuid] }), + }); + const result = + (await response.json()) as Partial; + return normalizeAidpOperationResponse< + AidpDocumentOperationItem, + AidpDocumentOperationItem + >(result); + } + + /** + * Download one document through the AIDP management backend. + */ + async downloadDoc(id: string, fileUuid: string): Promise { + const url = buildUrl(API_ENDPOINTS.aidpMgmt.downloadKbDocument(id), {}); + return fetchWithErrorHandling(url, { + method: "POST", + headers: { + ...getAuthHeaders(), + "Content-Type": "application/json", + }, + body: JSON.stringify({ file_uuid: fileUuid }), + }); + } } const aidpKnowledgeService = new AidpKnowledgeService(); diff --git a/frontend/ext_components/aidp/services/aidpUploadUtils.ts b/frontend/ext_components/aidp/services/aidpUploadUtils.ts new file mode 100644 index 000000000..98de81dfb --- /dev/null +++ b/frontend/ext_components/aidp/services/aidpUploadUtils.ts @@ -0,0 +1,15 @@ +import type { AidpUploadFailedItem } from "./aidpKnowledgeService"; + +export const getAidpUploadFailureDetails = ( + failedList: AidpUploadFailedItem[], + language: string, + fallbackMessage: string +) => { + const isChinese = language.startsWith("zh"); + return failedList.map((item) => { + const reason = isChinese + ? item.reason_zh || item.reason_en + : item.reason_en || item.reason_zh; + return `${item.file_name}: ${reason || fallbackMessage}`; + }); +}; diff --git a/frontend/lib/aidpDocumentStatus.ts b/frontend/lib/aidpDocumentStatus.ts new file mode 100644 index 000000000..2cff52e8a --- /dev/null +++ b/frontend/lib/aidpDocumentStatus.ts @@ -0,0 +1,128 @@ +/** + * AIDP knowledge-file status helpers. + * + * Intentionally dependency-free (no imports) so the behaviour can be unit + * tested directly with `node --test`, like the other standalone helpers under + * `lib/`. + * + * Background: AIDP accepts an upload before it has chunked/embedded/indexed the + * file, and the file may not show up in the history directory for a moment + * after that. The list therefore has to keep refreshing itself after an upload + * until every uploaded file is listed with a terminal status — that is what + * `findPendingUploadIds` decides. + */ + +/** Processing statuses AIDP reports for a knowledge file. */ +export const AIDP_DOCUMENT_STATUS = { + UPLOADING: "UPLOADING", + PROCESSING: "PROCESSING", + EXTRACTING: "EXTRACTING", + COMPLETED: "COMPLETED", + FAILED: "FAILED", +} as const; + +/** + * Terminal statuses: the file will not change again, so it is pointless to keep + * polling for it. + */ +export const AIDP_DOC_TERMINAL_STATUSES: readonly string[] = [ + AIDP_DOCUMENT_STATUS.COMPLETED, + AIDP_DOCUMENT_STATUS.FAILED, +]; + +/** + * In-progress statuses, in the order AIDP walks through them: the file is being + * uploaded, then chunked/embedded, then its content is extracted. Every one of + * them keeps the status column live and must keep the list polling. + */ +export const AIDP_DOC_IN_PROGRESS_STATUSES: readonly string[] = [ + AIDP_DOCUMENT_STATUS.UPLOADING, + AIDP_DOCUMENT_STATUS.PROCESSING, + AIDP_DOCUMENT_STATUS.EXTRACTING, +]; + +/** + * Polling interval of the AIDP document list while work is outstanding. + * Ingestion takes seconds to minutes, so ten seconds keeps the status column + * responsive without hammering the backend. + */ +export const AIDP_DOC_STATUS_POLL_MS = 10000; + +/** + * Upper bound on how long the list keeps refreshing for files the user just + * uploaded. Guards against polling forever when AIDP silently drops an upload + * (or never reports the file through the history endpoint). + */ +export const AIDP_DOC_UPLOAD_WATCH_TIMEOUT_MS = 5 * 60 * 1000; + +/** Normalize a status for comparison; AIDP capitalization is not guaranteed. */ +export const normalizeAidpDocStatus = (status?: string): string => + (status || "").trim().toUpperCase(); + +/** + * Whether the status means "still on its way in". + * + * Any status AIDP has not told us is terminal counts as in progress, so a new + * upstream state keeps the column live instead of freezing the list. + */ +export const isAidpDocProcessing = (status?: string): boolean => { + const normalized = normalizeAidpDocStatus(status); + if (!normalized) return false; + return !isAidpDocTerminal(normalized); +}; + +/** Whether the status means "finished, one way or the other". */ +export const isAidpDocTerminal = (status?: string): boolean => + AIDP_DOC_TERMINAL_STATUSES.includes(normalizeAidpDocStatus(status)); + +/** Structural view of a listed document, so this module stays import-free. */ +export interface AidpDocumentStatusView { + file_ino_no?: string | number | null; + status?: string; +} + +/** Structural view of an entry of the AIDP upload response `success_list`. */ +export interface AidpUploadedFileView { + file_ino_no?: string | number | null; +} + +/** + * Collect the identities of the files AIDP accepted in an upload response. + * + * Entries without a usable `file_ino_no` are skipped rather than coerced into + * the literal string `"undefined"`, which would never match a listed document + * and would therefore keep the watch alive until it times out. + */ +export const collectUploadedFileIds = ( + successList?: readonly AidpUploadedFileView[] +): string[] => { + const ids: string[] = []; + for (const item of successList ?? []) { + const raw = item?.file_ino_no; + if (raw === undefined || raw === null || raw === "") continue; + ids.push(String(raw)); + } + return ids; +}; + +/** + * Which of the uploaded files are still unresolved. + * + * A file counts as resolved only once it appears in the list with a terminal + * status. A file that is not listed at all stays pending deliberately: right + * after an upload AIDP may not report it yet, and treating "not listed" as done + * would stop the refresh exactly when the list is still stale — the case where + * the user cannot tell whether the upload worked. + */ +export const findPendingUploadIds = ( + uploadedIds: readonly string[], + documents: readonly AidpDocumentStatusView[] +): string[] => { + if (uploadedIds.length === 0) return []; + const resolved = new Set( + documents + .filter((doc) => isAidpDocTerminal(doc.status)) + .map((doc) => String(doc.file_ino_no)) + ); + return uploadedIds.filter((id) => !resolved.has(String(id))); +}; diff --git a/frontend/public/locales/en/common.json b/frontend/public/locales/en/common.json index fc1c438c8..81681aaa6 100644 --- a/frontend/public/locales/en/common.json +++ b/frontend/public/locales/en/common.json @@ -844,7 +844,37 @@ "knowledgeBase.error.syncFailed": "Failed to sync DataMate knowledge bases", "knowledgeBase.message.testingConnection": "Testing connection...", "knowledgeBase.message.testingSync": "Syncing knowledge bases...", + "knowledgeBase.page.title": "Knowledge Base", + "knowledgeBase.page.description": "Create, manage, and maintain your team's knowledge assets", + "knowledgeBase.page.all": "All Knowledge Bases", + "knowledgeBase.page.count": "{{count}} total", + "knowledgeBase.page.back": "Back to knowledge bases", + "knowledgeBase.create.subtitle": "Configure basic information and adjust it later at any time", + "knowledgeBase.create.field.name": "Knowledge base name", + "knowledgeBase.create.field.description": "Description", + "knowledgeBase.create.optional": "(Optional)", + "knowledgeBase.create.namePlaceholder": "e.g. Product Knowledge Hub", + "knowledgeBase.create.descriptionPlaceholder": "Briefly describe what this knowledge base contains", + "knowledgeBase.create.advancedSettings": "Advanced settings", + "knowledgeBase.create.submit": "Create and enter", + "knowledgeBase.create.field.embeddingModel": "Embedding model", + "knowledgeBase.create.field.groups": "User groups", + "knowledgeBase.create.field.permission": "Group permission", + "knowledgeBase.create.field.preserve": "Document copy", + "knowledgeBase.create.field.quota": "Storage quota", + "knowledgeBase.create.uploadTitle": "Upload documents to build your knowledge base", "knowledgeBase.list.title": "Knowledge Base List", + "knowledgeBase.personalCapacity.title": "Personal knowledge base capacity", + "knowledgeBase.personalCapacity.withQuota": "Used {{used}} / Total {{quota}}", + "knowledgeBase.personalCapacity.unlimited": "Used {{used}} / Unlimited", + "knowledgeBase.personalCapacity.available": "Available", + "knowledgeBase.personalCapacity.total": "Total capacity", + "knowledgeBase.personalCapacity.unlimitedValue": "Unlimited", + "knowledgeBase.personalCapacity.loadFailed": "Personal knowledge base capacity is temporarily unavailable", + "knowledgeBase.capacity.title": "Storage capacity", + "knowledgeBase.capacity.available": "Available", + "knowledgeBase.capacity.total": "Total capacity", + "knowledgeBase.capacity.unlimited": "Unlimited", "knowledgeBase.button.create": "Create", "knowledgeBase.button.sync": "Sync", "knowledgeBase.button.syncDataMate": "Sync DataMate Knowledge Bases", @@ -867,6 +897,11 @@ "knowledgeBase.search.placeholder": "Search knowledge base name", "knowledgeBase.filter.source.placeholder": "Filter by source", "knowledgeBase.filter.model.placeholder": "Filter by model", + "knowledgeBase.filter.title": "Filter knowledge bases", + "knowledgeBase.filter.button": "Filter", + "knowledgeBase.card.create": "Create knowledge base", + "knowledgeBase.card.createDescription": "Upload documents to build a knowledge asset", + "knowledgeBase.tag.updatedAt": "Updated {{date}}", "knowledgeBase.filter.clear": "Clear filters", "knowledgeBase.source.nexent": "{productName}", "knowledgeBase.source.datamate": "DataMate", @@ -963,6 +998,7 @@ "document.hint.uploadToCreate": "Please select files to upload to complete knowledge base creation", "document.hint.noDocuments": "No documents in this knowledge base, please upload documents", "document.table.header.name": "Document Name", + "document.table.header.tags": "Tags", "document.table.header.status": "Status", "document.table.header.size": "Size", "document.table.header.date": "Upload Date", @@ -4001,13 +4037,29 @@ "aidpKnowledge.fetchDocsFailed": "Failed to fetch documents", "aidpKnowledge.docFileName": "File Name", "aidpKnowledge.docType": "Type", + "aidpKnowledge.docStatus": "Status", + "aidpKnowledge.docStatusUploading": "Uploading", + "aidpKnowledge.docStatusProcessing": "Processing", + "aidpKnowledge.docStatusExtracting": "Extracting", + "aidpKnowledge.docStatusCompleted": "Completed", + "aidpKnowledge.docStatusFailed": "Failed", "aidpKnowledge.docSize": "Size", "aidpKnowledge.docCreatedAt": "Created At", + "aidpKnowledge.docActions": "Actions", + "aidpKnowledge.download": "Download", + "aidpKnowledge.delete": "Delete", + "aidpKnowledge.downloadSuccess": "File download started", + "aidpKnowledge.downloadFailed": "Failed to download file", + "aidpKnowledge.confirmDeleteDocTitle": "Delete document", + "aidpKnowledge.confirmDeleteDocContent": "Are you sure you want to delete this document? This cannot be undone.", + "aidpKnowledge.deleteDocSuccess": "Document deleted successfully", + "aidpKnowledge.deleteDocFailed": "Failed to delete document", "aidpKnowledge.noDocuments": "No documents yet", "aidpKnowledge.loadingDocs": "Loading documents...", "aidpKnowledge.uploadSuccess": "{{count}} document(s) uploaded successfully", "aidpKnowledge.uploadPartial": "{{success}} uploaded, {{failed}} failed", "aidpKnowledge.uploadFailed": "Upload failed", + "aidpKnowledge.uploadDuplicateFile": "\"{{fileName}}\" already exists. Please do not upload it again.", "aidpKnowledge.uploading": "Uploading...", "aidpKnowledge.uploadHint": "Click or drag files here to upload", "aidpKnowledge.uploadHintDetail": "Text (txt/json/markdown), Web (html), Docs (pdf/docx/doc/ppt/pptx), Sheets (xlsx/xls/csv), Images (png/jpeg/jpg/bmp)", diff --git a/frontend/public/locales/zh/common.json b/frontend/public/locales/zh/common.json index b5f66b9d1..fe9c0613f 100644 --- a/frontend/public/locales/zh/common.json +++ b/frontend/public/locales/zh/common.json @@ -814,7 +814,37 @@ "knowledgeBase.error.syncFailed": "同步 DataMate 知识库失败", "knowledgeBase.message.testingConnection": "正在测试连接...", "knowledgeBase.message.testingSync": "正在同步知识库...", + "knowledgeBase.page.title": "知识库", + "knowledgeBase.page.description": "创建、管理并维护团队的知识资产", + "knowledgeBase.page.all": "全部知识库", + "knowledgeBase.page.count": "共 {{count}} 个", + "knowledgeBase.page.back": "返回知识库", + "knowledgeBase.create.subtitle": "配置基本信息,后续可随时调整", + "knowledgeBase.create.field.name": "知识库名称", + "knowledgeBase.create.field.description": "描述", + "knowledgeBase.create.optional": "(选填)", + "knowledgeBase.create.namePlaceholder": "例如:产品知识中心", + "knowledgeBase.create.descriptionPlaceholder": "简要说明这个知识库包含什么内容", + "knowledgeBase.create.advancedSettings": "高级设置", + "knowledgeBase.create.submit": "创建并进入", + "knowledgeBase.create.field.embeddingModel": "向量模型", + "knowledgeBase.create.field.groups": "所属用户组", + "knowledgeBase.create.field.permission": "组内权限", + "knowledgeBase.create.field.preserve": "文档副本", + "knowledgeBase.create.field.quota": "存储配额", + "knowledgeBase.create.uploadTitle": "上传文档,构建知识库", "knowledgeBase.list.title": "知识库列表", + "knowledgeBase.personalCapacity.title": "个人知识库容量", + "knowledgeBase.personalCapacity.withQuota": "已用 {{used}} / 总容量 {{quota}}", + "knowledgeBase.personalCapacity.unlimited": "已用 {{used}} / 无限制", + "knowledgeBase.personalCapacity.available": "可用容量", + "knowledgeBase.personalCapacity.total": "总容量", + "knowledgeBase.personalCapacity.unlimitedValue": "无限制", + "knowledgeBase.personalCapacity.loadFailed": "个人知识库容量暂时无法加载", + "knowledgeBase.capacity.title": "存储容量", + "knowledgeBase.capacity.available": "可用容量", + "knowledgeBase.capacity.total": "总容量", + "knowledgeBase.capacity.unlimited": "无限制", "knowledgeBase.button.create": "创建", "knowledgeBase.button.sync": "同步", "knowledgeBase.button.syncDataMate": "同步DataMate知识库", @@ -837,6 +867,11 @@ "knowledgeBase.search.placeholder": "搜索知识库名称", "knowledgeBase.filter.source.placeholder": "筛选来源", "knowledgeBase.filter.model.placeholder": "筛选模型", + "knowledgeBase.filter.title": "筛选知识库", + "knowledgeBase.filter.button": "筛选", + "knowledgeBase.card.create": "新建知识库", + "knowledgeBase.card.createDescription": "上传文档,构建专属知识资产", + "knowledgeBase.tag.updatedAt": "更新于{{date}}", "knowledgeBase.source.nexent": "{productName}", "knowledgeBase.source.datamate": "DataMate", "knowledgeBase.source.dify": "Dify", @@ -931,6 +966,7 @@ "document.hint.uploadToCreate": "请选择文件上传以完成知识库创建", "document.hint.noDocuments": "该知识库中暂无文档,请上传文档", "document.table.header.name": "文档名称", + "document.table.header.tags": "标签", "document.table.header.status": "状态", "document.table.header.size": "大小", "document.table.header.date": "上传日期", @@ -3716,13 +3752,29 @@ "aidpKnowledge.fetchDocsFailed": "获取文档列表失败", "aidpKnowledge.docFileName": "文件名", "aidpKnowledge.docType": "类型", + "aidpKnowledge.docStatus": "状态", + "aidpKnowledge.docStatusUploading": "上传中", + "aidpKnowledge.docStatusProcessing": "处理中", + "aidpKnowledge.docStatusExtracting": "提取中", + "aidpKnowledge.docStatusCompleted": "已完成", + "aidpKnowledge.docStatusFailed": "失败", "aidpKnowledge.docSize": "大小", "aidpKnowledge.docCreatedAt": "创建时间", + "aidpKnowledge.docActions": "操作", + "aidpKnowledge.download": "下载", + "aidpKnowledge.delete": "删除", + "aidpKnowledge.downloadSuccess": "文件下载已开始", + "aidpKnowledge.downloadFailed": "文件下载失败", + "aidpKnowledge.confirmDeleteDocTitle": "删除文档", + "aidpKnowledge.confirmDeleteDocContent": "确定删除该文档吗?此操作无法撤销。", + "aidpKnowledge.deleteDocSuccess": "文档删除成功", + "aidpKnowledge.deleteDocFailed": "文档删除失败", "aidpKnowledge.noDocuments": "暂无文档", "aidpKnowledge.loadingDocs": "正在加载文档...", "aidpKnowledge.uploadSuccess": "成功上传 {{count}} 个文档", "aidpKnowledge.uploadPartial": "{{success}} 个上传成功,{{failed}} 个失败", "aidpKnowledge.uploadFailed": "上传失败", + "aidpKnowledge.uploadDuplicateFile": "\"{{fileName}}\":文件已存在,请勿重复上传。", "aidpKnowledge.uploading": "正在上传...", "aidpKnowledge.uploadHint": "点击或拖拽文件到此处上传", "aidpKnowledge.uploadHintDetail": "支持文本(txt/json/markdown)、网页(html)、文档(pdf/docx/doc/ppt/pptx)、表格(xlsx/xls/csv)、图片(png/jpeg/jpg/bmp)", diff --git a/frontend/services/api.ts b/frontend/services/api.ts index a6fa45e0c..b1c133628 100644 --- a/frontend/services/api.ts +++ b/frontend/services/api.ts @@ -377,6 +377,10 @@ export const API_ENDPOINTS = { kbDetail: (id: string) => `${API_BASE_URL}/aidp-mgmt/knowledge-bases/${id}`, kbDocuments: (id: string) => `${API_BASE_URL}/aidp-mgmt/knowledge-bases/${id}/documents`, + removeKbDocuments: (id: string) => + `${API_BASE_URL}/aidp-mgmt/knowledge-bases/${id}/documents/remove`, + downloadKbDocument: (id: string) => + `${API_BASE_URL}/aidp-mgmt/knowledge-bases/${id}/documents/download`, models: `${API_BASE_URL}/aidp-mgmt/models`, /** PATCH endpoint for the per-KB in-group permission. */ kbPermission: (id: string) => diff --git a/frontend/tests/aidpDocumentStatus.test.ts b/frontend/tests/aidpDocumentStatus.test.ts new file mode 100644 index 000000000..5f420c67f --- /dev/null +++ b/frontend/tests/aidpDocumentStatus.test.ts @@ -0,0 +1,140 @@ +import assert from "node:assert/strict"; +import test from "node:test"; + +import { + AIDP_DOCUMENT_STATUS, + AIDP_DOC_IN_PROGRESS_STATUSES, + collectUploadedFileIds, + findPendingUploadIds, + isAidpDocProcessing, + isAidpDocTerminal, + normalizeAidpDocStatus, + // @ts-expect-error -- Node requires the extension for this standalone test. +} from "../lib/aidpDocumentStatus.ts"; + +test("normalizes status casing and surrounding whitespace", () => { + assert.equal(normalizeAidpDocStatus(" processing "), "PROCESSING"); + assert.equal(normalizeAidpDocStatus("Completed"), "COMPLETED"); + assert.equal(normalizeAidpDocStatus(undefined), ""); + assert.equal(normalizeAidpDocStatus(" "), ""); +}); + +test("treats a status without a terminal outcome as still processing", () => { + assert.equal(isAidpDocProcessing("PROCESSING"), true); + assert.equal(isAidpDocProcessing("processing"), true); + assert.equal(isAidpDocProcessing(AIDP_DOCUMENT_STATUS.COMPLETED), false); + assert.equal(isAidpDocProcessing(undefined), false); + + assert.equal(isAidpDocTerminal("COMPLETED"), true); + assert.equal(isAidpDocTerminal("failed"), true); + assert.equal(isAidpDocTerminal("PROCESSING"), false); + // An unknown status is not terminal: polling must not stop on a status the + // UI does not understand. + assert.equal(isAidpDocTerminal("QUEUED"), false); + assert.equal(isAidpDocTerminal(undefined), false); +}); + +test("keeps every upload stage out of the terminal set", () => { + assert.deepEqual(AIDP_DOC_IN_PROGRESS_STATUSES, [ + "UPLOADING", + "PROCESSING", + "EXTRACTING", + ]); + for (const status of AIDP_DOC_IN_PROGRESS_STATUSES) { + assert.equal( + isAidpDocTerminal(status), + false, + `${status} must not be terminal` + ); + assert.equal( + isAidpDocProcessing(status), + true, + `${status} must count as processing` + ); + } + // Case-insensitive, like every other status comparison. + assert.equal(isAidpDocProcessing("extracting"), true); + assert.equal(isAidpDocProcessing(" uploading "), true); +}); + +test("keeps an upload pending while AIDP has not listed it yet", () => { + // The reported bug: right after an upload the file is not in the list at all, + // so the refresh must keep running instead of stopping immediately. + assert.deepEqual(findPendingUploadIds(["file-1"], []), ["file-1"]); +}); + +test("keeps an upload pending while it is still processing", () => { + const documents = [{ file_ino_no: "file-1", status: "PROCESSING" }]; + assert.deepEqual(findPendingUploadIds(["file-1"], documents), ["file-1"]); +}); + +test("settles an upload once it reaches a terminal status", () => { + assert.deepEqual( + findPendingUploadIds( + ["file-1"], + [{ file_ino_no: "file-1", status: "COMPLETED" }] + ), + [] + ); + assert.deepEqual( + findPendingUploadIds( + ["file-1"], + [{ file_ino_no: "file-1", status: "FAILED" }] + ), + [] + ); +}); + +test("matches numeric upload ids against string document ids", () => { + // AIDP returns numbers from the upload endpoint and strings from the history + // listing, so a strict comparison would never settle the watch. The ids are + // read back through the same helper the components use. + const uploadedIds = collectUploadedFileIds([{ file_ino_no: 17001 }]); + assert.deepEqual(uploadedIds, ["17001"]); + assert.deepEqual( + findPendingUploadIds(uploadedIds, [ + { file_ino_no: "17001", status: "COMPLETED" }, + ]), + [] + ); +}); + +test("reports only the uploads that are still unresolved", () => { + const documents = [ + { file_ino_no: "file-1", status: "COMPLETED" }, + { file_ino_no: "file-2", status: "PROCESSING" }, + { file_ino_no: "file-other", status: "FAILED" }, + ]; + assert.deepEqual( + findPendingUploadIds(["file-1", "file-2", "file-3"], documents), + ["file-2", "file-3"] + ); +}); + +test("returns nothing to wait for when there was no upload", () => { + assert.deepEqual(findPendingUploadIds([], [{ file_ino_no: "file-1" }]), []); +}); + +test("collects uploaded file ids from the upload response", () => { + assert.deepEqual( + collectUploadedFileIds([{ file_ino_no: 17001 }, { file_ino_no: "17002" }]), + ["17001", "17002"] + ); +}); + +test("skips upload entries without a usable id", () => { + // A missing id must not become the literal string "undefined", which would + // never match a listed document and would keep the watch alive until timeout. + assert.deepEqual( + collectUploadedFileIds([ + {}, + { file_ino_no: undefined }, + { file_ino_no: "" }, + { file_ino_no: null }, + { file_ino_no: 17003 }, + ]), + ["17003"] + ); + assert.deepEqual(collectUploadedFileIds(undefined), []); + assert.deepEqual(collectUploadedFileIds([]), []); +}); diff --git a/test/ext_components/aidp/mock_servers/aidp_mgmt_mock_server.py b/test/ext_components/aidp/mock_servers/aidp_mgmt_mock_server.py index 742133bdf..676ae6e7a 100644 --- a/test/ext_components/aidp/mock_servers/aidp_mgmt_mock_server.py +++ b/test/ext_components/aidp/mock_servers/aidp_mgmt_mock_server.py @@ -9,8 +9,22 @@ - DELETE /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{id} (delete) - POST /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{id}/KnowledgeFiles/Upload (upload docs) - GET /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{id}/KnowledgeFiles (list docs) + - POST /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{id}/KnowledgeFiles/Remove (remove docs) + - POST /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{id}/KnowledgeFiles/Download (download doc) + - GET /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{id}/Channels (ingestion channels) + - POST /KnowledgeBase/Tenants/{tenant}/KnowledgeBases/{id}/KnowledgeFiles/History (all-status file history) - POST /KnowledgeBase/Tenants/{tenant}/Retrieval/FusionSearch (search - preserved from reference) +Document status simulation (drives the "processing status" UI): + * Uploaded documents start as ``PROCESSING`` and flip to ``COMPLETED`` once + ``_PROCESSING_SECONDS`` have elapsed, so polling behaviour can be observed + end to end. Tune it with ``POST /_mock/processing-seconds?seconds=N``. + * ``POST /_mock/doc-status`` (body ``{kds_id, file_ino_no, status}``) forces one + document into any status without waiting for the timer, including the + non-terminal ``UPLOADING`` / ``EXTRACTING`` stages. + * ``GET .../KnowledgeFiles`` keeps returning COMPLETED documents only (mirrors + real AIDP), while ``POST .../KnowledgeFiles/History`` returns every status. + Knowledge base + document state is persisted to ``_state/knowledge_bases.json`` (next to this file). On restart the mock loads the file, so KBs created by tests or frontend sessions survive across restarts without re-creation @@ -22,16 +36,18 @@ import argparse import json import logging +import mimetypes import os import time import uuid from pathlib import Path from typing import Any, Dict, List, Literal, Optional +from urllib.parse import quote from fastapi import FastAPI, File, Header, HTTPException, Query, UploadFile from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel, Field -from starlette.responses import JSONResponse +from starlette.responses import JSONResponse, StreamingResponse logger = logging.getLogger("aidp_mgmt_mock") logging.basicConfig( @@ -57,6 +73,23 @@ _KB_PREFIX = f"/KnowledgeBase/Tenants/{TENANT}/KnowledgeBases" _MODELS_PREFIX = f"/ModelService/Tenants/{TENANT}/Service" +# Document status constants mirroring the real AIDP vocabulary. +STATUS_PROCESSING = "PROCESSING" +STATUS_COMPLETED = "COMPLETED" +STATUS_FAILED = "FAILED" +_TERMINAL_STATUSES = {STATUS_COMPLETED, STATUS_FAILED} + +# One ingestion channel per knowledge base, exposed at the KB-scoped path the +# real AIDP uses (``.../KnowledgeBases/{kds_id}/Channels``). ``src_dir`` embeds +# the KB id, which is how this mock maps a History ``dir_path`` back to its +# documents. The tenant-scoped variant is deliberately NOT served: a wrong path +# in the adapter must fail here exactly as it fails against AIDP. +_CHANNEL_ROOT = "/aidp/knowledge" + +# Seconds an uploaded document stays PROCESSING before turning COMPLETED. +# Overridable at runtime through POST /_mock/processing-seconds. +_PROCESSING_SECONDS = 8.0 + # Directory for persisted runtime state. Lives next to this file so the mock # is self-contained (no absolute paths) and stays out of version control via # ``.gitignore``. Only KB + document state is persisted; failure-injection @@ -118,14 +151,16 @@ def _seed_initial_data() -> None: # Seed some documents for the FAQ KB so list_docs is non-empty by default. _DOCUMENTS_BY_KB["aidp-kb-faq"] = [ { - "file_ino_no": "file-faq-001", + "file_uuid": "00000000-0000-4000-8000-000000000001", + "file_ino_no": 1001, "file_name": "常见问题汇总.txt", "file_size": 2048, "file_type": "txt", "create_time": 1718000400, }, { - "file_ino_no": "file-faq-002", + "file_uuid": "00000000-0000-4000-8000-000000000002", + "file_ino_no": 1002, "file_name": "troubleshooting.md", "file_size": 4096, "file_type": "md", @@ -192,7 +227,99 @@ def _load_state() -> None: ) +def _ensure_document_uuids() -> None: + """Backfill stable UUIDs for state created before UUID support existed.""" + for kds_id, documents in _DOCUMENTS_BY_KB.items(): + if not isinstance(documents, list): + continue + for document in documents: + if not isinstance(document, dict) or document.get("file_uuid"): + continue + file_ino_no = str(document.get("file_ino_no") or uuid.uuid4()) + document["file_uuid"] = str( + uuid.uuid5(uuid.NAMESPACE_URL, f"mock-aidp:{kds_id}:{file_ino_no}") + ) + + _load_state() +_ensure_document_uuids() +_save_state() + + +def _public_document(document: Dict[str, Any]) -> Dict[str, Any]: + """Return document metadata without any mock-only private fields.""" + return {key: value for key, value in document.items() if not key.startswith("_")} + + +def _find_document(kds_id: str, file_uuid: str) -> Optional[Dict[str, Any]]: + return next( + ( + document + for document in _DOCUMENTS_BY_KB.get(kds_id, []) + if document.get("file_uuid") == file_uuid + ), + None, + ) + + +def _document_content(document: Dict[str, Any]) -> bytes: + """Build deterministic mock content for a document download.""" + return ( + f"Mock AIDP content for {document.get('file_name', 'download')}\n" + ).encode("utf-8") + + +def _content_disposition(filename: str) -> str: + """Build a standard ASCII fallback plus RFC 5987 UTF-8 filename header.""" + ascii_name = "".join( + char if 32 <= ord(char) < 127 and char not in {'"', "\\"} else "_" + for char in filename + ) + return f'attachment; filename="{ascii_name}"; filename*=UTF-8\'\'{quote(filename)}' + + +# ============================================================================= +# Document status helpers +# ============================================================================= +def _channel_src_dir(kds_id: str) -> str: + """Return the source directory of a knowledge base's ingestion channel.""" + return f"{_CHANNEL_ROOT}/{kds_id}" + + +def _kds_id_from_dir_path(dir_path: Optional[str]) -> Optional[str]: + """Reverse ``_channel_src_dir`` so a History request maps back to one KB.""" + if not isinstance(dir_path, str): + return None + prefix = f"{_CHANNEL_ROOT}/" + if not dir_path.startswith(prefix): + return None + return dir_path[len(prefix):].strip("/") or None + + +def _doc_effective_status(doc: Dict[str, Any]) -> str: + """Return the document's current status, advancing the processing timer. + + Documents persisted before status simulation existed (and the seed data) + carry no status at all and are treated as already ingested. + """ + status = doc.get("status") + if status is None or status == "": + return STATUS_COMPLETED + if status == STATUS_PROCESSING: + deadline = doc.get("processing_until") + if isinstance(deadline, (int, float)) and time.time() < deadline: + return STATUS_PROCESSING + return STATUS_COMPLETED + return str(status) + + +def _visible_in_completed_listing(doc: Dict[str, Any]) -> bool: + """Whether a document appears in the legacy completed-files listing. + + Real AIDP only exposes ingested files there; files still being processed are + invisible, which is exactly the behaviour the history endpoint replaces. + """ + return _doc_effective_status(doc) == STATUS_COMPLETED # ============================================================================= @@ -211,6 +338,23 @@ class UpdateKbBody(BaseModel): description: Optional[str] = None +class DocHistoryBody(BaseModel): + """Body of POST .../KnowledgeFiles/History (channel + directory scoped).""" + + fs_id: Optional[str] = None + dir_path: Optional[str] = None + + +class DocStatusBody(BaseModel): + """Body of POST /_mock/doc-status (test control over one document).""" + + kds_id: str + file_ino_no: str + status: Literal[ + "UPLOADING", "PROCESSING", "EXTRACTING", "COMPLETED", "FAILED" + ] = "FAILED" + + class MetadataCondition(BaseModel): logical_operator: Literal["and", "or"] = "and" conditions: List[Dict[str, Any]] = Field(default_factory=list) @@ -230,6 +374,14 @@ class FusionSearchRequest(BaseModel): metadata_condition: Optional[MetadataCondition] = None +class RemoveFilesBody(BaseModel): + file_uuids: List[uuid.UUID] = Field(..., min_length=1) + + +class DownloadFileBody(BaseModel): + file_uuid: uuid.UUID = Field(...) + + # ============================================================================= # Auth helper # ============================================================================= @@ -322,6 +474,44 @@ def reset_failures() -> JSONResponse: }) +@app.post("/_mock/processing-seconds") +def set_processing_seconds( + seconds: float = Query(8.0, ge=0.0, le=600.0, description="Seconds a new upload stays PROCESSING"), +) -> JSONResponse: + """Tune how long newly uploaded documents stay PROCESSING. + + Set 0 to make uploads complete immediately, or a large value to keep the + frontend's status polling running while you inspect it. + """ + global _PROCESSING_SECONDS + _PROCESSING_SECONDS = seconds + logger.info("MOCK CONFIG processing seconds = %s", seconds) + return JSONResponse(content={"processing_seconds": _PROCESSING_SECONDS}) + + +@app.post("/_mock/doc-status") +def force_doc_status(body: DocStatusBody) -> JSONResponse: + """Force one document into a given status (used to render a stage in the UI).""" + docs = _DOCUMENTS_BY_KB.get(body.kds_id, []) + for doc in docs: + # Compare as strings: document ids are numeric in some state files and + # strings in others, and callers should not have to care. + if str(doc.get("file_ino_no")) == str(body.file_ino_no): + doc["status"] = body.status + doc.pop("processing_until", None) + _save_state() + logger.info( + "MOCK CONFIG kds_id=%s file=%s status=%s", + body.kds_id, body.file_ino_no, body.status, + ) + return JSONResponse(content={"doc": {**doc, "status": body.status}}) + + raise HTTPException( + status_code=404, + detail=f"Document {body.file_ino_no} not found in {body.kds_id}", + ) + + # ============================================================================= # Knowledge Base CRUD # ============================================================================= @@ -331,12 +521,26 @@ def reset_failures() -> JSONResponse: def list_knowledge_bases( page: int = Query(1, ge=1), page_size: int = Query(20, ge=1, le=100), + keyword: Optional[str] = Query(default=None), authorization: Optional[str] = Header(default=None), ) -> JSONResponse: - """List knowledge bases with pagination + next_link (matches AIDP shape).""" + """List knowledge bases with pagination + next_link (matches AIDP shape). + + ``keyword`` narrows the set by knowledge-base name (case-insensitive), the + way the real AIDP list endpoint does. Without this the mock returned every + KB for a filtered request, which made the search feature look broken in + local runs once the page stopped filtering client-side. + """ _check_auth(authorization) all_items = list(_KNOWLEDGE_BASES.values()) + normalized_keyword = (keyword or "").strip().lower() + if normalized_keyword: + all_items = [ + kb + for kb in all_items + if normalized_keyword in str(kb.get("kds_name") or "").lower() + ] start = (page - 1) * page_size end = start + page_size # Enrich each item with document_count (same as detail endpoint does) @@ -380,7 +584,11 @@ def count_documents( _check_auth(authorization) if kds_id not in _KNOWLEDGE_BASES: raise HTTPException(status_code=404, detail=f"Knowledge base {kds_id} not found") - count = len(_DOCUMENTS_BY_KB.get(kds_id, [])) + # Only ingested files count on the legacy endpoint, matching the list. + count = len([ + doc for doc in _DOCUMENTS_BY_KB.get(kds_id, []) + if _visible_in_completed_listing(doc) + ]) logger.info("COUNT DOCS kds_id=%s count=%d", kds_id, count) return JSONResponse(content={"count": count}) @@ -503,13 +711,26 @@ async def upload_documents( for f in files: try: content = await f.read() - file_ino_no = f"file-{uuid.uuid4().hex[:12]}" + file_ino_no = max( + ( + document["file_ino_no"] + for document in _DOCUMENTS_BY_KB.get(kds_id, []) + if isinstance(document.get("file_ino_no"), int) + ), + default=0, + ) + 1 doc = { + "file_uuid": str(uuid.uuid4()), "file_ino_no": file_ino_no, "file_name": f.filename or "unknown", "file_size": len(content), "file_type": (f.filename.rsplit(".", 1)[-1] if f.filename and "." in f.filename else "bin"), "create_time": int(time.time()), + # Uploaded files enter the ingestion pipeline immediately: the + # legacy list endpoint hides them until the timer elapses, while + # the history endpoint reports them as PROCESSING. + "status": STATUS_PROCESSING, + "processing_until": time.time() + _PROCESSING_SECONDS, } _DOCUMENTS_BY_KB.setdefault(kds_id, []).append(doc) success_docs.append(doc) @@ -535,6 +756,7 @@ async def upload_documents( "file_type": doc["file_type"], "file_size": doc["file_size"], "file_ino_no": doc["file_ino_no"], + "file_uuid": doc["file_uuid"], "first_upload_time": doc["create_time"], } for doc in success_docs @@ -550,16 +772,24 @@ def list_documents( page_size: int = Query(20, ge=1, le=100), authorization: Optional[str] = Header(default=None), ) -> JSONResponse: - """List documents in a knowledge base with pagination.""" + """List documents in a knowledge base with pagination. + + Only ingested (COMPLETED) files are returned: real AIDP hides files that are + still being chunked/embedded here, which is why the frontend used to show + nothing right after an upload. + """ _check_auth(authorization) if kds_id not in _KNOWLEDGE_BASES: raise HTTPException(status_code=404, detail=f"Knowledge base {kds_id} not found") - all_docs = _DOCUMENTS_BY_KB.get(kds_id, []) + all_docs = [ + doc for doc in _DOCUMENTS_BY_KB.get(kds_id, []) + if _visible_in_completed_listing(doc) + ] start = (page - 1) * page_size end = start + page_size - items = all_docs[start:end] + items = [_public_document(doc) for doc in all_docs[start:end]] # Real AIDP returns `next_link` as the authoritative "more pages exist" # signal. When there are no more docs, next_link is simply absent. @@ -576,6 +806,172 @@ def list_documents( }) +@app.post(f"{_KB_PREFIX}/{{kds_id}}/KnowledgeFiles/Remove") +def remove_documents( + kds_id: str, + body: RemoveFilesBody, + authorization: Optional[str] = Header(default=None), +) -> JSONResponse: + """Remove documents by file UUID and return per-file success/failure lists.""" + _check_auth(authorization) + + if kds_id not in _KNOWLEDGE_BASES: + raise HTTPException(status_code=404, detail=f"Knowledge base {kds_id} not found") + + documents = _DOCUMENTS_BY_KB.setdefault(kds_id, []) + remaining = list(documents) + success_list: List[Dict[str, str]] = [] + failed_list: List[Dict[str, str]] = [] + for raw_file_uuid in body.file_uuids: + file_uuid = str(raw_file_uuid) + matched = next( + (document for document in remaining if document.get("file_uuid") == file_uuid), + None, + ) + if matched is None: + failed_list.append({"file_uuid": file_uuid}) + continue + remaining.remove(matched) + success_list.append({"file_uuid": file_uuid}) + + _DOCUMENTS_BY_KB[kds_id] = remaining + _save_state() + logger.info( + "REMOVE DOCS kds_id=%s total=%d success=%d failed=%d", + kds_id, + len(body.file_uuids), + len(success_list), + len(failed_list), + ) + return JSONResponse(content={ + "summary": { + "total": len(body.file_uuids), + "success": len(success_list), + "failed": len(failed_list), + }, + "success_list": success_list, + "failed_list": failed_list, + }) + + +@app.post(f"{_KB_PREFIX}/{{kds_id}}/KnowledgeFiles/Download") +def download_document( + kds_id: str, + body: DownloadFileBody, + authorization: Optional[str] = Header(default=None), +) -> StreamingResponse: + """Return deterministic binary content for a document download.""" + _check_auth(authorization) + + if kds_id not in _KNOWLEDGE_BASES: + raise HTTPException(status_code=404, detail=f"Knowledge base {kds_id} not found") + + file_uuid = str(body.file_uuid) + document = _find_document(kds_id, file_uuid) + if document is None: + raise HTTPException(status_code=404, detail=f"File {file_uuid} not found") + + filename = str(document.get("file_name") or "download") + content = _document_content(document) + content_type = mimetypes.guess_type(filename)[0] or "application/octet-stream" + response_headers = { + "Content-Disposition": _content_disposition(filename), + "X-File-Size": str(len(content)), + } + async def content_stream(): + for offset in range(0, len(content), 8 * 1024): + yield content[offset : offset + 8 * 1024] + + return StreamingResponse( + content_stream(), + media_type=content_type, + headers=response_headers, + ) + + +# ============================================================================= +# Ingestion channels + knowledge-file history +# ============================================================================= + + +@app.get(f"{_KB_PREFIX}/{{kds_id}}/Channels") +def list_channels( + kds_id: str, + authorization: Optional[str] = Header(default=None), +) -> JSONResponse: + """List the ingestion channels feeding one knowledge base. + + The catalog is knowledge-base scoped, so a channel exposes the file-system + id and source directory the history endpoint is addressed with, plus + ``kds_id`` so a caller can associate the channel with its KB. + """ + _check_auth(authorization) + + kb = _KNOWLEDGE_BASES.get(kds_id) + if kb is None: + logger.info("LIST CHANNELS unknown kds_id=%s -> 404", kds_id) + return JSONResponse( + status_code=404, + content={"detail": f"Knowledge base {kds_id} not found"}, + ) + + items = [ + { + "fs_id": f"mock-fs-{kds_id}", + "src_dir": _channel_src_dir(kds_id), + "kds_id": kds_id, + "name": kb.get("kds_name"), + } + ] + logger.info("LIST CHANNELS kds_id=%s returned=%d", kds_id, len(items)) + return JSONResponse(content={"value": items}) + + +@app.post(f"{_KB_PREFIX}/{{kds_id}}/KnowledgeFiles/History") +def knowledge_file_history( + kds_id: str, + body: DocHistoryBody, + authorization: Optional[str] = Header(default=None), +) -> JSONResponse: + """List every file in a channel directory, whatever its processing status. + + The endpoint is knowledge-base scoped, like the channel catalog and the + document list; the body still addresses the request to one channel directory + of that KB. Documents that are still PROCESSING are included with their live + status, and a directory pointing outside the KB answers with an empty list. + """ + _check_auth(authorization) + + if kds_id not in _KNOWLEDGE_BASES: + logger.info("FILE HISTORY unknown kds_id=%s -> 404", kds_id) + return JSONResponse( + status_code=404, + content={"detail": f"Knowledge base {kds_id} not found"}, + ) + + if _kds_id_from_dir_path(body.dir_path) != kds_id: + logger.info( + "FILE HISTORY dir_path=%r does not belong to kds_id=%s -> empty", + body.dir_path, + kds_id, + ) + return JSONResponse(content={"value": []}) + + items = [ + { + **doc, + "dir_path": _channel_src_dir(kds_id), + "status": _doc_effective_status(doc), + } + for doc in _DOCUMENTS_BY_KB.get(kds_id, []) + ] + logger.info( + "FILE HISTORY kds_id=%s fs_id=%s dir_path=%s returned=%d", + kds_id, body.fs_id, body.dir_path, len(items), + ) + return JSONResponse(content={"value": items}) + + # ============================================================================= # FusionSearch (preserved from reference mock) # ============================================================================= diff --git a/test/ext_components/aidp/test_aidp_access_service.py b/test/ext_components/aidp/test_aidp_access_service.py index a8a19fd83..07c2d345f 100644 --- a/test/ext_components/aidp/test_aidp_access_service.py +++ b/test/ext_components/aidp/test_aidp_access_service.py @@ -370,3 +370,25 @@ def test_accessible_row_missing_kb_id_is_skipped(): assert snapshot.accessible_ids == [] assert snapshot.name_to_id == {} + + +def test_channels_cache_is_keyed_per_knowledge_base(): + """Channels are KB-scoped, so another KB must not reuse the cached answer.""" + loader = MagicMock(return_value={"value": [{"fs_id": "fs-1", "src_dir": "/dir/1"}]}) + + first = service.get_cached_aidp_channels( + "https://channels.example", "key", "kb-1", loader=loader + ) + second = service.get_cached_aidp_channels( + "https://channels.example", "key", "kb-1", loader=loader + ) + other = service.get_cached_aidp_channels( + "https://channels.example", "key", "kb-2", loader=loader + ) + + assert first == [{"fs_id": "fs-1", "src_dir": "/dir/1"}] + assert second == first + assert other == first + # One upstream load per KB: the repeat call is served from the cache, and the + # second KB gets its own entry instead of the first KB's channel. + assert loader.call_count == 2 diff --git a/test/ext_components/aidp/test_aidp_mgmt_app.py b/test/ext_components/aidp/test_aidp_mgmt_app.py index 00ca05099..bb845e2cb 100644 --- a/test/ext_components/aidp/test_aidp_mgmt_app.py +++ b/test/ext_components/aidp/test_aidp_mgmt_app.py @@ -21,6 +21,7 @@ from typing import Any from unittest.mock import MagicMock, patch +import httpx import pytest from fastapi import FastAPI from fastapi.testclient import TestClient @@ -46,6 +47,11 @@ def _mod(name): nexent_storage_factory = _mod("nexent.storage.storage_client_factory") nexent_storage_factory.create_storage_client_from_config = MagicMock() +services_pkg = _mod("services") +services_pkg.__path__ = [os.path.join(BACKEND_DIR, "services")] +tag_management_service = _mod("services.tag_management_service") +tag_management_service.TagManagementService = MagicMock() + class _MinIOStorageConfig: def __init__(self, **kwargs): @@ -55,7 +61,7 @@ def __init__(self, **kwargs): nexent_storage_factory.MinIOStorageConfig = _MinIOStorageConfig for mod in (nexent_pkg, nexent_utils, nexent_http_mgr, nexent_storage, - nexent_storage_factory): + nexent_storage_factory, services_pkg, tag_management_service): sys.modules.setdefault(mod.__name__, mod) # Register non-prefixed ``database`` / ``database.client`` stubs so that @@ -568,6 +574,7 @@ def test_list_intersects_before_pagination_and_fetches_only_visible_detail(self) user_id=USER_ID, tenant_id=TENANT_ID, aidp_tenant_id="aidp", + keyword=None, ) mock_detail.assert_called_once_with( aidp_mgmt_app.AIDP_SERVER_URL, @@ -575,6 +582,44 @@ def test_list_intersects_before_pagination_and_fetches_only_visible_detail(self) "kb-2", ) + def test_list_forwards_trimmed_keyword_to_aidp(self): + """A keyword is trimmed before it reaches AIDP.""" + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from types import SimpleNamespace + + with patch.object( + aidp_mgmt_app, + "resolve_current_aidp_access", + return_value=SimpleNamespace(accessible_rows=[]), + ) as mock_snapshot: + response = client.get( + "/aidp-mgmt/knowledge-bases?keyword=%20report%20", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + assert mock_snapshot.call_args.kwargs["keyword"] == "report" + + def test_list_treats_blank_keyword_as_unfiltered(self): + """A whitespace-only keyword must not narrow the listing.""" + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from types import SimpleNamespace + + with patch.object( + aidp_mgmt_app, + "resolve_current_aidp_access", + return_value=SimpleNamespace(accessible_rows=[]), + ) as mock_snapshot: + response = client.get( + "/aidp-mgmt/knowledge-bases?keyword=%20%20", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + assert mock_snapshot.call_args.kwargs["keyword"] is None + def test_list_skips_detail_when_catalog_row_has_card_metadata(self): client = _client() from ext_components.aidp.apps import aidp_mgmt_app @@ -750,8 +795,11 @@ def test_list_documents_uses_count_api(self): from ext_components.aidp.apps import aidp_mgmt_app from ext_components.aidp.services import aidp_permission_service + # No channel is resolvable, so the request exercises the completed-files + # fallback (the history path has its own tests below). with patch.object(aidp_permission_service, "require_permission", return_value=MagicMock(permission="READ_ONLY")), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", return_value=[]), \ patch.object(aidp_mgmt_app, "list_aidp_docs_impl", return_value={"value": [{"name": "a"}]}), \ patch.object(aidp_mgmt_app, "count_aidp_docs_impl", return_value=42): @@ -763,6 +811,7 @@ def test_list_documents_uses_count_api(self): body = response.json() assert body["total_count"] == 42 assert body["has_more"] is True + assert body["processing_count"] == 0 def test_list_and_count_requests_run_concurrently(self): client = _client() @@ -783,6 +832,10 @@ def count_docs(*_args, **_kwargs): aidp_permission_service, "require_permission", return_value=MagicMock(permission="READ_ONLY"), + ), patch.object( + aidp_mgmt_app, + "get_cached_aidp_channels", + return_value=[], ), patch.object( aidp_mgmt_app, "list_aidp_docs_impl", @@ -801,6 +854,153 @@ def count_docs(*_args, **_kwargs): assert response.json()["total_count"] == 1 +# --- Remove/download documents ------------------------------------------- + + +class TestAidpDocumentFileOperations: + def test_remove_forwards_only_uuids(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + aidp_result = { + "summary": {"total": 2, "success": 1, "failed": 1}, + "success_list": [{"file_uuid": "00000000-0000-4000-8000-000000000001"}], + "failed_list": [{"file_uuid": "00000000-0000-4000-8000-000000000002"}], + } + with patch.object( + aidp_permission_service, + "require_permission", + return_value=MagicMock(permission="EDIT"), + ), patch.object( + aidp_mgmt_app, + "remove_aidp_docs_impl", + return_value=aidp_result, + ) as mock_remove: + response = client.post( + "/aidp-mgmt/knowledge-bases/kb-1/documents/remove", + headers=_bearer(), + json={ + "file_uuids": [ + "00000000-0000-4000-8000-000000000001", + "00000000-0000-4000-8000-000000000002", + ] + }, + ) + + assert response.status_code == HTTPStatus.OK + assert response.json() == aidp_result + assert mock_remove.call_args.args[3] == [ + "00000000-0000-4000-8000-000000000001", + "00000000-0000-4000-8000-000000000002", + ] + + def test_remove_requires_file_uuids(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + + with patch.object(aidp_mgmt_app, "remove_aidp_docs_impl") as mock_remove: + response = client.post( + "/aidp-mgmt/knowledge-bases/kb-1/documents/remove", + headers=_bearer(), + json={}, + ) + + assert response.status_code == HTTPStatus.UNPROCESSABLE_ENTITY + mock_remove.assert_not_called() + + def test_remove_requires_standard_file_uuid(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + + with patch.object(aidp_mgmt_app, "remove_aidp_docs_impl") as mock_remove: + response = client.post( + "/aidp-mgmt/knowledge-bases/kb-1/documents/remove", + headers=_bearer(), + json={"file_uuids": ["uuid-1"]}, + ) + + assert response.status_code == HTTPStatus.UNPROCESSABLE_ENTITY + mock_remove.assert_not_called() + + def test_download_returns_binary_response(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + mock_response = httpx.Response( + 200, + headers={ + "Content-Type": "text/plain", + "Content-Disposition": 'attachment; filename="a.txt"', + "X-File-Size": "5", + }, + content=b"hello", + request=httpx.Request("POST", SERVER_URL), + ) + + async def stream_document(*args): + return mock_response + + with patch.object( + aidp_permission_service, + "require_permission", + return_value=MagicMock(permission="READ_ONLY"), + ), patch.object( + aidp_mgmt_app, + "stream_aidp_doc_impl", + side_effect=stream_document, + ) as mock_download: + response = client.post( + "/aidp-mgmt/knowledge-bases/kb-1/documents/download", + headers=_bearer(), + json={"file_uuid": "00000000-0000-4000-8000-000000000001"}, + ) + + assert response.status_code == HTTPStatus.OK + assert response.content == b"hello" + assert response.headers["content-type"].startswith("text/plain") + assert response.headers["content-disposition"] == 'attachment; filename="a.txt"' + assert response.headers["x-file-size"] == "5" + assert mock_download.call_args.args == ( + SERVER_URL, + API_KEY, + "kb-1", + "00000000-0000-4000-8000-000000000001", + ) + assert "x-file-name" not in response.headers + + def test_remove_does_not_invalidate_cache_when_no_file_succeeds(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + aidp_result = { + "summary": {"total": 1, "success": 0, "failed": 1}, + "success_list": [], + "failed_list": [{"file_uuid": "00000000-0000-4000-8000-000000000001"}], + } + with patch.object( + aidp_permission_service, + "require_permission", + return_value=MagicMock(permission="EDIT"), + ), patch.object( + aidp_mgmt_app, + "remove_aidp_docs_impl", + return_value=aidp_result, + ), patch.object(aidp_mgmt_app, "invalidate_aidp_kb_detail_cache") as mock_kb_cache, patch.object( + aidp_mgmt_app, "invalidate_aidp_doc_count_cache" + ) as mock_count_cache: + response = client.post( + "/aidp-mgmt/knowledge-bases/kb-1/documents/remove", + headers=_bearer(), + json={"file_uuids": ["00000000-0000-4000-8000-000000000001"]}, + ) + + assert response.status_code == HTTPStatus.OK + mock_kb_cache.assert_not_called() + mock_count_cache.assert_not_called() + # --- Models list (auth only, no per-KB permission) ------------------------ @@ -1178,6 +1378,7 @@ def test_list_docs_falls_back_when_count_api_fails(self): with patch.object(aidp_permission_service, "require_permission", return_value=MagicMock(permission="READ_ONLY")), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", return_value=[]), \ patch.object(aidp_mgmt_app, "list_aidp_docs_impl", return_value={"value": [{"a": 1}, {"a": 2}]}), \ patch.object(aidp_mgmt_app, "count_aidp_docs_impl", @@ -1357,3 +1558,282 @@ def test_update_skips_sync_when_no_name_change(self): assert response.status_code == HTTPStatus.OK # kds_name is None/empty in both result and body -> sync skipped mock_update_perm.assert_not_called() + + +# --------------------------------------------------------------------------- +# Document list - all-status history data source +# --------------------------------------------------------------------------- + + +class TestListDocumentsHistory: + """The document list prefers the all-status history over completed files. + + AIDP only exposes ingested files through ``.../KnowledgeFiles``, so the list + used to stay empty (and the processing state invisible) while an upload was + still being chunked/embedded. These tests pin the new source: channels + + history, in-process pagination, live statuses, and the graceful fallback + that keeps the endpoint working on AIDP builds without those endpoints. + """ + + _CHANNELS = [ + {"fs_id": "fs-1", "src_dir": "/aidp/knowledge/kb-1", "kds_id": "kb-1"} + ] + + @staticmethod + def _read_only(): + return MagicMock(permission="READ_ONLY") + + def test_history_source_report_statuses_and_processing_count(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + history = { + "value": [ + {"file_ino_no": "f-old", "file_name": "old.txt", + "first_upload_time": 1718000000, "status": "COMPLETED"}, + {"file_ino_no": "f-new", "file_name": "new.pdf", + "first_upload_time": 1718000900, "status": "PROCESSING"}, + {"file_ino_no": "f-bad", "file_name": "bad.txt", + "first_upload_time": 1718000800, "status": "FAILED"}, + ] + } + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + return_value=self._CHANNELS), \ + patch.object(aidp_mgmt_app, "list_aidp_doc_history_impl", + return_value=history) as mock_history, \ + patch.object(aidp_mgmt_app, "list_aidp_docs_impl") as mock_completed, \ + patch.object(aidp_mgmt_app, "count_aidp_docs_impl") as mock_count: + response = client.get( + "/aidp-mgmt/knowledge-bases/kb-1/documents", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + body = response.json() + # Newest upload first, so an accepted file is visible on page 1. + assert [item["file_ino_no"] for item in body["value"]] == [ + "f-new", "f-bad", "f-old", + ] + assert [item["status"] for item in body["value"]] == [ + "PROCESSING", "FAILED", "COMPLETED", + ] + assert body["total_count"] == 3 + assert body["has_more"] is False + assert body["total_reliable"] is True + assert body["processing_count"] == 1 + # The history payload is authoritative: the completed-files listing and + # its Count endpoint must not be hit at all. + mock_history.assert_called_once() + # The call carries the resolved channel plus the KB the path is scoped to. + assert mock_history.call_args.args[2:] == ( + "fs-1", "/aidp/knowledge/kb-1", "kb-1", + ) + mock_completed.assert_not_called() + mock_count.assert_not_called() + + @pytest.mark.parametrize("page,expected_count,expected_has_more", [ + (1, 10, True), + (2, 2, False), + ]) + def test_history_source_paginates_in_process(self, page, expected_count, expected_has_more): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + history = { + "value": [ + {"file_ino_no": f"f-{index}", "file_name": f"{index}.txt", + "first_upload_time": 1718000000 + index, "status": "COMPLETED"} + for index in range(12) + ] + } + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + return_value=self._CHANNELS), \ + patch.object(aidp_mgmt_app, "list_aidp_doc_history_impl", + return_value=history): + response = client.get( + f"/aidp-mgmt/knowledge-bases/kb-1/documents?page={page}&page_size=10", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + body = response.json() + assert len(body["value"]) == expected_count + assert body["total_count"] == 12 + assert body["has_more"] is expected_has_more + assert body["processing_count"] == 0 + + def test_items_without_numeric_timestamp_sort_last(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + history = { + "value": [ + {"file_ino_no": "no-time", "file_name": "iso.txt", + "created_at": "2024-06-10T00:00:00Z", "status": "COMPLETED"}, + {"file_ino_no": "timestamped", "file_name": "ts.txt", + "first_upload_time": 1718000000, "status": "COMPLETED"}, + ] + } + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + return_value=self._CHANNELS), \ + patch.object(aidp_mgmt_app, "list_aidp_doc_history_impl", + return_value=history): + response = client.get( + "/aidp-mgmt/knowledge-bases/kb-1/documents", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + assert [item["file_ino_no"] for item in response.json()["value"]] == [ + "timestamped", "no-time", + ] + + def test_history_source_ignores_non_dict_items(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + history = { + "value": [ + "bogus", + {"file_ino_no": "f-1", "file_name": "ok.txt", + "first_upload_time": 1718000000, "status": "PROCESSING"}, + ] + } + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + return_value=self._CHANNELS), \ + patch.object(aidp_mgmt_app, "list_aidp_doc_history_impl", + return_value=history): + response = client.get( + "/aidp-mgmt/knowledge-bases/kb-1/documents", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + body = response.json() + assert [item["file_ino_no"] for item in body["value"]] == ["f-1"] + assert body["total_count"] == 1 + + def test_history_failure_falls_back_to_completed_listing(self): + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + return_value=self._CHANNELS), \ + patch.object(aidp_mgmt_app, "list_aidp_doc_history_impl", + side_effect=AppException(ErrorCode.AIDP_SERVICE_ERROR, "history down")), \ + patch.object(aidp_mgmt_app, "list_aidp_docs_impl", + return_value={"value": [{"file_name": "done.txt"}]}), \ + patch.object(aidp_mgmt_app, "count_aidp_docs_impl", return_value=1): + response = client.get( + "/aidp-mgmt/knowledge-bases/kb-1/documents", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + body = response.json() + assert body["value"] == [{"file_name": "done.txt"}] + assert body["total_count"] == 1 + # Nothing is known to be processing when the fallback runs. + assert body["processing_count"] == 0 + + def test_empty_history_falls_back_to_kb_scoped_listing(self): + """An empty channel directory must not blank a KB that has files.""" + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + return_value=self._CHANNELS), \ + patch.object(aidp_mgmt_app, "list_aidp_doc_history_impl", + return_value={"value": []}), \ + patch.object(aidp_mgmt_app, "list_aidp_docs_impl", + return_value={"value": [{"file_name": "done.txt"}]}), \ + patch.object(aidp_mgmt_app, "count_aidp_docs_impl", return_value=1): + response = client.get( + "/aidp-mgmt/knowledge-bases/kb-1/documents", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + body = response.json() + assert body["value"] == [{"file_name": "done.txt"}] + assert body["total_count"] == 1 + + def test_history_source_counts_every_in_progress_stage(self): + """Uploading/extracting keep polling alive just like processing does.""" + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + history = { + "value": [ + {"file_ino_no": "f-up", "file_name": "up.txt", + "first_upload_time": 1718000000, "status": "UPLOADING"}, + {"file_ino_no": "f-ex", "file_name": "ex.pdf", + "first_upload_time": 1718000100, "status": "EXTRACTING"}, + {"file_ino_no": "f-pr", "file_name": "pr.pdf", + "first_upload_time": 1718000200, "status": "PROCESSING"}, + {"file_ino_no": "f-done", "file_name": "done.txt", + "first_upload_time": 1718000300, "status": "COMPLETED"}, + ] + } + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + return_value=self._CHANNELS), \ + patch.object(aidp_mgmt_app, "list_aidp_doc_history_impl", + return_value=history): + response = client.get( + "/aidp-mgmt/knowledge-bases/kb-1/documents", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + body = response.json() + assert body["processing_count"] == 3 + assert [item["status"] for item in body["value"]] == [ + "COMPLETED", "PROCESSING", "EXTRACTING", "UPLOADING", + ] + + def test_unexpected_history_error_falls_back_to_completed_listing(self): + """A non-AppException bug in the history path must not break the list.""" + client = _client() + from ext_components.aidp.apps import aidp_mgmt_app + from ext_components.aidp.services import aidp_permission_service + + with patch.object(aidp_permission_service, "require_permission", + return_value=self._read_only()), \ + patch.object(aidp_mgmt_app, "get_cached_aidp_channels", + side_effect=TypeError("unexpected")), \ + patch.object(aidp_mgmt_app, "list_aidp_docs_impl", + return_value={"value": [{"file_name": "done.txt"}]}), \ + patch.object(aidp_mgmt_app, "count_aidp_docs_impl", return_value=1): + response = client.get( + "/aidp-mgmt/knowledge-bases/kb-1/documents", + headers=_bearer(), + ) + + assert response.status_code == HTTPStatus.OK + assert response.json()["value"] == [{"file_name": "done.txt"}] diff --git a/test/ext_components/aidp/test_aidp_service.py b/test/ext_components/aidp/test_aidp_service.py index e8a6c6790..46decb832 100644 --- a/test/ext_components/aidp/test_aidp_service.py +++ b/test/ext_components/aidp/test_aidp_service.py @@ -1,8 +1,9 @@ import importlib.util +import logging import os import sys from types import ModuleType -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import httpx import pytest @@ -277,6 +278,115 @@ def test_fetches_pages_from_count_and_uses_configured_tenant(self, aidp_service_ "http://127.0.0.1:30081/KnowledgeBase/Tenants/aidp/KnowledgeBases?page=2&page_size=100", ] + def test_forwards_url_encoded_keyword_on_every_page(self, aidp_service_module): + """The keyword rides along on every list request, percent-encoded.""" + aidp_service_module.count_aidp_kbs_impl.return_value = 101 + page1 = MagicMock(status_code=200) + page1.raise_for_status.return_value = None + page1.json.return_value = {"value": [{"kds_id": "kb-1"}], "next_link": None} + page2 = MagicMock(status_code=200) + page2.raise_for_status.return_value = None + page2.json.return_value = {"value": [{"kds_id": "kb-1"}], "next_link": None} + mock_client = MagicMock() + mock_client.get.side_effect = [page1, page2] + mock_manager = MagicMock() + mock_manager.get_sync_client.return_value = mock_client + aidp_service_module.http_client_manager = mock_manager + + aidp_service_module.fetch_all_aidp_knowledge_bases_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + keyword="季度 报告", + ) + + encoded = "%E5%AD%A3%E5%BA%A6%20%E6%8A%A5%E5%91%8A" + requested_urls = [call.args[0] for call in mock_client.get.call_args_list] + assert requested_urls == [ + "http://127.0.0.1:30081/KnowledgeBase/Tenants/aidp/KnowledgeBases" + f"?page=1&page_size=100&keyword={encoded}", + "http://127.0.0.1:30081/KnowledgeBase/Tenants/aidp/KnowledgeBases" + f"?page=2&page_size=100&keyword={encoded}", + ] + + def test_omits_keyword_parameter_when_blank(self, aidp_service_module): + """A blank keyword must produce the same URL as an unfiltered call.""" + aidp_service_module.count_aidp_kbs_impl.return_value = 1 + page1 = MagicMock(status_code=200) + page1.raise_for_status.return_value = None + page1.json.return_value = {"value": [{"kds_id": "kb-1"}], "next_link": None} + mock_client = MagicMock() + mock_client.get.return_value = page1 + mock_manager = MagicMock() + mock_manager.get_sync_client.return_value = mock_client + aidp_service_module.http_client_manager = mock_manager + + aidp_service_module.fetch_all_aidp_knowledge_bases_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + keyword=" ", + ) + + assert mock_client.get.call_args.args[0] == ( + "http://127.0.0.1:30081/KnowledgeBase/Tenants/aidp/KnowledgeBases" + "?page=1&page_size=100" + ) + + def test_keyword_ignored_by_upstream_is_applied_locally(self, aidp_service_module): + """A build that ignores ``?keyword=`` must not return the full catalog.""" + aidp_service_module.count_aidp_kbs_impl.return_value = 2 + page1 = MagicMock(status_code=200) + page1.raise_for_status.return_value = None + page1.json.return_value = { + "value": [ + {"kds_id": "kb-faq", "kds_name": "AIDP FAQ"}, + {"kds_id": "kb-java", "kds_name": "Java Migration Notes"}, + ], + "next_link": None, + } + mock_client = MagicMock() + mock_client.get.return_value = page1 + mock_manager = MagicMock() + mock_manager.get_sync_client.return_value = mock_client + aidp_service_module.http_client_manager = mock_manager + + result = aidp_service_module.fetch_all_aidp_knowledge_bases_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + keyword="faq", + ) + + assert [item["kds_id"] for item in result["value"]] == ["kb-faq"] + assert result["total_count"] == 1 + + def test_upstream_filtered_result_is_left_untouched(self, aidp_service_module): + """Once AIDP narrowed the set, its own matches stay authoritative.""" + aidp_service_module.count_aidp_kbs_impl.return_value = 101 + page1 = MagicMock(status_code=200) + page1.raise_for_status.return_value = None + # A match made on AIDP's own terms (for example tokenized) that a plain + # substring check would miss — it must survive the fallback. + page1.json.return_value = { + "value": [{"kds_id": "kb-1", "kds_name": "Quarterly"}], + "next_link": None, + } + page2 = MagicMock(status_code=200) + page2.raise_for_status.return_value = None + page2.json.return_value = {"value": [], "next_link": None} + mock_client = MagicMock() + mock_client.get.side_effect = [page1, page2] + mock_manager = MagicMock() + mock_manager.get_sync_client.return_value = mock_client + aidp_service_module.http_client_manager = mock_manager + + result = aidp_service_module.fetch_all_aidp_knowledge_bases_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + keyword="reports", + ) + + assert [item["kds_id"] for item in result["value"]] == ["kb-1"] + assert result["total_count"] == 1 + def test_deduplicates_kds_ids_across_pages(self, aidp_service_module): aidp_service_module.count_aidp_kbs_impl.return_value = 101 page1 = MagicMock(status_code=200) @@ -1091,6 +1201,16 @@ def _setup_mock_client(aidp_service_module, method="get", response=None, side_ef return mock_client +def _setup_mock_async_client(aidp_service_module, response=None, side_effect=None): + """Create and wire an async mock client into the service module manager.""" + mock_client = MagicMock() + mock_client.send = AsyncMock(side_effect=side_effect, return_value=response) + mock_manager = MagicMock() + mock_manager.get_async_client.return_value = mock_client + aidp_service_module.http_client_manager = mock_manager + return mock_client + + def _make_http_error(status_code, method="GET"): """Create an httpx.HTTPStatusError with given status code.""" request = httpx.Request(method, "http://127.0.0.1:30081") @@ -1995,7 +2115,7 @@ def test_invalid_config(self, aidp_service_module, server_url, api_key): def test_success_normalizes_docs(self, aidp_service_module): mock_resp = _make_success_response({ "value": [ - {"name": "doc1", "first_upload_time": 1700000000}, + {"name": "doc1", "file_uuid": "uuid-1", "first_upload_time": 1700000000}, {"name": "doc2", "create_time": 1700100000, "update_time": 1700200000}, ], "total_count": 2, @@ -2012,6 +2132,7 @@ def test_success_normalizes_docs(self, aidp_service_module): assert len(result["value"]) == 2 # Normalization adds created_at / updated_at assert result["value"][0]["created_at"] is not None + assert result["value"][0]["file_uuid"] == "uuid-1" assert result["value"][1]["updated_at"] is not None def test_success_non_list_value_not_normalized(self, aidp_service_module): @@ -2081,6 +2202,147 @@ def test_json_parse_value_error(self, aidp_service_module): assert exc_info.value.error_code == ErrorCode.AIDP_RESPONSE_ERROR +# --------------------------------------------------------------------------- +# remove_aidp_docs_impl / download_aidp_doc_impl tests +# --------------------------------------------------------------------------- +class TestAidpDocumentFileOperations: + def test_remove_sends_uuid_array_and_preserves_partial_result( + self, aidp_service_module + ): + expected = { + "summary": {"total": 2, "success": 1, "failed": 1}, + "success_list": [{"file_uuid": "uuid-1"}], + "failed_list": [{"file_uuid": "uuid-2"}], + } + mock_client = _setup_mock_client( + aidp_service_module, + method="post", + response=_make_success_response(expected), + ) + + result = aidp_service_module.remove_aidp_docs_impl( + "http://127.0.0.1:30081", + "jwt-token", + "kb-1", + ["uuid-1", "uuid-2"], + ) + + assert result == expected + call = mock_client.post.call_args + assert call.args[0].endswith( + "/KnowledgeBase/Tenants/aidp/KnowledgeBases/kb-1/KnowledgeFiles/Remove" + ) + assert call.kwargs["json"] == {"file_uuids": ["uuid-1", "uuid-2"]} + + def test_remove_maps_request_error(self, aidp_service_module): + request = httpx.Request("POST", "http://127.0.0.1:30081") + _setup_mock_client( + aidp_service_module, + method="post", + side_effect=httpx.RequestError("network down", request=request), + ) + + with pytest.raises(AppException) as exc_info: + aidp_service_module.remove_aidp_docs_impl( + "http://127.0.0.1:30081", "jwt-token", "kb-1", ["uuid-1"] + ) + assert exc_info.value.error_code == ErrorCode.AIDP_CONNECTION_ERROR + + def test_remove_maps_invalid_json(self, aidp_service_module): + mock_response = _make_success_response({}) + mock_response.json.side_effect = ValueError("bad json") + _setup_mock_client(aidp_service_module, method="post", response=mock_response) + + with pytest.raises(AppException) as exc_info: + aidp_service_module.remove_aidp_docs_impl( + "http://127.0.0.1:30081", "jwt-token", "kb-1", ["uuid-1"] + ) + assert exc_info.value.error_code == ErrorCode.AIDP_RESPONSE_ERROR + + @pytest.mark.parametrize("status_code", [401, 403, 500]) + def test_remove_maps_upstream_http_errors(self, aidp_service_module, status_code): + _setup_mock_client( + aidp_service_module, + method="post", + side_effect=_make_http_error(status_code, "POST"), + ) + + with pytest.raises(AppException) as exc_info: + aidp_service_module.remove_aidp_docs_impl( + "http://127.0.0.1:30081", "jwt-token", "kb-1", ["uuid-1"] + ) + expected_code = ( + ErrorCode.AIDP_AUTH_ERROR + if status_code in (401, 403) + else ErrorCode.AIDP_SERVICE_ERROR + ) + assert exc_info.value.error_code == expected_code + + @pytest.mark.asyncio + async def test_download_streams_binary_content_and_headers(self, aidp_service_module): + mock_response = httpx.Response( + 200, + headers={ + "Content-Type": "text/plain", + "Content-Disposition": 'attachment; filename="a.txt"', + "X-File-Size": "16", + }, + content=b"downloaded bytes", + request=httpx.Request("POST", "http://127.0.0.1:30081"), + ) + mock_client = _setup_mock_async_client(aidp_service_module, response=mock_response) + + response = await aidp_service_module.stream_aidp_doc_impl( + "http://127.0.0.1:30081", "jwt-token", "kb-1", "uuid-1" + ) + chunks = [chunk async for chunk in response.aiter_bytes()] + + assert b"".join(chunks) == b"downloaded bytes" + assert response.headers["Content-Type"] == "text/plain" + assert response.headers["Content-Disposition"] == 'attachment; filename="a.txt"' + assert response.headers["X-File-Size"] == "16" + request = mock_client.build_request.call_args + assert request.args[0] == "POST" + assert request.args[1].endswith( + "/KnowledgeBase/Tenants/aidp/KnowledgeBases/kb-1/KnowledgeFiles/Download" + ) + assert request.kwargs["json"] == {"file_uuid": "uuid-1"} + await response.aclose() + assert mock_response.is_closed + + @pytest.mark.asyncio + async def test_download_maps_request_error(self, aidp_service_module): + request = httpx.Request("POST", "http://127.0.0.1:30081") + _setup_mock_async_client( + aidp_service_module, + side_effect=httpx.RequestError("network down", request=request), + ) + + with pytest.raises(AppException) as exc_info: + await aidp_service_module.stream_aidp_doc_impl( + "http://127.0.0.1:30081", "jwt-token", "kb-1", "uuid-1" + ) + assert exc_info.value.error_code == ErrorCode.AIDP_CONNECTION_ERROR + + @pytest.mark.asyncio + async def test_download_maps_upstream_http_error(self, aidp_service_module): + response = httpx.Response( + 404, + json={"error": "file not found"}, + request=httpx.Request("POST", "http://127.0.0.1:30081"), + ) + _setup_mock_async_client( + aidp_service_module, + response=response, + ) + + with pytest.raises(AppException) as exc_info: + await aidp_service_module.stream_aidp_doc_impl( + "http://127.0.0.1:30081", "jwt-token", "kb-1", "uuid-1" + ) + assert exc_info.value.error_code == ErrorCode.AIDP_SERVICE_ERROR + + # --------------------------------------------------------------------------- # list_aidp_models_impl tests # --------------------------------------------------------------------------- @@ -2202,3 +2464,620 @@ def test_json_parse_value_error(self, aidp_service_module): server_url="http://127.0.0.1:30081", api_key="jwt-token" ) assert exc_info.value.error_code == ErrorCode.AIDP_RESPONSE_ERROR + + +# --------------------------------------------------------------------------- +# select_aidp_channel tests +# --------------------------------------------------------------------------- +class TestSelectAidpChannel: + """Tests for select_aidp_channel (channel -> history addressing helper).""" + + def test_matches_channel_by_explicit_kds_id(self, aidp_service_module): + channels = [ + {"fs_id": "fs-a", "src_dir": "/dir/a", "kds_id": "kb-a"}, + {"fs_id": "fs-b", "src_dir": "/dir/b", "kds_id": "kb-b"}, + ] + assert aidp_service_module.select_aidp_channel(channels, "kb-b") == { + "fs_id": "fs-b", + "src_dir": "/dir/b", + } + + def test_matches_channel_by_source_directory(self, aidp_service_module): + channels = [ + {"fs_id": "fs-a", "src_dir": "/dir/a"}, + {"fs_id": "fs-b", "src_dir": "/knowledge/kb-b"}, + ] + assert aidp_service_module.select_aidp_channel(channels, "kb-b") == { + "fs_id": "fs-b", + "src_dir": "/knowledge/kb-b", + } + + def test_matches_channel_by_id_list(self, aidp_service_module): + channels = [ + {"fs_id": "fs-a", "src_dir": "/dir/a", "knowledge_base_ids": ["kb-x", "kb-y"]}, + ] + assert aidp_service_module.select_aidp_channel(channels, "kb-y")["fs_id"] == "fs-a" + + def test_accepts_alias_field_names(self, aidp_service_module): + channels = [{"fsId": "fs-alias", "srcDir": "/dir/alias"}] + assert aidp_service_module.select_aidp_channel(channels, "kb-1") == { + "fs_id": "fs-alias", + "src_dir": "/dir/alias", + } + + def test_falls_back_to_first_usable_channel(self, aidp_service_module): + """Single-channel deployments are addressed by the only channel there is.""" + channels = [ + {"fs_id": "fs-only", "src_dir": "/dir/only"}, + ] + assert aidp_service_module.select_aidp_channel(channels, "unrelated-kb") == { + "fs_id": "fs-only", + "src_dir": "/dir/only", + } + + def test_reads_nested_addressing_fields(self, aidp_service_module): + """A channel that nests fs_id/src_dir still resolves.""" + channels = [ + { + "id": "channel-2", + "config": {"fs_id": "fs-nested", "src_dir": "/knowledge/kb-nested"}, + }, + ] + assert aidp_service_module.select_aidp_channel(channels, "kb-nested") == { + "fs_id": "fs-nested", + "src_dir": "/knowledge/kb-nested", + } + + def test_matches_channel_by_undocumented_field(self, aidp_service_module): + """A channel naming its knowledge base in any field is still matched.""" + channels = [ + {"fs_id": "fs-a", "src_dir": "/dir/a", "channel_name": "kb-named"}, + {"fs_id": "fs-b", "src_dir": "/dir/b", "channel_name": "other"}, + ] + assert aidp_service_module.select_aidp_channel(channels, "kb-named") == { + "fs_id": "fs-a", + "src_dir": "/dir/a", + } + + def test_skips_entries_missing_fs_id_or_dir(self, aidp_service_module): + channels = [ + {"fs_id": "", "src_dir": "/dir/empty"}, + {"fs_id": "fs-a", "src_dir": ""}, + {"kds_id": "kb-a"}, + {"fs_id": "fs-ok", "src_dir": "/dir/ok"}, + ] + assert aidp_service_module.select_aidp_channel(channels, "kb-a") == { + "fs_id": "fs-ok", + "src_dir": "/dir/ok", + } + + def test_returns_none_when_nothing_usable(self, aidp_service_module): + assert aidp_service_module.select_aidp_channel([], "kb-1") is None + assert aidp_service_module.select_aidp_channel( + [{"fs_id": "fs-a"}, "not-a-dict"], "kb-1" + ) is None + + +# --------------------------------------------------------------------------- +# list_aidp_channels_impl tests +# --------------------------------------------------------------------------- +class TestListAidpChannelsImpl: + """Tests for list_aidp_channels_impl (GET .../KnowledgeBases/{id}/Channels).""" + + _KB = "kb-1" + + @pytest.mark.parametrize( + "server_url,api_key", + [("", "token"), ("ftp://bad", "token"), ("http://ok", "")], + ) + def test_invalid_config(self, aidp_service_module, server_url, api_key): + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_channels_impl( + server_url=server_url, api_key=api_key, kds_id=self._KB + ) + assert exc_info.value.error_code == ErrorCode.AIDP_CONFIG_INVALID + + @pytest.mark.parametrize("kds_id", ["", " "]) + def test_missing_knowledge_base_id_raises_without_request( + self, aidp_service_module, kds_id + ): + """The catalog is KB-scoped: a request without a KB id cannot be built.""" + mock_client = _setup_mock_client(aidp_service_module, method="get") + + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=kds_id, + ) + + assert exc_info.value.error_code == ErrorCode.AIDP_CONFIG_INVALID + mock_client.get.assert_not_called() + + def test_success_requests_knowledge_base_scoped_path(self, aidp_service_module): + mock_resp = _make_success_response({ + "value": [{"fs_id": "fs-1", "src_dir": "/dir/1"}], + }) + mock_client = _setup_mock_client( + aidp_service_module, method="get", response=mock_resp + ) + + result = aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + + assert result["value"] == [{"fs_id": "fs-1", "src_dir": "/dir/1"}] + call_args = mock_client.get.call_args + assert call_args[0][0] == ( + "http://127.0.0.1:30081/KnowledgeBase/Tenants/aidp/KnowledgeBases/kb-1/Channels" + ) + assert call_args.kwargs["headers"]["Authorization"] == "Bearer jwt-token" + + def test_tenant_override_is_forwarded(self, aidp_service_module): + mock_resp = _make_success_response({"value": []}) + mock_client = _setup_mock_client( + aidp_service_module, method="get", response=mock_resp + ) + + aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + tenant_id="other", + ) + + assert ( + "/KnowledgeBase/Tenants/other/KnowledgeBases/kb-1/Channels" + in mock_client.get.call_args[0][0] + ) + + def test_non_dict_entries_are_dropped(self, aidp_service_module): + mock_resp = _make_success_response({"value": ["junk", {"fs_id": "fs-1"}]}) + _setup_mock_client(aidp_service_module, method="get", response=mock_resp) + + result = aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + + assert result["value"] == [{"fs_id": "fs-1"}] + + @pytest.mark.parametrize("key", ["data", "result", "records", "items", "list"]) + def test_list_envelopes_are_normalized_to_value( + self, aidp_service_module, key + ): + """A renamed list wrapper must not read as 'no channels upstream'.""" + mock_resp = _make_success_response({key: [{"fs_id": "fs-1", "src_dir": "/dir/1"}]}) + _setup_mock_client(aidp_service_module, method="get", response=mock_resp) + + result = aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + + assert result["value"] == [{"fs_id": "fs-1", "src_dir": "/dir/1"}] + + def test_bare_list_response_is_accepted(self, aidp_service_module): + mock_resp = _make_success_response([{"fs_id": "fs-1", "src_dir": "/dir/1"}]) + _setup_mock_client(aidp_service_module, method="get", response=mock_resp) + + result = aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + + assert result["value"] == [{"fs_id": "fs-1", "src_dir": "/dir/1"}] + + @pytest.mark.parametrize("payload", [{"value": "not-a-list"}, "junk", 123]) + def test_payload_without_any_list_raises(self, aidp_service_module, payload): + mock_resp = _make_success_response(payload) + _setup_mock_client(aidp_service_module, method="get", response=mock_resp) + + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + assert exc_info.value.error_code == ErrorCode.AIDP_RESPONSE_ERROR + + @pytest.mark.parametrize("status_code,expected", [ + (401, ErrorCode.AIDP_AUTH_ERROR), + (403, ErrorCode.AIDP_AUTH_ERROR), + (429, ErrorCode.AIDP_RATE_LIMIT), + (500, ErrorCode.AIDP_SERVICE_ERROR), + ]) + def test_http_errors_are_mapped(self, aidp_service_module, status_code, expected): + _setup_mock_client( + aidp_service_module, method="get", + side_effect=_make_http_error(status_code), + ) + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + assert exc_info.value.error_code == expected + + def test_connection_error(self, aidp_service_module): + request = httpx.Request("GET", "http://127.0.0.1:30081") + _setup_mock_client( + aidp_service_module, method="get", + side_effect=httpx.RequestError("network down", request=request), + ) + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + assert exc_info.value.error_code == ErrorCode.AIDP_CONNECTION_ERROR + + def test_json_parse_value_error(self, aidp_service_module): + mock_resp = _make_success_response({}) + mock_resp.json.side_effect = ValueError("bad json") + _setup_mock_client(aidp_service_module, method="get", response=mock_resp) + + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_channels_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + kds_id=self._KB, + ) + assert exc_info.value.error_code == ErrorCode.AIDP_RESPONSE_ERROR + + +# --------------------------------------------------------------------------- +# list_aidp_doc_history_impl tests +# --------------------------------------------------------------------------- +class TestListAidpDocHistoryImpl: + """Tests for list_aidp_doc_history_impl (KB-scoped KnowledgeFiles/History).""" + + _KB = "kb-1" + + @pytest.mark.parametrize("fs_id,dir_path", [ + ("", "/dir/1"), + ("fs-1", ""), + (" ", "/dir/1"), + ("fs-1", " "), + ]) + def test_missing_address_params_raise_without_request( + self, aidp_service_module, fs_id, dir_path + ): + mock_client = _setup_mock_client(aidp_service_module, method="post") + + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id=fs_id, + dir_path=dir_path, + kds_id=self._KB, + ) + assert exc_info.value.error_code == ErrorCode.AIDP_CONFIG_INVALID + mock_client.post.assert_not_called() + + @pytest.mark.parametrize("kds_id", ["", " "]) + def test_missing_knowledge_base_id_raises_without_request( + self, aidp_service_module, kds_id + ): + """The endpoint is KB-scoped: a request without a KB id cannot be built.""" + mock_client = _setup_mock_client(aidp_service_module, method="post") + + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=kds_id, + ) + + assert exc_info.value.error_code == ErrorCode.AIDP_CONFIG_INVALID + mock_client.post.assert_not_called() + + def test_success_posts_body_and_normalizes_status(self, aidp_service_module): + mock_resp = _make_success_response({ + "value": [ + { + "file_ino_no": "f-1", + "file_name": "doc.pdf", + "first_upload_time": 1718000000, + "status": "processing", + }, + { + "file_ino_no": "f-2", + "file_name": "done.txt", + "first_upload_time": 1718000100, + "status": " COMPLETED ", + }, + ], + }) + mock_client = _setup_mock_client( + aidp_service_module, method="post", response=mock_resp + ) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/aidp/knowledge/kb-1", + kds_id=self._KB, + ) + + assert [item["status"] for item in result["value"]] == [ + "PROCESSING", "COMPLETED", + ] + # Timestamps are normalized like every other document payload. + assert result["value"][0]["created_at"] is not None + + call_args = mock_client.post.call_args + assert call_args[0][0] == ( + "http://127.0.0.1:30081/KnowledgeBase/Tenants/aidp" + "/KnowledgeBases/kb-1/KnowledgeFiles/History" + ) + assert call_args.kwargs["json"] == { + "fs_id": "fs-1", + "dir_path": "/aidp/knowledge/kb-1", + } + assert call_args.kwargs["headers"]["Authorization"] == "Bearer jwt-token" + + def test_tenant_override_is_forwarded(self, aidp_service_module): + mock_resp = _make_success_response({"value": []}) + mock_client = _setup_mock_client( + aidp_service_module, method="post", response=mock_resp + ) + + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + tenant_id="other", + ) + + assert ( + "/KnowledgeBase/Tenants/other/KnowledgeBases/kb-1/KnowledgeFiles/History" + in mock_client.post.call_args[0][0] + ) + + def test_items_without_status_keep_no_status_key(self, aidp_service_module): + mock_resp = _make_success_response({ + "value": [{"file_ino_no": "f-1", "file_name": "doc.txt"}], + }) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert "status" not in result["value"][0] + + def test_blank_status_is_not_reported(self, aidp_service_module): + mock_resp = _make_success_response({ + "value": [{"file_ino_no": "f-1", "status": " "}], + }) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert "status" not in result["value"][0] + + @pytest.mark.parametrize( + "alias", ["file_status", "fileStatus", "file_state", "doc_status"] + ) + def test_status_aliases_are_emitted_under_status( + self, aidp_service_module, alias + ): + """A deployment-specific status field still fills the status column.""" + mock_resp = _make_success_response({ + "value": [ + {"file_ino_no": "f-1", "file_name": "doc.pdf", alias: " processing "}, + ], + }) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert result["value"][0]["status"] == "PROCESSING" + + def test_canonical_status_wins_over_alias(self, aidp_service_module): + mock_resp = _make_success_response({ + "value": [ + { + "file_ino_no": "f-1", + "status": "COMPLETED", + "file_status": "FAILED", + }, + ], + }) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert result["value"][0]["status"] == "COMPLETED" + + def test_non_string_status_values_are_ignored(self, aidp_service_module): + """Numeric state vocabularies must not be mistaken for a file status.""" + mock_resp = _make_success_response({ + "value": [{"file_ino_no": "f-1", "file_state": 4, "status": None}], + }) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert "status" not in result["value"][0] + + def test_unreadable_status_payload_is_reported( + self, aidp_service_module, caplog + ): + """The real field name must be discoverable from the log.""" + mock_resp = _make_success_response({ + "value": [{"file_ino_no": "f-1", "upload_phase": "chunking"}], + }) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + with caplog.at_level(logging.WARNING): + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert "no readable status" in caplog.text + assert "upload_phase" in caplog.text + + @pytest.mark.parametrize( + "key", ["data", "result", "records", "items", "list"] + ) + def test_list_envelopes_are_normalized_to_value(self, aidp_service_module, key): + mock_resp = _make_success_response({ + key: [{"file_ino_no": "f-1", "status": "processing"}], + }) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert [item["status"] for item in result["value"]] == ["PROCESSING"] + + @pytest.mark.parametrize("payload", [{"value": "not-a-list"}, "junk"]) + def test_payload_without_any_list_raises(self, aidp_service_module, payload): + mock_resp = _make_success_response(payload) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + assert exc_info.value.error_code == ErrorCode.AIDP_RESPONSE_ERROR + + def test_response_without_list_carries_empty_value(self, aidp_service_module): + """An empty directory is a valid answer, not a payload error.""" + mock_resp = _make_success_response({"value": []}) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert result["value"] == [] + + def test_bare_list_response_is_accepted(self, aidp_service_module): + """Some builds answer with the array itself instead of an envelope.""" + mock_resp = _make_success_response( + [{"file_ino_no": "f-1", "status": "completed"}] + ) + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + result = aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + + assert [item["status"] for item in result["value"]] == ["COMPLETED"] + + @pytest.mark.parametrize("status_code,expected", [ + (401, ErrorCode.AIDP_AUTH_ERROR), + (403, ErrorCode.AIDP_AUTH_ERROR), + (429, ErrorCode.AIDP_RATE_LIMIT), + (500, ErrorCode.AIDP_SERVICE_ERROR), + ]) + def test_http_errors_are_mapped(self, aidp_service_module, status_code, expected): + _setup_mock_client( + aidp_service_module, method="post", + side_effect=_make_http_error(status_code, method="POST"), + ) + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + assert exc_info.value.error_code == expected + + def test_connection_error(self, aidp_service_module): + request = httpx.Request("POST", "http://127.0.0.1:30081") + _setup_mock_client( + aidp_service_module, method="post", + side_effect=httpx.RequestError("network down", request=request), + ) + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + assert exc_info.value.error_code == ErrorCode.AIDP_CONNECTION_ERROR + + def test_json_parse_value_error(self, aidp_service_module): + mock_resp = _make_success_response({}) + mock_resp.json.side_effect = ValueError("bad json") + _setup_mock_client(aidp_service_module, method="post", response=mock_resp) + + with pytest.raises(AppException) as exc_info: + aidp_service_module.list_aidp_doc_history_impl( + server_url="http://127.0.0.1:30081", + api_key="jwt-token", + fs_id="fs-1", + dir_path="/dir/1", + kds_id=self._KB, + ) + assert exc_info.value.error_code == ErrorCode.AIDP_RESPONSE_ERROR From 0a0577274cac3178e88c8e6b4dfae07ee7484b9d Mon Sep 17 00:00:00 2001 From: ljy Date: Sun, 20 Sep 2026 18:50:45 +0800 Subject: [PATCH 8/9] Fix: editing a ModelEngine model no longer flips ssl_verify to True - the update path only checked api_key emptiness while the create path also exempts open/router URLs (ModelEngine self-signed certs). The edit dialog prefills the real key and always submits it, so any edit silently broke connectivity. Exemption now checks the payload URL with a fallback to the stored record; batch-edit groups get the same protection --- backend/services/model_management_service.py | 14 ++++- .../services/test_model_management_service.py | 58 +++++++++++++++++++ 2 files changed, 71 insertions(+), 1 deletion(-) diff --git a/backend/services/model_management_service.py b/backend/services/model_management_service.py index 5928303c6..c35cf46f9 100644 --- a/backend/services/model_management_service.py +++ b/backend/services/model_management_service.py @@ -784,9 +784,21 @@ async def update_single_model_for_tenant( # Auto-set ssl_verify based on api_key if provided: # - Empty api_key -> ssl_verify=False + # - "open/router" URL (ModelEngine, self-signed certs) -> ssl_verify=False # - Otherwise -> ssl_verify=True + # The open/router exemption mirrors the create path + # (create_model_for_tenant): without it, editing a ModelEngine model + # submits the prefilled non-empty api_key and silently flips + # ssl_verify to True, breaking connectivity against its self-signed + # certificate. The URL is taken from the update payload when present + # and falls back to the stored record otherwise. if "api_key" in model_data: - if not model_data["api_key"]: + effective_base_url = ( + model_data.get("base_url") + or (existing_models[0].get("base_url") if existing_models else "") + or "" + ) + if not model_data["api_key"] or "open/router" in effective_base_url: model_data["ssl_verify"] = False else: model_data["ssl_verify"] = True diff --git a/test/backend/services/test_model_management_service.py b/test/backend/services/test_model_management_service.py index 29edd3ce6..d7cef61cd 100644 --- a/test/backend/services/test_model_management_service.py +++ b/test/backend/services/test_model_management_service.py @@ -1784,6 +1784,64 @@ async def test_update_single_model_for_tenant_empty_api_key_sets_ssl_verify_fals assert update_call[0][1]["ssl_verify"] is False +async def test_update_single_model_for_tenant_open_router_url_keeps_ssl_verify_false(): + """Editing a ModelEngine model must not flip ssl_verify to True. + + ModelEngine endpoints serve self-signed certificates, so their records + are created with ssl_verify=False. The edit dialog prefills the real + api_key and always submits it; without the open/router exemption the + update path would silently flip ssl_verify to True and break + connectivity. Mirrors the create-path exemption in + create_model_for_tenant. + """ + svc = import_svc() + + existing_models = [ + { + "model_id": 1, + "model_type": "llm", + "display_name": "name", + "base_url": "https://141.111.135.222:30012/open/router/v1", + }, + ] + model_data = { + "model_id": 1, + "display_name": "name", + "api_key": "my-secret-key", + } + + with mock.patch.object(svc, "get_models_by_display_name", return_value=existing_models), \ + mock.patch.object(svc, "update_model_record") as mock_update: + + await svc.update_single_model_for_tenant("u1", "t1", "name", model_data) + + update_call = mock_update.call_args + assert update_call[0][1]["ssl_verify"] is False + + +async def test_update_single_model_for_tenant_open_router_url_in_payload_keeps_ssl_verify_false(): + """The open/router exemption also applies when the update payload carries the URL itself.""" + svc = import_svc() + + existing_models = [ + {"model_id": 1, "model_type": "llm", "display_name": "name"}, + ] + model_data = { + "model_id": 1, + "display_name": "name", + "api_key": "my-secret-key", + "base_url": "https://example.com/open/router/v1", + } + + with mock.patch.object(svc, "get_models_by_display_name", return_value=existing_models), \ + mock.patch.object(svc, "update_model_record") as mock_update: + + await svc.update_single_model_for_tenant("u1", "t1", "name", model_data) + + update_call = mock_update.call_args + assert update_call[0][1]["ssl_verify"] is False + + async def test_update_single_model_for_tenant_generic_exception(): """Test that generic exceptions are caught and re-raised (covers lines 329-331).""" svc = import_svc() From 3125e4ebec48e270105bd27cd7bdc9a04d6a7b5f Mon Sep 17 00:00:00 2001 From: ljy Date: Sun, 20 Sep 2026 19:07:24 +0800 Subject: [PATCH 9/9] refactor: extract MODEL_ENGINE_URL_MARKER constant (SonarCloud S1192) and use a placeholder domain in test fixtures - no behavior change --- backend/consts/provider.py | 5 +++++ backend/services/model_management_service.py | 13 +++++++------ .../services/test_model_management_service.py | 3 ++- 3 files changed, 14 insertions(+), 7 deletions(-) diff --git a/backend/consts/provider.py b/backend/consts/provider.py index fe49332b7..63ffb15f2 100644 --- a/backend/consts/provider.py +++ b/backend/consts/provider.py @@ -26,3 +26,8 @@ class ProviderEnum(str, Enum): # ModelEngine # Base URL and API key are loaded from environment variables at runtime +# URL path segment identifying ModelEngine northbound endpoints +# (e.g. https://host:port/open/router/v1). Endpoints behind this marker +# serve self-signed certificates, so their records must keep +# ssl_verify=False on both the create and update paths. +MODEL_ENGINE_URL_MARKER = "open/router" diff --git a/backend/services/model_management_service.py b/backend/services/model_management_service.py index c35cf46f9..ff0b871a0 100644 --- a/backend/services/model_management_service.py +++ b/backend/services/model_management_service.py @@ -19,6 +19,7 @@ DASHSCOPE_BASE_URL, DASHSCOPE_REALTIME_BASE_URL, TOKENPONY_BASE_URL, + MODEL_ENGINE_URL_MARKER, ) from database.model_management_db import ( @@ -340,16 +341,16 @@ async def create_model_for_tenant(user_id: str, tenant_id: str, model_data: Dict ) # Auto-set ssl_verify based on api_key: # - Empty api_key (local/LAN services) -> ssl_verify=False - # - "open/router" URL -> ssl_verify=False + # - ModelEngine URL (self-signed certs) -> ssl_verify=False # - Otherwise -> ssl_verify=True model_api_key = model_data.get("api_key", "") - if not model_api_key or "open/router" in model_base_url: + if not model_api_key or MODEL_ENGINE_URL_MARKER in model_base_url: model_data["ssl_verify"] = False else: model_data["ssl_verify"] = True - # Set model_factory to modelengine when using open/router URL - if "open/router" in model_base_url: + # Set model_factory to modelengine when using a ModelEngine URL + if MODEL_ENGINE_URL_MARKER in model_base_url: model_data["model_factory"] = "modelengine" if model_data.get("model_type") in ("vlm", "vlm2", "vlm3", "vlm4"): @@ -784,7 +785,7 @@ async def update_single_model_for_tenant( # Auto-set ssl_verify based on api_key if provided: # - Empty api_key -> ssl_verify=False - # - "open/router" URL (ModelEngine, self-signed certs) -> ssl_verify=False + # - ModelEngine URL (self-signed certs) -> ssl_verify=False # - Otherwise -> ssl_verify=True # The open/router exemption mirrors the create path # (create_model_for_tenant): without it, editing a ModelEngine model @@ -798,7 +799,7 @@ async def update_single_model_for_tenant( or (existing_models[0].get("base_url") if existing_models else "") or "" ) - if not model_data["api_key"] or "open/router" in effective_base_url: + if not model_data["api_key"] or MODEL_ENGINE_URL_MARKER in effective_base_url: model_data["ssl_verify"] = False else: model_data["ssl_verify"] = True diff --git a/test/backend/services/test_model_management_service.py b/test/backend/services/test_model_management_service.py index d7cef61cd..45502ffda 100644 --- a/test/backend/services/test_model_management_service.py +++ b/test/backend/services/test_model_management_service.py @@ -211,6 +211,7 @@ class _ProviderEnum: consts_provider_mod.DASHSCOPE_REALTIME_BASE_URL = "wss://dashscope.aliyuncs.com/api-ws/v1/realtime" consts_provider_mod.DASHSCOPE_STT_BASE_URL = consts_provider_mod.DASHSCOPE_REALTIME_BASE_URL consts_provider_mod.TOKENPONY_BASE_URL = "https://api.tokenpony.cn/v1/" +consts_provider_mod.MODEL_ENGINE_URL_MARKER = "open/router" sys.modules["consts.provider"] = consts_provider_mod # Stub services.model_provider_service used by service @@ -1801,7 +1802,7 @@ async def test_update_single_model_for_tenant_open_router_url_keeps_ssl_verify_f "model_id": 1, "model_type": "llm", "display_name": "name", - "base_url": "https://141.111.135.222:30012/open/router/v1", + "base_url": "https://modelengine.example.com/open/router/v1", }, ] model_data = {