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/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/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/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/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/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/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/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/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/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/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/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/backend/services/model_management_service.py b/backend/services/model_management_service.py index 5928303c6..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,9 +785,21 @@ async def update_single_model_for_tenant( # Auto-set ssl_verify based on api_key if provided: # - Empty api_key -> 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 + # 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 MODEL_ENGINE_URL_MARKER in effective_base_url: model_data["ssl_verify"] = False else: model_data["ssl_verify"] = True 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/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/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/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/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]/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]/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/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts b/frontend/app/[locale]/newchat/adapter/remote-chat-model-adapter.ts index e40eead40..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); @@ -2034,11 +2096,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 }); @@ -2050,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" @@ -2058,6 +2126,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 ( @@ -2654,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/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/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/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/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/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/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/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/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 163c6fc20..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)}") @@ -1380,7 +1390,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 +1404,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/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/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/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/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) 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/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/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_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_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/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_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_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_model_management_service.py b/test/backend/services/test_model_management_service.py index 29edd3ce6..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 @@ -1784,6 +1785,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://modelengine.example.com/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() 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 = [] 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() 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"]) 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/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 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 452736a8b..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" @@ -4805,7 +4831,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 +4843,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( 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, + } + ]