From fe4c6f1a6906e8b465387598e702d1c987b25dd2 Mon Sep 17 00:00:00 2001 From: Jasonxia007 Date: Fri, 14 Aug 2026 10:01:42 +0800 Subject: [PATCH 1/5] =?UTF-8?q?=E2=9C=A8=20Parameterize=20file=20split=20f?= =?UTF-8?q?unction=20in=20data-process=20task=20=F0=9F=A7=AA=20Add=20telem?= =?UTF-8?q?etry=20data=20for=20knowledge=20base?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/apps/data_process_app.py | 3 +- backend/consts/const.py | 14 +- backend/consts/model.py | 1 + backend/data_process/ray_actors.py | 46 ++- backend/data_process/tasks.py | 177 ++++++++-- backend/data_process/worker.py | 34 +- backend/data_process_service.py | 8 +- backend/database/attachment_db.py | 4 + backend/services/data_process_service.py | 2 + backend/services/file_management_service.py | 2 + backend/utils/file_management_utils.py | 8 +- backend/utils/knowledge_telemetry.py | 318 ++++++++++++++++++ .../otel-collector-phoenix-config.yml | 48 ++- deploy/env/.env.example | 6 +- .../dockerfiles/data-process/Dockerfile | 2 +- deploy/k8s/deploy.sh | 6 +- .../nexent-common/templates/configmap.yaml | 6 +- .../nexent/charts/nexent-common/values.yaml | 8 +- .../templates/otel-collector-configmap.yaml | 38 ++- sdk/nexent/data_process/core.py | 54 ++- sdk/nexent/data_process/file_splitter.py | 63 +++- sdk/nexent/monitor/monitoring.py | 2 + test/backend/data_process/test_tasks.py | 156 ++++++++- test/backend/data_process/test_worker.py | 2 +- .../backend/utils/test_knowledge_telemetry.py | 63 ++++ test/sdk/data_process/test_core.py | 18 +- test/sdk/data_process/test_file_splitter.py | 53 +++ 27 files changed, 1053 insertions(+), 89 deletions(-) create mode 100644 backend/utils/knowledge_telemetry.py create mode 100644 test/backend/utils/test_knowledge_telemetry.py diff --git a/backend/apps/data_process_app.py b/backend/apps/data_process_app.py index 693eb987ea..9589817bac 100644 --- a/backend/apps/data_process_app.py +++ b/backend/apps/data_process_app.py @@ -58,7 +58,8 @@ async def create_task(request: TaskRequest, authorization: Optional[str] = Heade original_filename=request.original_filename, authorization=authorization, embedding_model_id=request.embedding_model_id, - tenant_id=request.tenant_id + tenant_id=request.tenant_id, + telemetry_context=getattr(request, "telemetry_context", {}) or {}, ) return JSONResponse(status_code=HTTPStatus.CREATED, content={"task_id": task_result.id}) diff --git a/backend/consts/const.py b/backend/consts/const.py index fffda132d0..33b0219ab3 100644 --- a/backend/consts/const.py +++ b/backend/consts/const.py @@ -262,15 +262,13 @@ class VectorDatabaseType(str, Enum): RAY_ACTOR_NUM_CPUS = int(os.getenv("RAY_ACTOR_NUM_CPUS", "2")) RAY_DASHBOARD_PORT = int(os.getenv("RAY_DASHBOARD_PORT", "8265")) RAY_DASHBOARD_HOST = os.getenv("RAY_DASHBOARD_HOST", "0.0.0.0") -RAY_NUM_CPUS = int(os.getenv("RAY_NUM_CPUS", "4")) -RAY_OBJECT_STORE_MEMORY_GB = float( - os.getenv("RAY_OBJECT_STORE_MEMORY_GB", "0.25")) +RAY_NUM_CPUS = DP_PART_PROCESSOR_COUNT * RAY_ACTOR_NUM_CPUS +RAY_OBJECT_STORE_MEMORY_GB = float(os.getenv("RAY_OBJECT_STORE_MEMORY_GB", "0.25")) RAY_TEMP_DIR = os.getenv("RAY_TEMP_DIR", "/tmp/ray") RAY_LOG_LEVEL = os.getenv("RAY_LOG_LEVEL", "INFO").upper() # Disable plasma preallocation to reduce idle memory usage # When set to false, Ray will allocate object store memory on-demand instead of preallocating -RAY_preallocate_plasma = os.getenv( - "RAY_preallocate_plasma", "false").lower() == "true" +RAY_preallocate_plasma = os.getenv("RAY_preallocate_plasma", "false").lower() == "true" # Service Control Flags @@ -299,13 +297,13 @@ class VectorDatabaseType(str, Enum): QUEUES = os.getenv("QUEUES", "process_q,process_part_q,forward_q") # Will be dynamically set based on PID if not provided WORKER_NAME = os.getenv("WORKER_NAME") -WORKER_CONCURRENCY = int(os.getenv("WORKER_CONCURRENCY", "4")) +WORKER_CONCURRENCY = DP_PART_PROCESSOR_COUNT + 1 +DP_PART_PROCESSOR_COUNT = int(os.getenv("DP_PART_PROCESSOR_COUNT", "3")) +DP_FILE_SPLIT_SIZE_MB = int(os.getenv("DP_FILE_SPLIT_SIZE_MB", "5")) RAY_WARM_ACTOR_POOL_SIZE_PART = int( os.getenv("RAY_WARM_ACTOR_POOL_SIZE_PART", "2")) RAY_WARM_ACTOR_POOL_SIZE_PROCESS = int( os.getenv("RAY_WARM_ACTOR_POOL_SIZE_PROCESS", "1")) -# Global Ray actor pool (shared by process_q/process_part_q workers) -RAY_GLOBAL_ACTOR_POOL_SIZE = int(os.getenv("RAY_GLOBAL_ACTOR_POOL_SIZE", "3")) RAY_ACTOR_WARM_TIMEOUT_S = float(os.getenv("RAY_ACTOR_WARM_TIMEOUT_S", "60")) RAY_GLOBAL_ACTOR_POOL_NAME = os.getenv( "RAY_GLOBAL_ACTOR_POOL_NAME", "nexent_global_data_processor_pool") diff --git a/backend/consts/model.py b/backend/consts/model.py index dca8bc1387..71d861d633 100644 --- a/backend/consts/model.py +++ b/backend/consts/model.py @@ -414,6 +414,7 @@ class TaskRequest(BaseModel): original_filename: Optional[str] = None embedding_model_id: Optional[int] = None tenant_id: Optional[str] = None + telemetry_context: Dict[str, str] = Field(default_factory=dict) additional_params: Dict[str, Any] = Field(default_factory=dict) diff --git a/backend/data_process/ray_actors.py b/backend/data_process/ray_actors.py index c3879c007e..19bb03def1 100644 --- a/backend/data_process/ray_actors.py +++ b/backend/data_process/ray_actors.py @@ -32,8 +32,25 @@ class DataProcessorRayActor: """ def __init__(self): + # Ray actors are independent processes and must initialize their own + # provider before DataProcessCore creates typed preprocessing spans. + try: + from utils.monitoring import monitoring_manager + + self._monitoring_manager = monitoring_manager + telemetry_enabled = monitoring_manager.is_enabled + except Exception: + self._monitoring_manager = None + telemetry_enabled = False + logger.warning( + "Knowledge telemetry initialization failed in Ray actor; processing will continue", + exc_info=True, + ) logger.info( - f"Ray actor initialized using {RAY_ACTOR_NUM_CPUS} CPU cores...") + "Ray actor initialized using %s CPU cores; telemetry_enabled=%s", + RAY_ACTOR_NUM_CPUS, + telemetry_enabled, + ) self._processor = DataProcessCore() def ping(self) -> bool: @@ -71,12 +88,22 @@ def _run_file_process( process_params: Dict[str, Any], log_subject: str, ) -> List[Dict[str, Any]]: - result = self._processor.file_process( - file_data=file_data, + from utils.knowledge_telemetry import knowledge_span + + with knowledge_span( + "knowledge.process.ray_actor", + "process.ray_actor", + telemetry_context=process_params.get("telemetry_context"), filename=filename, - chunking_strategy=chunking_strategy, - **process_params - ) + file_size_bytes=len(file_data), + task_id=process_params.get("task_id"), + ): + result = self._processor.file_process( + file_data=file_data, + filename=filename, + chunking_strategy=chunking_strategy, + **process_params + ) chunks, images_info = self._normalize_processor_result(result) if images_info: @@ -304,7 +331,8 @@ def split_file( source: str, destination: str, task_id: Optional[str] = None, - max_size: int = 5 * 1024 * 1024, + max_size: Optional[int] = None, + target_parts: Optional[int] = None, file_data: Optional[bytes] = None, **params ) -> List[bytes]: @@ -312,7 +340,8 @@ def split_file( Split file into parts using DataProcessCore.file_split and return raw bytes list. """ logger.info( - f"[RayActor] Splitting file: source='{source}', destination='{destination}', task_id='{task_id}', max_size={max_size}" + f"[RayActor] Splitting file: source='{source}', destination='{destination}', " + f"task_id='{task_id}', max_size={max_size}, target_parts={target_parts}" ) if file_data is None: @@ -336,6 +365,7 @@ def split_file( file_data=file_data, filename=source, max_size=max_size, + target_parts=target_parts, **params ) split_elapsed = time.perf_counter() - split_start diff --git a/backend/data_process/tasks.py b/backend/data_process/tasks.py index 82d0ed7072..31e1b6bdbc 100644 --- a/backend/data_process/tasks.py +++ b/backend/data_process/tasks.py @@ -20,6 +20,7 @@ from celery.result import allow_join_result from utils.file_management_utils import get_file_size +from utils.knowledge_telemetry import knowledge_span, set_span_attributes, trace_knowledge_operation from database.attachment_db import get_file_stream from database.knowledge_db import get_knowledge_record from services.redis_service import get_redis_service @@ -32,13 +33,14 @@ FORWARD_REDIS_RETRY_MAX, DP_REDIS_CHUNKS_WAIT_TIMEOUT_S, DP_REDIS_CHUNKS_POLL_INTERVAL_MS, + DP_FILE_SPLIT_SIZE_MB, + DP_PART_PROCESSOR_COUNT, RAY_ACTOR_NUM_CPUS, RAY_NUM_CPUS, DISABLE_RAY_DASHBOARD, ROOT_DIR, PER_WAVE_TIMEOUT, MAX_TIMEOUT, - RAY_GLOBAL_ACTOR_POOL_SIZE, RAY_ACTOR_WARM_TIMEOUT_S, RAY_GLOBAL_ACTOR_POOL_NAME, RAY_GLOBAL_ACTOR_POOL_NAMESPACE @@ -52,6 +54,16 @@ IMAGE_METADATA_PROCESS_SOURCE = "UniversalImageExtractor" +@trace_knowledge_operation("knowledge.minio.fetch", "minio.fetch") +def _fetch_minio_source(source: str) -> bytes: + file_stream = get_file_stream(source) + if file_stream is None: + raise FileNotFoundError(f"Unable to fetch file from URL: {source}") + data = file_stream.read() + set_span_attributes(file_size_bytes=len(data), stage="minio.fetch") + return data + + def _wait_for_split_ready(redis_key: str, timeout_s: int, poll_interval_ms: int) -> int: """ Wait until async split aggregation is marked ready in Redis. @@ -90,7 +102,8 @@ def _estimate_parallel_parts() -> int: except Exception: total_cpus = os.cpu_count() or 1 actor_cpus = max(1, int(RAY_ACTOR_NUM_CPUS)) - return max(1, total_cpus // actor_cpus) + cpu_capacity = max(1, total_cpus // actor_cpus) + return min(DP_PART_PROCESSOR_COUNT, cpu_capacity) def _compute_split_wait_timeout(parts_count: int) -> int: @@ -453,6 +466,7 @@ def _build_forward_cancelled_result(ctx: _ForwardContext) -> Dict[str, Any]: } +@trace_knowledge_operation("knowledge.forward.redis_read", "forward.redis_read") def _load_forward_chunks( self: Task, *, @@ -618,6 +632,7 @@ def _extract_error_code_from_es_response( return None +@trace_knowledge_operation("knowledge.forward.elasticsearch", "forward.elasticsearch") def _send_chunks_to_es( chunks: List[Dict[str, Any]], index_name: str, @@ -740,6 +755,19 @@ def ensure_pool(self, desired: int, max_allowed: int) -> int: desired = max(0, int(desired)) max_allowed = max(1, int(max_allowed)) desired = min(desired, max_allowed) + while len(self.actors) > desired: + actor = self.actors.pop() + try: + ray.kill(actor, no_restart=True) + except Exception: + logger.warning( + "[GlobalRayActorPoolManager] Failed to stop excess actor", + exc_info=True, + ) + if self.actors: + self.rr_index %= len(self.actors) + else: + self.rr_index = 0 missing = max(0, desired - len(self.actors)) for _ in range(missing): actor = self._create_and_warm_actor() @@ -796,7 +824,7 @@ def prewarm_ray_actors(target_size: Optional[int] = None) -> int: """ Ensure a global shared pool of warm Ray actors exists for low-latency task execution. """ - desired = RAY_GLOBAL_ACTOR_POOL_SIZE if target_size is None else max( + desired = DP_PART_PROCESSOR_COUNT if target_size is None else max( 0, int(target_size)) manager = _get_or_create_global_pool_manager() current_after = ray.get( @@ -851,6 +879,7 @@ def on_retry(self, exc, task_id, args, kwargs, einfo): @app.task(bind=True, base=LoggingTask, name='data_process.tasks.process_part', queue='process_part_q') +@trace_knowledge_operation("knowledge.process.part", "process.ray_part") def process_part( self, part_bytes: bytes, @@ -924,6 +953,7 @@ def aggregate_parts( @app.task(bind=True, base=LoggingTask, name='data_process.tasks.aggregate_store_chunks', queue='process_part_q') +@trace_knowledge_operation("knowledge.process.redis_aggregate", "process.redis") def aggregate_store_chunks( self, parts_results: List[Dict[str, Any]], @@ -996,6 +1026,7 @@ def aggregate_store_chunks( @app.task(bind=True, base=LoggingTask, name='data_process.tasks.forward_part', queue='forward_q') +@trace_knowledge_operation("knowledge.forward.batch", "forward.batch") def forward_part( self, chunks: List[Dict[str, Any]], @@ -1081,6 +1112,7 @@ def forward_part( @app.task(bind=True, base=LoggingTask, name='data_process.tasks.aggregate_forward_parts', queue='forward_q') +@trace_knowledge_operation("knowledge.forward.aggregate", "forward.aggregate") def aggregate_forward_parts( self, parts_results: List[Dict[str, Any]], @@ -1115,15 +1147,24 @@ def _split_file_for_processing( source_type: str, task_id: str, params: Dict[str, Any], + file_size_bytes: int, file_data: Optional[bytes] = None, ) -> List[bytes]: - max_size = 5 * 1024 * 1024 params.pop("max_size", None) + params.pop("target_parts", None) logger.info( - f"[{request_id}] PROCESS TASK: Splitting file before processing (max_size={max_size})") + f"[{request_id}] PROCESS TASK: Splitting file before processing " + f"(file_size={file_size_bytes}, target_parts={DP_PART_PROCESSOR_COUNT})") split_actor_get_start = time.perf_counter() - split_actor = _get_split_actor() + with knowledge_span( + "knowledge.process.split_actor_acquire", + "process.split_actor_acquire", + task_id=task_id, + source_type=source_type, + processor_count=DP_PART_PROCESSOR_COUNT, + ): + split_actor = _get_split_actor() split_actor_get_elapsed = time.perf_counter() - split_actor_get_start logger.info( f"[{request_id}] PROCESS TASK: split actor ready in {split_actor_get_elapsed:.3f}s") @@ -1133,14 +1174,24 @@ def _split_file_for_processing( "source": source, "destination": source_type, "task_id": task_id, - "max_size": max_size, + "target_parts": DP_PART_PROCESSOR_COUNT, **params, } if file_data is not None: split_kwargs["file_data"] = file_data - parts_ref = split_actor.split_file.remote(**split_kwargs) - parts = ray.get(parts_ref) + with knowledge_span( + "knowledge.process.file_split_rpc", + "process.file_split_rpc", + task_id=task_id, + source_type=source_type, + file_size_bytes=file_size_bytes, + processor_count=DP_PART_PROCESSOR_COUNT, + ) as span: + parts_ref = split_actor.split_file.remote(**split_kwargs) + parts = ray.get(parts_ref) + if span is not None: + span.set_attribute("file.parts_count", len(parts or [])) split_call_elapsed = time.perf_counter() - split_call_start logger.info( f"[{request_id}] PROCESS TASK: split_file RPC done in {split_call_elapsed:.3f}s " @@ -1228,20 +1279,38 @@ def _run_processing_for_parts( ).set(queue='process_part_q') logger.info( f"[{request_id}] PROCESS TASK: Dispatching {len(parts)} part tasks...") - chord(group_tasks)(callback) + with knowledge_span( + "knowledge.process.part_dispatch", + "process.part_dispatch", + task_id=task_id, + part_count=len(parts), + parallel_parts=_estimate_parallel_parts(), + queue_name="process_part_q", + ): + chord(group_tasks)(callback) split_wait_timeout = _compute_split_wait_timeout(len(parts)) logger.info( f"[{request_id}] PROCESS TASK: Waiting split aggregation, timeout={split_wait_timeout}s, " f"parts={len(parts)}, est_parallel={_estimate_parallel_parts()}") - split_chunk_count = _wait_for_split_ready( - redis_key=redis_key, - timeout_s=split_wait_timeout, + with knowledge_span( + "knowledge.process.part_wait", + "process.part_wait", + task_id=task_id, + part_count=len(parts), + parallel_parts=_estimate_parallel_parts(), + timeout_seconds=split_wait_timeout, poll_interval_ms=DP_REDIS_CHUNKS_POLL_INTERVAL_MS, - ) + ): + split_chunk_count = _wait_for_split_ready( + redis_key=redis_key, + timeout_s=split_wait_timeout, + poll_interval_ms=DP_REDIS_CHUNKS_POLL_INTERVAL_MS, + ) return True, None, split_chunk_count +@trace_knowledge_operation("knowledge.process.split_and_ray", "process.split_and_ray") def _process_source_with_split( request_id: str, source: str, @@ -1255,12 +1324,53 @@ def _process_source_with_split( params: Dict[str, Any], file_data: Optional[bytes] = None, ) -> Tuple[bool, Optional[List[Dict[str, Any]]], Optional[int]]: + file_size_bytes = ( + len(file_data) + if file_data is not None + else os.path.getsize(source) + ) + split_threshold_bytes = DP_FILE_SPLIT_SIZE_MB * 1024 * 1024 + if file_size_bytes <= split_threshold_bytes: + logger.info( + f"[{request_id}] PROCESS TASK: File size {file_size_bytes} does not exceed " + f"split threshold {split_threshold_bytes}; processing without FileSplitter") + process_actor = get_ray_actor() + if file_data is not None: + chunks_ref = process_actor.process_bytes.remote( + file_data, + original_filename or os.path.basename(source), + chunking_strategy, + task_id=task_id, + model_id=embedding_model_id, + tenant_id=tenant_id, + **params, + ) + else: + chunks_ref = process_actor.process_file.remote( + source, + chunking_strategy, + destination=source_type, + task_id=task_id, + model_id=embedding_model_id, + tenant_id=tenant_id, + **params, + ) + chunks = ray.get(chunks_ref) + if chunks: + redis_key = f"dp:{task_id}:chunks" + stored = ray.get( + process_actor.store_chunks_in_redis.remote(redis_key, chunks)) + if not stored: + raise RuntimeError("Failed to persist processed chunks in Redis") + return False, chunks, None + parts = _split_file_for_processing( request_id=request_id, source=source, source_type=source_type, task_id=task_id, params=params, + file_size_bytes=file_size_bytes, file_data=file_data, ) filename_for_processing = original_filename or os.path.basename(source) @@ -1289,7 +1399,10 @@ def _process_source_with_split( if not split_async: redis_key = f"dp:{task_id}:chunks" process_actor = get_ray_actor() - process_actor.store_chunks_in_redis.remote(redis_key, chunks) + store_ref = process_actor.store_chunks_in_redis.remote(redis_key, chunks) + stored = ray.get(store_ref) + if not stored: + raise RuntimeError("Failed to persist processed chunks in Redis") logger.info( f"[{request_id}] PROCESS TASK: Stored chunks in Redis at key '{redis_key}'") @@ -1318,6 +1431,7 @@ def _build_no_valid_chunks_error( @app.task(bind=True, base=LoggingTask, name='data_process.tasks.process', queue='process_q') +@trace_knowledge_operation("knowledge.process", "process") def process( self, source: str, @@ -1376,7 +1490,7 @@ def process( raise FileNotFoundError(f"File does not exist: {source}") file_size = os.path.getsize(source) - file_size_mb = file_size / (5 * 1024 * 1024) + file_size_mb = file_size / (1024 * 1024) logger.info( f"[{self.request.id}] PROCESS TASK: File size: {file_size_mb:.2f}MB") @@ -1405,11 +1519,7 @@ def process( # Measure MinIO fetch time in process worker logs for observability fetch_start = time.perf_counter() - file_stream = get_file_stream(source) - if file_stream is None: - raise FileNotFoundError( - f"Unable to fetch file from URL: {source}") - file_data = file_stream.read() + file_data = _fetch_minio_source(source) fetch_elapsed = time.perf_counter() - fetch_start logger.info( f"[{self.request.id}] PROCESS TASK: MinIO fetch done in {fetch_elapsed:.3f}s, " @@ -1476,6 +1586,7 @@ def process( logger.info( f"[{self.request.id}] PROCESS TASK: Chunk composition: total={chunk_count}, " f"image_metadata={image_metadata_chunk_count}, text={max(0, chunk_count - image_metadata_chunk_count)}") + set_span_attributes(chunk_count=chunk_count, stage="process.complete") # Update task state to SUCCESS after Ray processing completes # This transitions from STARTED (PROCESSING) to SUCCESS (WAIT_FOR_FORWARDING) @@ -1627,6 +1738,7 @@ def process( @app.task(bind=True, base=LoggingTask, name='data_process.tasks.forward', queue='forward_q') +@trace_knowledge_operation("knowledge.forward", "forward") def forward( self, processed_data: Dict, @@ -1634,7 +1746,8 @@ def forward( source: str, source_type: str = 'minio', original_filename: Optional[str] = None, - authorization: Optional[str] = None + authorization: Optional[str] = None, + telemetry_context: Optional[Dict[str, str]] = None, ) -> Dict: """ Vectorize and store processed chunks in Elasticsearch @@ -1686,6 +1799,7 @@ def forward( # Calculate total chunks for progress tracking total_chunks = len(chunks) if chunks else 0 + set_span_attributes(chunk_count=total_chunks, stage="forward.format") formatted_chunks = [] # Compute once per file to avoid repeated IO/MinIO calls inside loop file_size = get_file_size(source_type, original_source) if isinstance( @@ -1965,10 +2079,12 @@ def forward( name="data_process.tasks.cleanup_source", queue="forward_q", ) +@trace_knowledge_operation("knowledge.cleanup", "cleanup") def cleanup_source( self, forward_result: Dict[str, Any], authorization: Optional[str] = None, + telemetry_context: Optional[Dict[str, str]] = None, ) -> Dict[str, Any]: """ Conditionally delete the MinIO source file after successful indexing. @@ -2065,6 +2181,7 @@ def cleanup_source( return forward_result +@trace_knowledge_operation("knowledge.chain.submit", "chain.submit") def submit_process_forward_chain( *, source: str, @@ -2075,6 +2192,7 @@ def submit_process_forward_chain( authorization: Optional[str] = None, embedding_model_id: Optional[int] = None, tenant_id: Optional[str] = None, + telemetry_context: Optional[Dict[str, str]] = None, ) -> str: """ Build and enqueue a Celery chain: process -> forward. @@ -2090,16 +2208,21 @@ def submit_process_forward_chain( index_name=index_name, original_filename=original_filename, embedding_model_id=embedding_model_id, - tenant_id=tenant_id + tenant_id=tenant_id, + telemetry_context=telemetry_context or {}, ).set(queue='process_q'), forward.s( index_name=index_name, source=source, source_type=source_type, original_filename=original_filename, - authorization=authorization + authorization=authorization, + telemetry_context=telemetry_context or {}, + ).set(queue='forward_q'), + cleanup_source.s( + authorization=authorization, + telemetry_context=telemetry_context or {}, ).set(queue='forward_q'), - cleanup_source.s(authorization=authorization).set(queue='forward_q'), ) result = task_chain.apply_async() @@ -2120,7 +2243,8 @@ def process_and_forward( original_filename: Optional[str] = None, authorization: Optional[str] = None, embedding_model_id: Optional[int] = None, - tenant_id: Optional[str] = None + tenant_id: Optional[str] = None, + telemetry_context: Optional[Dict[str, str]] = None, ) -> str: """ Combined task that chains processing and forwarding @@ -2152,6 +2276,7 @@ def process_and_forward( authorization=authorization, embedding_model_id=embedding_model_id, tenant_id=tenant_id, + telemetry_context=telemetry_context or {}, ) if chain_id: logger.info(f"Created task chain ID: {chain_id}") diff --git a/backend/data_process/worker.py b/backend/data_process/worker.py index 48323869bf..775e7bd89f 100644 --- a/backend/data_process/worker.py +++ b/backend/data_process/worker.py @@ -45,7 +45,7 @@ REDIS_URL, WORKER_CONCURRENCY, WORKER_NAME, - RAY_GLOBAL_ACTOR_POOL_SIZE, + DP_PART_PROCESSOR_COUNT, ) from .app import app @@ -155,6 +155,21 @@ def setup_worker_process_resources(**kwargs): logger.info(f"⚙️ Initialize worker process {process_id}") try: + # Celery prefork children need their own OTLP provider/exporter. Importing + # monitoring here avoids inheriting a dead BatchSpanProcessor thread. + try: + from utils.monitoring import monitoring_manager + + logger.info( + "Knowledge telemetry initialized in worker process: enabled=%s", + monitoring_manager.is_enabled, + ) + except Exception: + logger.warning( + "Knowledge telemetry initialization failed; worker will continue", + exc_info=True, + ) + # Initialize process-specific resources # e.g. database connection pool, cache client, etc. @@ -211,7 +226,7 @@ def worker_ready_handler(**kwargs): # Prewarm a cluster-global shared actor pool once at startup. # Multiple workers may trigger this, but pool manager is idempotent. - target = RAY_GLOBAL_ACTOR_POOL_SIZE + target = DP_PART_PROCESSOR_COUNT def _prewarm_in_background(): try: @@ -345,6 +360,21 @@ def validate_redis_connection() -> bool: def start_worker(): """Start Celery worker with appropriate settings""" + # The current worker uses a thread pool, so worker_process_init is not + # guaranteed to fire. Initialize the exporter in the worker main process. + try: + from utils.monitoring import monitoring_manager + + logger.info( + "Knowledge telemetry initialized before worker start: enabled=%s", + monitoring_manager.is_enabled, + ) + except Exception: + logger.warning( + "Knowledge telemetry initialization failed; worker will continue", + exc_info=True, + ) + # Read from runtime env first, so launcher-assigned values always win. queues = QUEUES worker_name = WORKER_NAME diff --git a/backend/data_process_service.py b/backend/data_process_service.py index c162cfaaa8..1d955ebf80 100644 --- a/backend/data_process_service.py +++ b/backend/data_process_service.py @@ -19,7 +19,8 @@ from consts.const import ( REDIS_URL, REDIS_PORT, FLOWER_PORT, RAY_DASHBOARD_PORT, RAY_DASHBOARD_HOST, RAY_ACTOR_NUM_CPUS, RAY_NUM_CPUS, DISABLE_RAY_DASHBOARD, DISABLE_CELERY_FLOWER, - DOCKER_ENVIRONMENT, RAY_OBJECT_STORE_MEMORY_GB, RAY_preallocate_plasma, RAY_TEMP_DIR + DOCKER_ENVIRONMENT, RAY_OBJECT_STORE_MEMORY_GB, RAY_preallocate_plasma, RAY_TEMP_DIR, + DP_PART_PROCESSOR_COUNT, ) # Load environment variables @@ -198,7 +199,10 @@ def start_workers(self): # Calculate concurrency for the process-worker. Each worker will spawn an actor, # so we limit concurrency to avoid oversubscribing Ray's CPU resources. - process_worker_concurrency = max(1, total_cpus // ray_actor_num_cpus) + process_worker_concurrency = min( + DP_PART_PROCESSOR_COUNT, + max(1, total_cpus // ray_actor_num_cpus), + ) # For forward-worker, it's I/O bound. A higher concurrency is fine, but we can cap it # relative to CPU count to avoid creating excessive threads on small machines. diff --git a/backend/database/attachment_db.py b/backend/database/attachment_db.py index 1d9e403b73..5edf07dc79 100644 --- a/backend/database/attachment_db.py +++ b/backend/database/attachment_db.py @@ -9,6 +9,8 @@ from consts.const import NORTHBOUND_EXTERNAL_URL from urllib.parse import quote +from utils.knowledge_telemetry import set_span_attributes, trace_knowledge_operation + def _normalize_object_and_bucket(object_name: str, bucket: Optional[str] = None) -> Tuple[str, Optional[str]]: """ @@ -147,6 +149,7 @@ def upload_file( return response +@trace_knowledge_operation("knowledge.minio.upload", "minio.upload") def upload_fileobj( file_obj: BinaryIO, file_name: str, @@ -188,6 +191,7 @@ def upload_fileobj( # Upload file success, result = minio_client.upload_fileobj( file_obj, object_name, bucket) + set_span_attributes(file_size_bytes=file_size, stage="minio.upload") # Restore original position (if file is still open) try: diff --git a/backend/services/data_process_service.py b/backend/services/data_process_service.py index a7529127c2..4a9eab3c5c 100644 --- a/backend/services/data_process_service.py +++ b/backend/services/data_process_service.py @@ -541,6 +541,7 @@ async def create_batch_tasks_impl(self, authorization: Optional[str], request: B original_filename = source_config.get('original_filename') embedding_model_id = source_config.get('embedding_model_id') tenant_id = source_config.get('tenant_id') + telemetry_context = source_config.get('telemetry_context') or {} # Validate required fields if not source: @@ -561,6 +562,7 @@ async def create_batch_tasks_impl(self, authorization: Optional[str], request: B authorization=authorization, embedding_model_id=embedding_model_id, tenant_id=tenant_id, + telemetry_context=telemetry_context, ) if not chain_id: logger.error( diff --git a/backend/services/file_management_service.py b/backend/services/file_management_service.py index f56d76d914..236ad02afd 100644 --- a/backend/services/file_management_service.py +++ b/backend/services/file_management_service.py @@ -42,6 +42,7 @@ from services.vectordatabase_service import ElasticSearchService, get_vector_db_core from utils.config_utils import tenant_config_manager, get_model_name_from_config from utils.file_management_utils import save_upload_file +from utils.knowledge_telemetry import trace_knowledge_operation from nexent import MessageObserver from nexent.multi_modal.utils import parse_s3_url @@ -470,6 +471,7 @@ def make_unique_names(original_names: List[str], taken_lower: set) -> List[str]: errors, uploaded_file_paths, uploaded_filenames, quota_status) +@trace_knowledge_operation("knowledge.upload.batch", "upload") async def upload_to_minio( files: List[UploadFile], folder: str, diff --git a/backend/utils/file_management_utils.py b/backend/utils/file_management_utils.py index 83c3957e71..29fd621ff0 100644 --- a/backend/utils/file_management_utils.py +++ b/backend/utils/file_management_utils.py @@ -16,6 +16,7 @@ from consts.model import ProcessParams from database.attachment_db import get_file_size_from_minio from utils.auth_utils import get_current_user_id +from utils.knowledge_telemetry import inject_trace_context, set_span_attributes, trace_knowledge_operation logger = logging.getLogger("file_management_utils") @@ -39,6 +40,7 @@ async def save_upload_file(file: UploadFile, upload_path: Path) -> bool: return False +@trace_knowledge_operation("knowledge.process.submit", "process.submit") async def trigger_data_process(files: List[dict], process_params: ProcessParams): """Trigger data processing service to handle uploaded files""" try: @@ -57,6 +59,8 @@ async def trigger_data_process(files: List[dict], process_params: ProcessParams) headers = { "Authorization": f"Bearer {process_params.authorization}" } + telemetry_context = inject_trace_context() + headers.update(telemetry_context) # Build source data list if len(files) == 1: @@ -71,6 +75,7 @@ async def trigger_data_process(files: List[dict], process_params: ProcessParams) "embedding_model_id": embedding_model_id, "tenant_id": tenant_id } + payload["telemetry_context"] = telemetry_context try: async with httpx.AsyncClient() as client: @@ -102,6 +107,7 @@ async def trigger_data_process(files: List[dict], process_params: ProcessParams) "embedding_model_id": embedding_model_id, "tenant_id": tenant_id } + source["telemetry_context"] = telemetry_context sources.append(source) payload = {"sources": sources} @@ -111,6 +117,7 @@ async def trigger_data_process(files: List[dict], process_params: ProcessParams) response = await client.post(f"{DATA_PROCESS_SERVICE}/tasks/batch", headers=headers, json=payload, timeout=30.0) if response.status_code == 201: + set_span_attributes(stage="process.submitted") return response.json() else: logger.error( @@ -412,4 +419,3 @@ def _run_libreoffice_conversion(): raise FileNotFoundError( "LibreOffice is not installed or not available in PATH. " ) from e - diff --git a/backend/utils/knowledge_telemetry.py b/backend/utils/knowledge_telemetry.py new file mode 100644 index 0000000000..87cf31409e --- /dev/null +++ b/backend/utils/knowledge_telemetry.py @@ -0,0 +1,318 @@ +"""Best-effort telemetry helpers for the knowledge-base ingestion pipeline.""" + +from __future__ import annotations + +import functools +import hashlib +import inspect +import logging +import os +import time +from contextlib import contextmanager +from typing import Any, Dict, Iterator, Mapping, Optional + +logger = logging.getLogger(__name__) +BYTES_PER_MB = 1024 * 1024 +KNOWLEDGE_SPAN_DESCRIPTIONS = { + "knowledge.upload.batch": "`读取一批上传文件并组织 MinIO 上传。", + "knowledge.minio.upload": "`执行实际的 MinIO 对象写入。", + "knowledge.process.submit": "`从 backend 调用 data-process HTTP 接口,提交单文件或批量处理请求。", + "knowledge.chain.submit": "`创建并投递 Celery process → forward → cleanup 任务链。它只表示任务成功进入队列,不代表文件已经处理完成。", + "knowledge.process": "单个文件的顶层解析任务,负责获取文件、拆分、调度处理并统计 chunk。", + "knowledge.minio.fetch": "从 MinIO 下载待处理文件字节。", + "knowledge.process.split_and_ray": "判断是否拆分文件,并向 Ray 或 Celery 子任务分发解析工作。", + "knowledge.process.split_actor_acquire": "请求已预热的 Ray actor 执行文件拆分任务。", + "knowledge.process.file_split_rpc": "在选定的 Ray actor 中执行文件拆分任务。", + "knowledge.process.part_dispatch": "分发给 Celery 到并等待聚合。", + "knowledge.process.part_wait": "等待 Celery part tasks 和 Redis 聚合完成。", + "knowledge.process.part": "处理一个拆分后的文档分片。", + "knowledge.process.ray_actor": "在独立 Ray actor 进程内执行文档预处理。", + "knowledge.preprocess.split": "使用 FileSplitter 拆分大文件。", + "knowledge.preprocess.image_extract": "为多模态知识库提取图片及图片元数据。", + "knowledge.preprocess.typed": "根据文件类型调用不同解析器。", + "knowledge.process.redis_aggregate": "合并各分片的 chunk,并将结果写入 Redis,供 forward 阶段读取。", + "knowledge.forward": "顶层索引任务,读取、过滤、格式化 chunk,并决定同步或分批提交。", + "knowledge.forward.redis_read": "`从 Redis 读取 process 阶段生成的 chunk。", + "knowledge.forward.batch": "`处理并提交一个 chunk batch。", + "knowledge.forward.elasticsearch": "调用 Elasticsearch 接口写入向量。", + "knowledge.forward.aggregate": "`汇总多个 batch 的结果,验证提交数和索引数。", + "knowledge.cleanup": "索引成功后按清理策略删除上传源文件;取消、失败或策略不允许时会跳过删除。", +} + + +def _span_kind(name: str) -> str: + """Return the Phoenix/OpenInference kind appropriate for the operation.""" + tool_markers = (".minio.", ".redis_", ".redis_aggregate", ".elasticsearch") + return "TOOL" if any(marker in name for marker in tool_markers) else "CHAIN" + + +def _is_celery_retry(error: BaseException) -> bool: + """Recognize Celery's control-flow retry without requiring Celery in backend.""" + error_type = type(error) + return error_type.__name__ == "Retry" and error_type.__module__.startswith("celery") + +try: + from opentelemetry import context as otel_context + from opentelemetry import metrics, propagate, trace + from opentelemetry.trace import Status, StatusCode + + OTEL_AVAILABLE = True +except ImportError: # pragma: no cover - optional dependency + OTEL_AVAILABLE = False + + +def _safe_hash(value: Any) -> str: + if not value: + return "" + return hashlib.sha256(str(value).encode("utf-8", errors="ignore")).hexdigest()[:16] + + +def _extension(filename: Any) -> str: + return os.path.splitext(str(filename or ""))[1].lower()[:16] + + +def _safe_attributes(values: Mapping[str, Any]) -> Dict[str, Any]: + """Return low-cardinality, non-content attributes accepted by OTel.""" + attrs: Dict[str, Any] = {} + aliases = { + "task_id": "task.id", + "chain_id": "chain.id", + "index_name": "knowledge_base.id", + "tenant_id": "tenant.id_hash", + "source_type": "source.type", + "stage": "ingestion.stage", + "file_size": "file.size_mb", + "file_size_bytes": "file.size_mb", + "chunk_count": "chunk.count", + "chunks_count": "chunk.count", + "batch_index": "batch.index", + "total_batches": "batch.total", + "processor": "processor.name", + "part_count": "file.parts_count", + "processor_count": "processor.count", + "parallel_parts": "processor.parallel_count", + "timeout_seconds": "timeout.seconds", + "poll_interval_ms": "poll.interval_ms", + "queue_name": "messaging.destination.name", + } + for key, target in aliases.items(): + value = values.get(key) + if value is None: + continue + if key in {"tenant_id", "index_name"}: + value = _safe_hash(value) + if key in {"file_size", "file_size_bytes"}: + try: + value = round(float(value) / BYTES_PER_MB, 3) + except (TypeError, ValueError): + continue + if isinstance(value, (str, bool, int, float)): + attrs[target] = value + filename = values.get("original_filename") or values.get("filename") or values.get("file_name") + if filename: + attrs["file.extension"] = _extension(filename) + return attrs + + +def inject_trace_context() -> Dict[str, str]: + """Create a safe W3C propagation carrier for HTTP/Celery boundaries.""" + carrier: Dict[str, str] = {} + if OTEL_AVAILABLE: + try: + propagate.inject(carrier) + except Exception: + logger.debug("Unable to inject telemetry context", exc_info=True) + return carrier + + +def _resource_snapshot() -> Dict[str, Any]: + """Collect process and host resources without making psutil mandatory.""" + try: + import psutil + + process = psutil.Process() + memory = psutil.virtual_memory() + children = process.children(recursive=True) + process_rss = process.memory_info().rss + process_tree_rss = process_rss + sum( + child.memory_info().rss for child in children if child.is_running() + ) + snapshot = { + "process.rss_memory_mb": round(process_rss / BYTES_PER_MB, 3), + "process.cpu_percent": round(process.cpu_percent(interval=None), 3), + "process.thread_count": process.num_threads(), + "process_tree.rss_memory_mb": round(process_tree_rss / BYTES_PER_MB, 3), + "process_tree.child_count": len(children), + "host.used_memory_percent": round(memory.percent, 3), + "host.available_memory_mb": round(memory.available / BYTES_PER_MB, 3), + "host.cpu_percent": round(psutil.cpu_percent(interval=None), 3), + } + for path, key in ( + ("/sys/fs/cgroup/memory.current", "container.used_memory_mb"), + ("/sys/fs/cgroup/memory.max", "container.memory_limit_mb"), + ): + try: + with open(path, encoding="ascii") as cgroup_file: + value = cgroup_file.read().strip() + if value != "max": + snapshot[key] = round(int(value) / BYTES_PER_MB, 3) + except (OSError, ValueError): + pass + return snapshot + except Exception: + return {} + + +def _record_metrics(stage: str, duration_ms: float, snapshot: Mapping[str, Any]) -> None: + if not OTEL_AVAILABLE: + return + try: + meter = metrics.get_meter("nexent.knowledge_ingestion") + labels = {"ingestion.stage": stage} + meter.create_histogram("nexent.ingestion.stage.duration", unit="ms").record(duration_ms, labels) + if "process.rss_memory_mb" in snapshot: + meter.create_histogram("nexent.ingestion.process.rss", unit="By").record( + snapshot["process.rss_memory_mb"] * BYTES_PER_MB, labels + ) + if "process.cpu_percent" in snapshot: + meter.create_histogram("nexent.ingestion.process.cpu", unit="%").record( + snapshot["process.cpu_percent"], labels + ) + except Exception: + logger.debug("Unable to record ingestion resource metrics", exc_info=True) + + +@contextmanager +def knowledge_span(name: str, stage: str, **attributes: Any) -> Iterator[Any]: + """Create a non-blocking ingestion span, optionally continuing a remote trace.""" + if not OTEL_AVAILABLE: + yield None + return + + token = None + span_cm = None + span = None + started = time.perf_counter() + start_resources = _resource_snapshot() + try: + carrier = attributes.pop("telemetry_context", None) + if isinstance(carrier, Mapping): + remote_context = propagate.extract(dict(carrier)) + token = otel_context.attach(remote_context) + tracer = trace.get_tracer("nexent.knowledge_ingestion") + span_cm = tracer.start_as_current_span(name) + span = span_cm.__enter__() + span.set_attributes({ + "ingestion.stage": stage, + "ingestion.operation.description": KNOWLEDGE_SPAN_DESCRIPTIONS.get(name, stage), + "openinference.span.kind": _span_kind(name), + **_safe_attributes(attributes), + **{f"resource.start.{key}": value for key, value in start_resources.items()}, + }) + except Exception: + logger.debug("Telemetry span setup failed for %s", name, exc_info=True) + if span_cm is not None: + try: + span_cm.__exit__(None, None, None) + except Exception: + pass + if token is not None: + try: + otel_context.detach(token) + except Exception: + pass + yield None + return + + try: + yield span + except Exception as exc: + try: + if _is_celery_retry(exc): + span.set_attribute("ingestion.status", "retry") + span.set_attribute("retry.attempt", int(attributes.get("retry_attempt", 0))) + retry_delay = attributes.get("retry_delay_seconds", getattr(exc, "when", 0)) + if isinstance(retry_delay, (int, float)): + span.set_attribute("retry.delay_seconds", float(retry_delay)) + # Retry is expected control flow, not a failed ingestion operation. + span.set_status(Status(StatusCode.OK)) + else: + span.record_exception(exc) + span.set_status(Status(StatusCode.ERROR, type(exc).__name__)) + span.set_attribute("error.type", type(exc).__name__) + except Exception: + logger.debug("Unable to record ingestion failure", exc_info=True) + raise + else: + try: + span.set_status(Status(StatusCode.OK)) + span.set_attribute("ingestion.status", "success") + except Exception: + logger.debug("Unable to record ingestion success", exc_info=True) + finally: + end_resources = _resource_snapshot() + duration_ms = (time.perf_counter() - started) * 1000.0 + if span is not None: + try: + span.set_attribute("ingestion.duration_ms", duration_ms) + span.set_attributes({f"resource.end.{key}": value for key, value in end_resources.items()}) + except Exception: + logger.debug("Unable to attach ingestion resource snapshot", exc_info=True) + _record_metrics(stage, duration_ms, end_resources) + if span_cm is not None: + try: + span_cm.__exit__(None, None, None) + except Exception: + logger.debug("Telemetry span close failed", exc_info=True) + if token is not None: + try: + otel_context.detach(token) + except Exception: + logger.debug("Telemetry context detach failed", exc_info=True) + + +def trace_knowledge_operation(name: str, stage: str): + """Decorate sync or async ingestion operations without changing behavior.""" + def decorator(func): + signature = inspect.signature(func) + + def span_args(args, kwargs): + try: + bound = signature.bind_partial(*args, **kwargs) + values = dict(bound.arguments) + except Exception: + values = dict(kwargs) + request = values.get("self") + if request is not None and hasattr(request, "request"): + values.setdefault("task_id", getattr(request.request, "id", None)) + values.setdefault("retry_attempt", getattr(request.request, "retries", 0) + 1) + values["telemetry_context"] = values.get("telemetry_context") or values.get("params", {}).get( + "telemetry_context" + ) + return values + + if inspect.iscoroutinefunction(func): + @functools.wraps(func) + async def async_wrapper(*args, **kwargs): + with knowledge_span(name, stage, **span_args(args, kwargs)): + return await func(*args, **kwargs) + return async_wrapper + + @functools.wraps(func) + def wrapper(*args, **kwargs): + with knowledge_span(name, stage, **span_args(args, kwargs)): + return func(*args, **kwargs) + return wrapper + return decorator + + +def set_span_attributes(**attributes: Any) -> None: + """Attach safe attributes to the active span, if any.""" + if not OTEL_AVAILABLE: + return + try: + span = trace.get_current_span() + if span and span.is_recording(): + span.set_attributes(_safe_attributes(attributes)) + except Exception: + logger.debug("Unable to set ingestion span attributes", exc_info=True) diff --git a/deploy/docker/assets/monitoring/otel-collector-phoenix-config.yml b/deploy/docker/assets/monitoring/otel-collector-phoenix-config.yml index 0682a6e4dc..489f08ba8c 100644 --- a/deploy/docker/assets/monitoring/otel-collector-phoenix-config.yml +++ b/deploy/docker/assets/monitoring/otel-collector-phoenix-config.yml @@ -19,11 +19,34 @@ processors: attributes: - key: service.name value: nexent-backend - action: upsert + action: insert - key: service.version from_attribute: version action: insert + # Phoenix selects projects from this resource attribute. Applying it only + # to knowledge.* spans keeps generic HTTP/Elasticsearch polling elsewhere. + transform/knowledge_project: + error_mode: ignore + trace_statements: + - context: span + statements: + - 'set(resource.attributes["openinference.project.name"], "knowledge-base-monitor") where IsMatch(span.name, "^knowledge\\.")' + + # Filter conditions describe what to drop. Two pipelines prevent resource + # attributes shared by a mixed OTLP batch from moving unrelated spans. + filter/knowledge_only: + error_mode: ignore + traces: + span: + - 'not IsMatch(span.name, "^knowledge\\.")' + + filter/non_knowledge: + error_mode: ignore + traces: + span: + - 'IsMatch(span.name, "^knowledge\\.")' + exporters: debug: verbosity: normal @@ -49,13 +72,34 @@ exporters: max_interval: 30s # 最大重试间隔不超过 30s max_elapsed_time: 300s # 一条数据最多重试 5 分钟,超过则彻底放弃并丢弃 + # Phoenix 15.5+ gives this header precedence over resource attributes. + otlphttp/phoenix_knowledge: + endpoint: http://phoenix:6006 + headers: + x-project-name: knowledge-base-monitor + timeout: 5s + sending_queue: + enabled: true + num_consumers: 4 + queue_size: 2000 + retry_on_failure: + enabled: true + initial_interval: 1s + max_interval: 30s + max_elapsed_time: 300s + service: pipelines: traces: receivers: [otlp] - processors: [memory_limiter, resource, batch] + processors: [memory_limiter, resource, filter/non_knowledge, batch] exporters: [otlphttp/phoenix, debug] + traces/knowledge: + receivers: [otlp] + processors: [memory_limiter, resource, filter/knowledge_only, transform/knowledge_project, batch] + exporters: [otlphttp/phoenix_knowledge, debug] + metrics: receivers: [otlp] processors: [memory_limiter, resource, batch] diff --git a/deploy/env/.env.example b/deploy/env/.env.example index 902e72d7c8..141022c420 100644 --- a/deploy/env/.env.example +++ b/deploy/env/.env.example @@ -140,7 +140,6 @@ FLOWER_PORT=5555 RAY_ACTOR_NUM_CPUS=2 RAY_DASHBOARD_PORT=8265 RAY_DASHBOARD_HOST=0.0.0.0 -RAY_NUM_CPUS=4 RAY_OBJECT_STORE_MEMORY_GB=0.25 RAY_TEMP_DIR=/tmp/ray RAY_LOG_LEVEL=INFO @@ -158,9 +157,10 @@ CELERY_TASK_TIME_LIMIT=3600 ELASTICSEARCH_REQUEST_TIMEOUT=30 # Worker Configuration -QUEUES=process_q,forward_q +QUEUES=process_q,process_part_q,forward_q WORKER_NAME= -WORKER_CONCURRENCY=4 +DP_PART_PROCESSOR_COUNT=3 +DP_FILE_SPLIT_SIZE_MB=5 # Skills Configuration SKILLS_PATH=/mnt/nexent-data/skills diff --git a/deploy/images/dockerfiles/data-process/Dockerfile b/deploy/images/dockerfiles/data-process/Dockerfile index bafa5334ca..c86534e277 100644 --- a/deploy/images/dockerfiles/data-process/Dockerfile +++ b/deploy/images/dockerfiles/data-process/Dockerfile @@ -139,7 +139,7 @@ RUN --mount=type=cache,id=nexent-data-process-uv-${TARGETARCH},target=/root/.cac fi && \ uv venv .venv && \ uv pip install --python .venv/bin/python --link-mode copy $mirror_index_args $torch_args ".[data-process]" && \ - uv pip install --python .venv/bin/python --link-mode copy $mirror_index_args $torch_args "/opt/sdk[data-process]" && \ + uv pip install --python .venv/bin/python --link-mode copy $mirror_index_args $torch_args "/opt/sdk[data-process,performance]" && \ if [ "$DATA_PROCESS_DEPENDENCY_VARIANT" = "cpu" ]; then \ .venv/bin/python -c 'import importlib.metadata as metadata, importlib.util, sys; blocked = sorted(name for name in ((dist.metadata.get("Name") or "").lower() for dist in metadata.distributions()) if name == "triton" or name.startswith("nvidia-") or name.startswith("cuda-")); blocked and sys.exit("CPU data-process image must not install CUDA packages: " + ", ".join(blocked)); spec = importlib.util.find_spec("torch"); torch = __import__("torch") if spec else None; torch is not None and torch.cuda.is_available() and sys.exit("CPU data-process image unexpectedly reports CUDA availability"); print(f"Using CPU PyTorch {torch.__version__}") if torch else None'; \ fi diff --git a/deploy/k8s/deploy.sh b/deploy/k8s/deploy.sh index 65eeb97d86..d09ead0878 100755 --- a/deploy/k8s/deploy.sh +++ b/deploy/k8s/deploy.sh @@ -481,7 +481,6 @@ render_k8s_runtime_config_values() { printf ' rayDashboardPort: %s\n' "$(yaml_quote "$(env_or_default RAY_DASHBOARD_PORT "8265")")" printf ' rayDashboardHost: %s\n' "$(yaml_quote "$(env_or_default RAY_DASHBOARD_HOST "0.0.0.0")")" printf ' rayActorNumCpus: %s\n' "$(yaml_quote "$(env_or_default RAY_ACTOR_NUM_CPUS "2")")" - printf ' rayNumCpus: %s\n' "$(yaml_quote "$(env_or_default RAY_NUM_CPUS "4")")" printf ' rayObjectStoreMemoryGb: %s\n' "$(yaml_quote "$(env_or_default RAY_OBJECT_STORE_MEMORY_GB "0.25")")" printf ' rayTempDir: %s\n' "$(yaml_quote "$(env_or_default RAY_TEMP_DIR "/tmp/ray")")" printf ' rayLogLevel: %s\n' "$(yaml_quote "$(env_or_default RAY_LOG_LEVEL "INFO")")" @@ -492,9 +491,10 @@ render_k8s_runtime_config_values() { printf ' celeryWorkerPrefetchMultiplier: %s\n' "$(yaml_quote "$(env_or_default CELERY_WORKER_PREFETCH_MULTIPLIER "1")")" printf ' celeryTaskTimeLimit: %s\n' "$(yaml_quote "$(env_or_default CELERY_TASK_TIME_LIMIT "3600")")" printf ' elasticsearchRequestTimeout: %s\n' "$(yaml_quote "$(env_or_default ELASTICSEARCH_REQUEST_TIMEOUT "30")")" - printf ' queues: %s\n' "$(yaml_quote "$(env_or_default QUEUES "process_q,forward_q")")" + printf ' queues: %s\n' "$(yaml_quote "$(env_or_default QUEUES "process_q,process_part_q,forward_q")")" + printf ' partProcessorCount: %s\n' "$(yaml_quote "$(env_or_default DP_PART_PROCESSOR_COUNT "3")")" + printf ' fileSplitSizeMb: %s\n' "$(yaml_quote "$(env_or_default DP_FILE_SPLIT_SIZE_MB "5")")" printf ' workerName: %s\n' "$(yaml_quote "$(env_or_default WORKER_NAME "")")" - printf ' workerConcurrency: %s\n' "$(yaml_quote "$(env_or_default WORKER_CONCURRENCY "4")")" echo " oauth:" printf ' githubClientId: %s\n' "$(yaml_quote "$(env_or_default GITHUB_OAUTH_CLIENT_ID "")")" printf ' githubClientSecret: %s\n' "$(yaml_quote "$(env_or_default GITHUB_OAUTH_CLIENT_SECRET "")")" diff --git a/deploy/k8s/helm/nexent/charts/nexent-common/templates/configmap.yaml b/deploy/k8s/helm/nexent/charts/nexent-common/templates/configmap.yaml index de8fedcc38..df3369612a 100644 --- a/deploy/k8s/helm/nexent/charts/nexent-common/templates/configmap.yaml +++ b/deploy/k8s/helm/nexent/charts/nexent-common/templates/configmap.yaml @@ -103,7 +103,9 @@ data: RAY_DASHBOARD_PORT: {{ .Values.config.dataProcess.rayDashboardPort | quote }} RAY_DASHBOARD_HOST: {{ .Values.config.dataProcess.rayDashboardHost | quote }} RAY_ACTOR_NUM_CPUS: {{ .Values.config.dataProcess.rayActorNumCpus | quote }} - RAY_NUM_CPUS: {{ .Values.config.dataProcess.rayNumCpus | quote }} + DP_PART_PROCESSOR_COUNT: {{ .Values.config.dataProcess.partProcessorCount | quote }} + DP_FILE_SPLIT_SIZE_MB: {{ .Values.config.dataProcess.fileSplitSizeMb | quote }} + RAY_NUM_CPUS: {{ default (mul (int .Values.config.dataProcess.partProcessorCount) (int .Values.config.dataProcess.rayActorNumCpus)) .Values.config.dataProcess.rayNumCpus | quote }} RAY_OBJECT_STORE_MEMORY_GB: {{ .Values.config.dataProcess.rayObjectStoreMemoryGb | quote }} RAY_TEMP_DIR: {{ .Values.config.dataProcess.rayTempDir | quote }} RAY_LOG_LEVEL: {{ .Values.config.dataProcess.rayLogLevel | quote }} @@ -116,7 +118,7 @@ data: ELASTICSEARCH_REQUEST_TIMEOUT: {{ .Values.config.dataProcess.elasticsearchRequestTimeout | quote }} QUEUES: {{ .Values.config.dataProcess.queues | quote }} WORKER_NAME: {{ .Values.config.dataProcess.workerName | quote }} - WORKER_CONCURRENCY: {{ .Values.config.dataProcess.workerConcurrency | quote }} + WORKER_CONCURRENCY: {{ default (add (int .Values.config.dataProcess.partProcessorCount) 1) .Values.config.dataProcess.workerConcurrency | quote }} # Telemetry and Monitoring Configuration ENABLE_TELEMETRY: {{ ternary (get $monitoring "enabled") .Values.config.telemetry.enabled (hasKey $monitoring "enabled") | quote }} diff --git a/deploy/k8s/helm/nexent/charts/nexent-common/values.yaml b/deploy/k8s/helm/nexent/charts/nexent-common/values.yaml index b8f8aded1e..8f0d6430d2 100644 --- a/deploy/k8s/helm/nexent/charts/nexent-common/values.yaml +++ b/deploy/k8s/helm/nexent/charts/nexent-common/values.yaml @@ -113,7 +113,9 @@ config: rayDashboardPort: "8265" rayDashboardHost: "0.0.0.0" rayActorNumCpus: "2" - rayNumCpus: "4" + partProcessorCount: "3" + fileSplitSizeMb: "5" + rayNumCpus: "" rayObjectStoreMemoryGb: "0.25" rayTempDir: "/tmp/ray" rayLogLevel: "INFO" @@ -124,9 +126,9 @@ config: celeryWorkerPrefetchMultiplier: "1" celeryTaskTimeLimit: "3600" elasticsearchRequestTimeout: "30" - queues: "process_q,forward_q" + queues: "process_q,process_part_q,forward_q" workerName: "" - workerConcurrency: "4" + workerConcurrency: "" telemetry: enabled: "false" provider: "otlp" diff --git a/deploy/k8s/helm/nexent/charts/nexent-monitoring/templates/otel-collector-configmap.yaml b/deploy/k8s/helm/nexent/charts/nexent-monitoring/templates/otel-collector-configmap.yaml index 74bab1ba60..7598aa7664 100644 --- a/deploy/k8s/helm/nexent/charts/nexent-monitoring/templates/otel-collector-configmap.yaml +++ b/deploy/k8s/helm/nexent/charts/nexent-monitoring/templates/otel-collector-configmap.yaml @@ -66,10 +66,26 @@ data: attributes: - key: service.name value: nexent-backend - action: upsert + action: insert - key: service.version from_attribute: version action: insert + transform/knowledge_project: + error_mode: ignore + trace_statements: + - context: span + statements: + - 'set(resource.attributes["openinference.project.name"], "knowledge-base-monitor") where IsMatch(span.name, "^knowledge\\.")' + filter/knowledge_only: + error_mode: ignore + traces: + span: + - 'not IsMatch(span.name, "^knowledge\\.")' + filter/non_knowledge: + error_mode: ignore + traces: + span: + - 'IsMatch(span.name, "^knowledge\\.")' exporters: debug: verbosity: normal @@ -85,12 +101,30 @@ data: initial_interval: 1s max_interval: 30s max_elapsed_time: 300s + otlphttp/phoenix_knowledge: + endpoint: http://nexent-phoenix:6006 + headers: + x-project-name: knowledge-base-monitor + timeout: 5s + sending_queue: + enabled: true + num_consumers: 4 + queue_size: 2000 + retry_on_failure: + enabled: true + initial_interval: 1s + max_interval: 30s + max_elapsed_time: 300s service: pipelines: traces: receivers: [otlp] - processors: [memory_limiter, resource, batch] + processors: [memory_limiter, resource, filter/non_knowledge, batch] exporters: [otlphttp/phoenix, debug] + traces/knowledge: + receivers: [otlp] + processors: [memory_limiter, resource, filter/knowledge_only, transform/knowledge_project, batch] + exporters: [otlphttp/phoenix_knowledge, debug] metrics: receivers: [otlp] processors: [memory_limiter, resource, batch] diff --git a/sdk/nexent/data_process/core.py b/sdk/nexent/data_process/core.py index b8f366fab1..1f8b2a32be 100644 --- a/sdk/nexent/data_process/core.py +++ b/sdk/nexent/data_process/core.py @@ -9,10 +9,12 @@ from .file_splitter import FileSplitter from .openpyxl_processor import OpenPyxlProcessor from .unstructured_processor import UnstructuredProcessor +from nexent.monitor import get_monitoring_manager logger = logging.getLogger("data_process.core") logger.setLevel(logging.INFO) +monitoring_manager = get_monitoring_manager() class DataProcessCore: @@ -117,9 +119,19 @@ def file_process( if not processor_instance: raise ValueError(f"Unsupported processor: {processor_name}") + extension = os.path.splitext(filename)[1].lower() if extract_image_processor_instance: - img_info = extract_image_processor_instance.process_file( - file_data, chunking_strategy, filename, **params) + with monitoring_manager.trace_operation( + "knowledge.preprocess.image_extract", + **{ + "ingestion.stage": "preprocess.image_extract", + "file.extension": extension, + "file.size_bytes": len(file_data), + "processor.name": extractor, + }, + ): + img_info = extract_image_processor_instance.process_file( + file_data, chunking_strategy, filename, **params) else: img_info = [] @@ -127,7 +139,20 @@ def file_process( logger.info( f"Processing in-memory file: {filename} with {processor_name} processor") try: - return processor_instance.process_file(file_data, chunking_strategy, filename=filename, **params), img_info + with monitoring_manager.trace_operation( + "knowledge.preprocess.typed", + **{ + "ingestion.stage": "preprocess.typed", + "file.extension": extension, + "file.size_bytes": len(file_data), + "processor.name": processor_name, + }, + ): + chunks = processor_instance.process_file( + file_data, chunking_strategy, filename=filename, **params + ) + monitoring_manager.set_span_attributes(**{"chunk.count": len(chunks)}) + return chunks, img_info except Exception as e: logger.error(f"File processing failed for {filename}: {str(e)}") raise @@ -146,7 +171,7 @@ def file_split( file_data: File content byte data filename: Filename splitter: Optional splitter name (reserved for future use) - **params: Additional splitter parameters (e.g., max_size, encoding, libreoffice_path) + **params: Additional splitter parameters (e.g., max_size, target_parts, encoding, libreoffice_path) Returns: List of BytesIO parts @@ -165,10 +190,27 @@ def file_split( logger.error(f"Splitter not found: {splitter_name}") return [BytesIO(file_data)] - max_size = params.pop("max_size", 5 * 1024 * 1024) + max_size = params.pop("max_size", None) + target_parts = params.pop("target_parts", None) try: - parts = splitter_instance.file_process(file_data, filename, max_size=max_size, **params) + with monitoring_manager.trace_operation( + "knowledge.preprocess.split", + **{ + "ingestion.stage": "preprocess.split", + "file.extension": ext, + "file.size_bytes": len(file_data), + "processor.name": splitter_name, + }, + ): + parts = splitter_instance.file_process( + file_data, + filename, + max_size=max_size, + target_parts=target_parts, + **params, + ) + monitoring_manager.set_span_attributes(**{"file.parts_count": len(parts)}) if not isinstance(parts, list) or not all(isinstance(p, BytesIO) for p in parts): logger.error("Invalid split result format: expected List[BytesIO]") return [BytesIO(file_data)] diff --git a/sdk/nexent/data_process/file_splitter.py b/sdk/nexent/data_process/file_splitter.py index 3572e76035..fc3de720bc 100644 --- a/sdk/nexent/data_process/file_splitter.py +++ b/sdk/nexent/data_process/file_splitter.py @@ -12,6 +12,12 @@ class FileSplitter: + @staticmethod + def _resolve_max_size(file_data, max_size=None, target_parts=None): + if target_parts is not None: + return max(1, math.ceil(len(file_data) / max(1, int(target_parts)))) + return max(1, int(max_size or 5 * 1024 * 1024)) + def split_csv_by_size(self, csv_bytes, max_size, encoding="utf-8"): text = csv_bytes.decode(encoding) reader = list(csv.reader(StringIO(text))) @@ -349,6 +355,30 @@ def split_range(start, end): return result + def split_pdf_by_parts(self, pdf_bytes, target_parts): + from pypdf import PdfReader, PdfWriter + + reader = PdfReader(BytesIO(pdf_bytes)) + total_pages = len(reader.pages) + if total_pages == 0: + return [] + + group_count = min(max(1, int(target_parts)), total_pages) + base_pages, extra_pages = divmod(total_pages, group_count) + result = [] + start = 0 + for group_index in range(group_count): + page_count = base_pages + (1 if group_index < extra_pages else 0) + end = start + page_count + writer = PdfWriter() + for page_index in range(start, end): + writer.add_page(reader.pages[page_index]) + buffer = BytesIO() + writer.write(buffer) + result.append(BytesIO(buffer.getvalue())) + start = end + return result + def split_txt_by_size(self, txt_bytes, max_size, encoding="utf-8"): buffer = BytesIO(txt_bytes) @@ -456,7 +486,9 @@ def _convert_bytes_with_libreoffice( with open(output_path, "rb") as f: return f.read() - def file_process(self, file_data, filename, max_size, **kwargs) -> List[BytesIO]: + def file_process( + self, file_data, filename, max_size=None, target_parts=None, **kwargs + ) -> List[BytesIO]: ext = os.path.splitext(filename)[1].lower() if ext in {".doc", ".docx"}: @@ -464,7 +496,13 @@ def file_process(self, file_data, filename, max_size, **kwargs) -> List[BytesIO] pdf_bytes = self._convert_bytes_with_libreoffice( file_data, ext, ".pdf", libreoffice_path=libreoffice_path ) - pdf_parts = self.split_pdf_by_size(pdf_bytes, max_size=max_size) + if target_parts is not None: + pdf_parts = self.split_pdf_by_parts(pdf_bytes, target_parts) + else: + pdf_parts = self.split_pdf_by_size( + pdf_bytes, + max_size=self._resolve_max_size(pdf_bytes, max_size=max_size), + ) # If no actual split happened, keep original Word bytes as-is. if not pdf_parts or len(pdf_parts) == 1: @@ -474,36 +512,41 @@ def file_process(self, file_data, filename, max_size, **kwargs) -> List[BytesIO] # while filenames remain as Word (handled by caller). return pdf_parts + effective_max_size = self._resolve_max_size( + file_data, max_size=max_size, target_parts=target_parts) + if ext == ".csv": return self.split_csv_by_size( file_data, - max_size=max_size, + max_size=effective_max_size, encoding=kwargs.get("encoding", "utf-8"), ) if ext == ".epub": - return self.split_epub_by_size(file_data, max_size=max_size) + return self.split_epub_by_size(file_data, max_size=effective_max_size) if ext in {".xlsx", ".xls"}: - return self.split_excel(file_data, max_size=max_size) + return self.split_excel(file_data, max_size=effective_max_size) if ext == ".json": - return self.split_json_stream(file_data, max_size=max_size) + return self.split_json_stream(file_data, max_size=effective_max_size) if ext == ".md": - return self.split_markdown(file_data, max_size=max_size) + return self.split_markdown(file_data, max_size=effective_max_size) if ext == ".pdf": - return self.split_pdf_by_size(file_data, max_size=max_size) + if target_parts is not None: + return self.split_pdf_by_parts(file_data, target_parts) + return self.split_pdf_by_size(file_data, max_size=effective_max_size) if ext == ".txt": return self.split_txt_by_size( file_data, - max_size=max_size, + max_size=effective_max_size, encoding=kwargs.get("encoding", "utf-8"), ) if ext == ".xml": - return self.split_xml_by_size(file_data, max_size=max_size) + return self.split_xml_by_size(file_data, max_size=effective_max_size) raise ValueError(f"Unsupported file extension: {ext}") diff --git a/sdk/nexent/monitor/monitoring.py b/sdk/nexent/monitor/monitoring.py index 765c09abb9..41dfee046d 100644 --- a/sdk/nexent/monitor/monitoring.py +++ b/sdk/nexent/monitor/monitoring.py @@ -964,6 +964,8 @@ def trace_operation( span.set_attribute("error.type", type(e).__name__) span.set_attribute("error.message", str(e)) raise + else: + span.set_status(Status(StatusCode.OK)) def set_openinference_output( self, diff --git a/test/backend/data_process/test_tasks.py b/test/backend/data_process/test_tasks.py index 53dc267f85..1838901383 100644 --- a/test/backend/data_process/test_tasks.py +++ b/test/backend/data_process/test_tasks.py @@ -370,13 +370,14 @@ def connect(self, func): const_mod.DATA_PROCESS_SERVICE = "http://data-process" const_mod.RAY_ACTOR_NUM_CPUS = 1 const_mod.RAY_NUM_CPUS = 4 + const_mod.DP_PART_PROCESSOR_COUNT = 3 + const_mod.DP_FILE_SPLIT_SIZE_MB = 5 const_mod.FORWARD_REDIS_RETRY_DELAY_S = 0 const_mod.FORWARD_REDIS_RETRY_MAX = 1 const_mod.DP_REDIS_CHUNKS_WAIT_TIMEOUT_S = 30 const_mod.DP_REDIS_CHUNKS_POLL_INTERVAL_MS = 200 const_mod.PER_WAVE_TIMEOUT = 30 const_mod.MAX_TIMEOUT = 1800 - const_mod.RAY_GLOBAL_ACTOR_POOL_SIZE = 3 const_mod.RAY_ACTOR_WARM_TIMEOUT_S = 60 const_mod.RAY_GLOBAL_ACTOR_POOL_NAME = "nexent_global_data_processor_pool" const_mod.RAY_GLOBAL_ACTOR_POOL_NAMESPACE = "nexent-data-process" @@ -613,7 +614,6 @@ class _CeleryTaskShim: "RAY_GLOBAL_ACTOR_POOL_NAME", "RAY_GLOBAL_ACTOR_POOL_NAMESPACE", "RAY_ACTOR_WARM_TIMEOUT_S", - "RAY_GLOBAL_ACTOR_POOL_SIZE", "RAY_ACTOR_NUM_CPUS", "ROOT_DIR", "PER_WAVE_TIMEOUT", @@ -974,7 +974,7 @@ def test_process_minio_path(monkeypatch): class FakeActor: def __init__(self): - self.process_file = types.SimpleNamespace( + self.process_bytes = types.SimpleNamespace( remote=lambda *a, **k: "ref") self.store_chunks_in_redis = types.SimpleNamespace( remote=lambda *a, **k: None) @@ -2384,7 +2384,7 @@ def test_process_url_source_with_many_chunks(monkeypatch): class FakeActor: def __init__(self): - self.process_file = types.SimpleNamespace( + self.process_bytes = types.SimpleNamespace( remote=lambda *a, **k: "ref_url") self.store_chunks_in_redis = types.SimpleNamespace( remote=lambda *a, **k: None) @@ -2508,7 +2508,11 @@ def test_estimate_parallel_parts_and_batch_helpers(monkeypatch): tasks, _ = import_tasks_with_fake_ray(monkeypatch) monkeypatch.setattr(tasks, "RAY_NUM_CPUS", 8) monkeypatch.setattr(tasks, "RAY_ACTOR_NUM_CPUS", 2) - assert tasks._estimate_parallel_parts() == 4 + monkeypatch.setattr(tasks, "DP_PART_PROCESSOR_COUNT", 3) + assert tasks._estimate_parallel_parts() == 3 + + monkeypatch.setattr(tasks, "RAY_NUM_CPUS", 4) + assert tasks._estimate_parallel_parts() == 2 batches = [[{"a": 1}], [{"a": 2}]] assert tasks._get_next_available_batch_index(batches, 0, batch_size=2) == 0 @@ -2516,6 +2520,125 @@ def test_estimate_parallel_parts_and_batch_helpers(monkeypatch): tasks._get_next_available_batch_index([[1], [2]], 0, batch_size=1) +def test_split_file_for_processing_targets_processor_count(monkeypatch): + tasks, fake_ray = import_tasks_with_fake_ray(monkeypatch) + captured = {} + spans = [] + + class CapturedSpan: + def __init__(self): + self.attributes = {} + + def set_attribute(self, name, value): + self.attributes[name] = value + + @contextmanager + def capture_span(name, stage, **attributes): + captured_span = CapturedSpan() + spans.append((name, stage, attributes, captured_span)) + yield captured_span + + class Actor: + split_file = types.SimpleNamespace( + remote=lambda **kwargs: captured.update(kwargs) or "parts-ref") + + monkeypatch.setattr(tasks, "DP_PART_PROCESSOR_COUNT", 4) + monkeypatch.setattr(tasks, "_get_split_actor", lambda: Actor()) + monkeypatch.setattr(tasks, "knowledge_span", capture_span) + fake_ray.get_returns = {"parts-ref": [b"part"]} + + params = {"max_size": 1, "encoding": "utf-8"} + parts = tasks._split_file_for_processing( + request_id="req", + source="file.txt", + source_type="local", + task_id="task", + params=params, + file_size_bytes=10 * 1024 * 1024, + ) + + assert parts == [b"part"] + assert captured["target_parts"] == 4 + assert "max_size" not in captured + assert "max_size" not in params + assert [span[0] for span in spans] == [ + "knowledge.process.split_actor_acquire", + "knowledge.process.file_split_rpc", + ] + assert spans[1][2]["processor_count"] == 4 + assert spans[1][3].attributes["file.parts_count"] == 1 + + +def test_process_source_below_split_threshold_skips_splitter(monkeypatch): + tasks, fake_ray = import_tasks_with_fake_ray(monkeypatch) + + class Actor: + process_bytes = types.SimpleNamespace(remote=lambda *args, **kwargs: "chunks-ref") + store_chunks_in_redis = types.SimpleNamespace(remote=lambda *args, **kwargs: "store-ref") + + monkeypatch.setattr(tasks, "DP_FILE_SPLIT_SIZE_MB", 5) + monkeypatch.setattr(tasks, "get_ray_actor", lambda: Actor()) + monkeypatch.setattr( + tasks, + "_split_file_for_processing", + lambda **kwargs: (_ for _ in ()).throw(AssertionError("splitter called")), + ) + fake_ray.get_returns = { + "chunks-ref": [{"content": "chunk"}], + "store-ref": True, + } + + result = tasks._process_source_with_split( + request_id="req", + source="file.txt", + source_type="minio", + task_id="task", + chunking_strategy="basic", + index_name="idx", + original_filename="file.txt", + embedding_model_id=None, + tenant_id=None, + params={}, + file_data=b"small file", + ) + + assert result == (False, [{"content": "chunk"}], None) + + +def test_process_source_above_split_threshold_uses_splitter(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + captured = {} + + monkeypatch.setattr(tasks, "DP_FILE_SPLIT_SIZE_MB", 1) + monkeypatch.setattr( + tasks, + "_split_file_for_processing", + lambda **kwargs: captured.update(kwargs) or [b"a", b"b"], + ) + monkeypatch.setattr( + tasks, + "_run_processing_for_parts", + lambda **kwargs: (True, None, 2), + ) + + result = tasks._process_source_with_split( + request_id="req", + source="file.txt", + source_type="minio", + task_id="task", + chunking_strategy="basic", + index_name="idx", + original_filename="file.txt", + embedding_model_id=None, + tenant_id=None, + params={}, + file_data=b"x" * (1024 * 1024 + 1), + ) + + assert result == (True, None, 2) + assert captured["file_size_bytes"] == 1024 * 1024 + 1 + + def test_extract_error_code_from_es_response_detail_string(monkeypatch): tasks, _ = import_tasks_with_fake_ray(monkeypatch) parsed = {"detail": "{\"error_code\":\"es_detail_code\"}"} @@ -2572,6 +2695,15 @@ def __init__(self): manager = tasks.GlobalRayActorPoolManager(warm_timeout_s=1) assert manager.ensure_pool(desired=2, max_allowed=3) == 2 assert manager.get_actor() is not None + killed = [] + monkeypatch.setattr( + tasks.ray, + "kill", + lambda actor, **kwargs: killed.append(actor), + raising=False, + ) + assert manager.ensure_pool(desired=1, max_allowed=3) == 1 + assert len(killed) == 1 def test_global_pool_manager_warm_fail(monkeypatch): @@ -2807,6 +2939,14 @@ def __init__(self): assert split_chunk_count is None captured = {} + spans = [] + + @contextmanager + def capture_span(name, stage, **attributes): + spans.append((name, stage, attributes)) + yield None + + monkeypatch.setattr(tasks, "knowledge_span", capture_span) monkeypatch.setattr(tasks, "process_part", types.SimpleNamespace( s=lambda **kwargs: types.SimpleNamespace(kwargs=kwargs))) monkeypatch.setattr(tasks, "aggregate_store_chunks", types.SimpleNamespace( @@ -2836,6 +2976,12 @@ def __init__(self): assert chunks2 is None assert split_chunk_count2 == 6 assert len(captured["group"]) == 3 + assert [span[0] for span in spans] == [ + "knowledge.process.part_dispatch", + "knowledge.process.part_wait", + ] + assert spans[0][2]["part_count"] == 3 + assert spans[1][2]["timeout_seconds"] == 9 def test_process_split_async_redis_image_metadata_count(monkeypatch, tmp_path): diff --git a/test/backend/data_process/test_worker.py b/test/backend/data_process/test_worker.py index 444d9c799e..f88fd588a3 100644 --- a/test/backend/data_process/test_worker.py +++ b/test/backend/data_process/test_worker.py @@ -68,9 +68,9 @@ def setup_mocks_for_worker(mocker, initialized=False): const_mod.DP_REDIS_CHUNKS_POLL_INTERVAL_MS = 100 const_mod.RAY_ACTOR_NUM_CPUS = 1 const_mod.RAY_NUM_CPUS = 4 + const_mod.DP_PART_PROCESSOR_COUNT = 3 const_mod.PER_WAVE_TIMEOUT = 300 const_mod.MAX_TIMEOUT = 3600 - const_mod.RAY_GLOBAL_ACTOR_POOL_SIZE = 10 const_mod.RAY_ACTOR_WARM_TIMEOUT_S = 60 const_mod.RAY_GLOBAL_ACTOR_POOL_NAME = "global_actor_pool" const_mod.RAY_GLOBAL_ACTOR_POOL_NAMESPACE = "nexent" diff --git a/test/backend/utils/test_knowledge_telemetry.py b/test/backend/utils/test_knowledge_telemetry.py new file mode 100644 index 0000000000..0b8cfe61ee --- /dev/null +++ b/test/backend/utils/test_knowledge_telemetry.py @@ -0,0 +1,63 @@ +from unittest.mock import MagicMock, patch + +import pytest + +from backend.utils import knowledge_telemetry + + +def test_span_kind_distinguishes_storage_tools_from_chains(): + assert knowledge_telemetry._span_kind("knowledge.minio.fetch") == "TOOL" + assert knowledge_telemetry._span_kind("knowledge.forward.redis_read") == "TOOL" + assert knowledge_telemetry._span_kind("knowledge.forward.elasticsearch") == "TOOL" + assert knowledge_telemetry._span_kind("knowledge.process") == "CHAIN" + + +def test_safe_attributes_maps_part_diagnostics(): + attrs = knowledge_telemetry._safe_attributes({ + "part_count": 4, + "processor_count": 4, + "parallel_parts": 3, + "timeout_seconds": 300, + "poll_interval_ms": 200, + "queue_name": "process_part_q", + }) + + assert attrs == { + "file.parts_count": 4, + "processor.count": 4, + "processor.parallel_count": 3, + "timeout.seconds": 300, + "poll.interval_ms": 200, + "messaging.destination.name": "process_part_q", + } + + +def test_knowledge_span_marks_celery_retry_without_error(): + Retry = type("Retry", (Exception,), {"__module__": "celery.exceptions"}) + span = MagicMock() + span_cm = MagicMock() + span_cm.__enter__.return_value = span + + with ( + patch.object(knowledge_telemetry, "OTEL_AVAILABLE", True), + patch.object(knowledge_telemetry.trace, "get_tracer") as get_tracer, + patch.object(knowledge_telemetry, "_resource_snapshot", return_value={}), + patch.object(knowledge_telemetry, "_record_metrics"), + ): + get_tracer.return_value.start_as_current_span.return_value = span_cm + with pytest.raises(Retry): + with knowledge_telemetry.knowledge_span( + "knowledge.forward.redis_read", + "forward.redis_read", + retry_attempt=2, + retry_delay_seconds=5, + ): + raise Retry() + + span.record_exception.assert_not_called() + span.set_attribute.assert_any_call("ingestion.status", "retry") + span.set_attribute.assert_any_call("retry.attempt", 2) + span.set_attribute.assert_any_call("retry.delay_seconds", 5.0) + span.set_status.assert_called_with( + knowledge_telemetry.Status(knowledge_telemetry.StatusCode.OK) + ) diff --git a/test/sdk/data_process/test_core.py b/test/sdk/data_process/test_core.py index e0edced14e..5576afe48e 100644 --- a/test/sdk/data_process/test_core.py +++ b/test/sdk/data_process/test_core.py @@ -403,8 +403,8 @@ def test_file_split_unsupported_extension_returns_original_bytes(self, core): assert isinstance(parts[0], BytesIO) assert parts[0].getvalue() == data - def test_file_split_uses_splitter_with_default_max_size(self, core): - """file_split should call FileSplitter with default max_size when omitted.""" + def test_file_split_passes_optional_split_parameters(self, core): + """file_split should pass optional split parameters explicitly.""" splitter = Mock() splitter.file_process.return_value = [BytesIO(b"p1"), BytesIO(b"p2")] core.processors["FileSplitter"] = splitter @@ -413,7 +413,19 @@ def test_file_split_uses_splitter_with_default_max_size(self, core): assert len(parts) == 2 splitter.file_process.assert_called_once_with( - b"csv-data", "data.csv", max_size=5 * 1024 * 1024 + b"csv-data", "data.csv", max_size=None, target_parts=None + ) + + def test_file_split_passes_target_parts(self, core): + splitter = Mock() + splitter.file_process.return_value = [BytesIO(b"p1"), BytesIO(b"p2")] + core.processors["FileSplitter"] = splitter + + parts = core.file_split(b"csv-data", "data.csv", target_parts=2) + + assert len(parts) == 2 + splitter.file_process.assert_called_once_with( + b"csv-data", "data.csv", max_size=None, target_parts=2 ) def test_file_split_invalid_split_result_falls_back(self, core): diff --git a/test/sdk/data_process/test_file_splitter.py b/test/sdk/data_process/test_file_splitter.py index 5c44131d78..e645bf10e0 100644 --- a/test/sdk/data_process/test_file_splitter.py +++ b/test/sdk/data_process/test_file_splitter.py @@ -48,6 +48,59 @@ def test_file_process_docx_multi_parts_returns_pdf_parts(monkeypatch): assert parts == expected_parts +def test_file_process_docx_target_parts_uses_converted_pdf(monkeypatch): + splitter = FileSplitter() + captured = {} + expected_parts = [BytesIO(b"p1"), BytesIO(b"p2"), BytesIO(b"p3")] + monkeypatch.setattr( + splitter, + "_convert_bytes_with_libreoffice", + lambda *args, **kwargs: b"converted-pdf-bytes", + ) + + def _split_pdf(pdf_bytes, target_parts): + captured["pdf_bytes"] = pdf_bytes + captured["target_parts"] = target_parts + return expected_parts + + monkeypatch.setattr(splitter, "split_pdf_by_parts", _split_pdf) + + parts = splitter.file_process( + b"compressed-word-bytes", "sample.docx", target_parts=3) + + assert parts == expected_parts + assert captured == { + "pdf_bytes": b"converted-pdf-bytes", + "target_parts": 3, + } + + +def test_split_pdf_by_parts_caps_output_at_target(monkeypatch): + splitter = FileSplitter() + + class FakeReader: + def __init__(self, *_args, **_kwargs): + self.pages = [object() for _ in range(11)] + + class FakeWriter: + def __init__(self): + self.pages = [] + + def add_page(self, page): + self.pages.append(page) + + def write(self, buffer): + buffer.write(b"x" * len(self.pages)) + + monkeypatch.setattr("pypdf.PdfReader", FakeReader) + monkeypatch.setattr("pypdf.PdfWriter", FakeWriter) + + parts = splitter.split_pdf_by_parts(b"%PDF", target_parts=5) + + assert len(parts) == 5 + assert [len(part.getvalue()) for part in parts] == [3, 2, 2, 2, 2] + + def test_file_process_csv_routes_to_split_csv(monkeypatch): splitter = FileSplitter() captured = {} From 0682800bb57b2a990ba00890ef1d79d67e4287d6 Mon Sep 17 00:00:00 2001 From: Jasonxia007 Date: Fri, 14 Aug 2026 11:05:01 +0800 Subject: [PATCH 2/5] =?UTF-8?q?=F0=9F=A7=AA=20Add=20test=20files?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- test/backend/data_process/test_ray_actors.py | 331 +++++ test/backend/data_process/test_tasks.py | 596 +++++++++ test/backend/data_process/test_worker.py | 426 ++++++ .../test_data_process_service_entrypoint.py | 269 ++++ test/backend/test_model_consts.py | 1188 +++++++++++++++++ .../utils/test_file_management_utils.py | 18 + .../backend/utils/test_knowledge_telemetry.py | 206 ++- test/backend/utils/test_monitoring.py | 72 +- test/sdk/data_process/test_core.py | 16 + test/sdk/data_process/test_file_splitter.py | 92 ++ .../test_file_splitter_coverage.py | 181 +++ 11 files changed, 3375 insertions(+), 20 deletions(-) create mode 100644 test/backend/test_data_process_service_entrypoint.py create mode 100644 test/sdk/data_process/test_file_splitter_coverage.py diff --git a/test/backend/data_process/test_ray_actors.py b/test/backend/data_process/test_ray_actors.py index 79a2f5bb95..0a2579930a 100644 --- a/test/backend/data_process/test_ray_actors.py +++ b/test/backend/data_process/test_ray_actors.py @@ -716,3 +716,334 @@ def file_split(self, *a, **k): actor = ray_actors.DataProcessorRayActor() assert actor.split_file("x.txt", "local", file_data=b"abc") == [] + +def test_ping_returns_true(monkeypatch): + """Test that ping() returns True for health check.""" + ray_actors = import_module(monkeypatch) + actor = ray_actors.DataProcessorRayActor() + assert actor.ping() is True + + +def test_normalize_processor_result_variants(monkeypatch): + """Test _normalize_processor_result handles various return types.""" + ray_actors = import_module(monkeypatch) + actor = ray_actors.DataProcessorRayActor() + + # Tuple with both chunks and images + result1 = ([{"content": "a"}], [{"image": "b"}]) + chunks, images = actor._normalize_processor_result(result1) + assert chunks == [{"content": "a"}] + assert images == [{"image": "b"}] + + # Empty tuple + result2 = ([], []) + chunks, images = actor._normalize_processor_result(result2) + assert chunks == [] + assert images == [] + + # None result + result3 = None + chunks, images = actor._normalize_processor_result(result3) + assert chunks == [] + assert images == [] + + # List only (not a tuple) + result4 = [{"content": "list-only"}] + chunks, images = actor._normalize_processor_result(result4) + assert chunks == [{"content": "list-only"}] + assert images == [] + + # Empty list + result5 = [] + chunks, images = actor._normalize_processor_result(result5) + assert chunks == [] + assert images == [] + + +def test_validate_chunks_variants(monkeypatch): + """Test _validate_chunks handles edge cases.""" + ray_actors = import_module(monkeypatch) + actor = ray_actors.DataProcessorRayActor() + + # None chunks + result = actor._validate_chunks(None, "source.txt") + assert result == [] + + # Non-list type + result = actor._validate_chunks("string", "source.txt") + assert result == [] + + # Empty list + result = actor._validate_chunks([], "source.txt") + assert result == [] + + # Valid list + valid_chunks = [{"content": "valid"}] + result = actor._validate_chunks(valid_chunks, "source.txt") + assert result == valid_chunks + + +def test_append_image_chunks_skips_invalid_entries(monkeypatch): + """Test _append_image_chunks skips non-dict and missing-image_bytes entries.""" + ray_actors = import_module(monkeypatch) + + class CoreWithBadImages: + def file_process(self, *a, **k): + return ( + [{"content": "text", "metadata": {}}], + [ + {"not": "dict"}, # Not a dict + {"image_bytes": b"img"}, # Missing image_format + ], + ) + + monkeypatch.setattr(ray_actors, "DataProcessCore", CoreWithBadImages) + monkeypatch.setattr( + ray_actors, + "upload_fileobj", + lambda file_obj, file_name, prefix=None: {"object_name": f"{prefix}/{file_name}"}, + ) + monkeypatch.setattr( + ray_actors, + "build_s3_url", + lambda object_name: f"s3://bucket/{object_name}", + ) + + actor = ray_actors.DataProcessorRayActor() + chunks = [{"content": "text", "metadata": {}}] + images = [ + {"not": "dict"}, + {"image_bytes": b"img"}, + ] + actor._append_image_chunks("source.pdf", chunks, images) + # Only valid text chunk should remain, no image chunks added + assert len(chunks) == 1 + assert chunks[0]["content"] == "text" + + +def test_apply_model_paths_sets_correct_keys(monkeypatch): + """Test _apply_model_paths sets the required model path keys.""" + ray_actors = import_module(monkeypatch) + actor = ray_actors.DataProcessorRayActor() + params = {} + actor._apply_model_paths(params) + assert "table_transformer_model_path" in params + assert "unstructured_default_model_initialize_params_json_path" in params + assert params["table_transformer_model_path"] == "/models/table" + assert params["unstructured_default_model_initialize_params_json_path"] == "/models/unstructured.json" + + +def test_process_bytes_with_minio_source(monkeypatch): + """Test process_bytes with minio source fetching file data.""" + ray_actors = import_module(monkeypatch) + + class CoreRecords: + captured = {} + + def __init__(self): + pass + + def file_process(self, file_data, filename, chunking_strategy, **params): + CoreRecords.captured = { + "file_data": file_data, + "filename": filename, + "chunking_strategy": chunking_strategy, + } + return [{"content": "processed", "metadata": {}}] + + monkeypatch.setattr(ray_actors, "DataProcessCore", CoreRecords) + actor = ray_actors.DataProcessorRayActor() + + # With file_data provided directly + chunks = actor.process_bytes( + b"file bytes content", + "test.pdf", + "basic", + task_id="task-123", + model_id=5, + tenant_id="tenant-1" + ) + assert len(chunks) == 1 + assert CoreRecords.captured["filename"] == "test.pdf" + assert CoreRecords.captured["chunking_strategy"] == "basic" + + +def test_split_file_logs_timing_and_parts(monkeypatch, caplog): + """Test split_file logs timing and part statistics.""" + ray_actors = import_module(monkeypatch) + + class PartBytes: + def __init__(self, data): + self._data = data + + def getvalue(self): + return self._data + + class CoreWithSplit: + def file_split(self, *a, **k): + # Return 3 parts with different sizes + return [ + PartBytes(b"part1 data here"), + PartBytes(b"part2 data"), + PartBytes(b"part3"), + ] + + monkeypatch.setattr(ray_actors, "DataProcessCore", CoreWithSplit) + actor = ray_actors.DataProcessorRayActor() + parts = actor.split_file("large.pdf", "local", file_data=b"large file content") + + assert len(parts) == 3 + assert parts[0] == b"part1 data here" + assert parts[1] == b"part2 data" + assert parts[2] == b"part3" + + +def test_split_file_handles_exception_in_getvalue(monkeypatch): + """Test split_file continues when part.getvalue() raises.""" + ray_actors = import_module(monkeypatch) + + class PartGood: + def getvalue(self): + return b"good part" + + class PartBad: + def getvalue(self): + raise RuntimeError("getvalue failed") + + class CoreWithBadParts: + def file_split(self, *a, **k): + return [PartGood(), PartBad(), PartGood()] + + monkeypatch.setattr(ray_actors, "DataProcessCore", CoreWithBadParts) + actor = ray_actors.DataProcessorRayActor() + parts = actor.split_file("test.pdf", "local", file_data=b"content") + + # Only good parts should be returned + assert len(parts) == 2 + assert parts[0] == b"good part" + assert parts[1] == b"good part" + + +def test_store_chunks_in_redis_with_various_chunks(monkeypatch): + """Test store_chunks_in_redis with various chunk inputs.""" + ray_actors = import_module(monkeypatch) + monkeypatch.setattr(ray_actors, "REDIS_BACKEND_URL", "redis://test") + + fake_client = FakeRedisClient() + fake_redis_module = types.SimpleNamespace( + Redis=types.SimpleNamespace(from_url=lambda *a, **k: fake_client) + ) + monkeypatch.setitem(sys.modules, "redis", fake_redis_module) + + actor = ray_actors.DataProcessorRayActor() + + # Empty list + ok = actor.store_chunks_in_redis("k-empty", []) + assert ok is True + assert json.loads(fake_client.get("k-empty")) == [] + + # List with various types + ok = actor.store_chunks_in_redis("k-mixed", [ + {"content": "text", "metadata": {"key": "value"}}, + {"content": "text2", "numbers": [1, 2, 3]}, + ]) + assert ok is True + stored = json.loads(fake_client.get("k-mixed")) + assert len(stored) == 2 + + # Verify expiration was set + assert "k-empty" in fake_client.expirations + assert fake_client.expirations["k-empty"] == 2 * 60 * 60 + + +def test_prepare_process_params(monkeypatch): + """Test _prepare_process_params applies model paths and chunk sizes.""" + ray_actors = import_module(monkeypatch) + + class RecorderCore: + captured_params = None + + def __init__(self): + pass + + def file_process(self, file_data, filename, chunking_strategy, **params): + RecorderCore.captured_params = params + return [{"content": "x", "metadata": {}}] + + monkeypatch.setattr(ray_actors, "DataProcessCore", RecorderCore) + monkeypatch.setattr( + ray_actors, + "get_model_by_model_id", + lambda model_id, tenant_id=None: { + "expected_chunk_size": 500, + "maximum_chunk_size": 1000, + "display_name": "test-model", + "model_type": "embedding", + }, + ) + + actor = ray_actors.DataProcessorRayActor() + params = {"extra_key": "extra_value"} + result = actor._prepare_process_params( + task_id="task-1", + model_id=5, + tenant_id="tenant-1", + params=params, + ) + + assert result["task_id"] == "task-1" + assert result["new_after_n_chars"] == 500 + assert result["max_characters"] == 1000 + assert result["model_type"] == "embedding" + assert result["table_transformer_model_path"] == "/models/table" + assert result["extra_key"] == "extra_value" + + +def test_run_file_process_with_telemetry_context(monkeypatch): + """Test _run_file_process uses knowledge_span for telemetry.""" + ray_actors = import_module(monkeypatch) + + captured_spans = [] + + class MockKnowledgeSpan: + def __init__(self, name, operation, **kwargs): + self.name = name + self.operation = operation + self.kwargs = kwargs + + def __enter__(self): + return self + + def __exit__(self, *args): + pass + + class RecordingCore: + def __init__(self): + self.calls = [] + + def file_process(self, file_data, filename, chunking_strategy, **params): + self.calls.append((filename, chunking_strategy, params)) + return [{"content": "test content", "metadata": {"creation_date": "2024-01-01"}}] + + monkeypatch.setattr(ray_actors, "DataProcessCore", RecordingCore) + monkeypatch.setattr( + ray_actors, + "knowledge_span", + MockKnowledgeSpan, + ) + + actor = ray_actors.DataProcessorRayActor() + result = actor._run_file_process( + file_data=b"test data", + filename="test.txt", + chunking_strategy="basic", + process_params={ + "telemetry_context": {"trace_id": "abc123"}, + "task_id": "task-1", + }, + log_subject="test", + ) + + assert len(result) == 1 + assert result[0]["content"] == "test content" + diff --git a/test/backend/data_process/test_tasks.py b/test/backend/data_process/test_tasks.py index 1838901383..1716265acc 100644 --- a/test/backend/data_process/test_tasks.py +++ b/test/backend/data_process/test_tasks.py @@ -3129,3 +3129,599 @@ def test_get_all_task_ids_uses_scan_instead_of_keys(monkeypatch): "task-1", "task-2", ] + + +def test_extract_error_code_various_formats(monkeypatch): + """Test extract_error_code handles different error formats.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # From parsed_error dict + result = tasks.extract_error_code( + "Some error message", + parsed_error={"error_code": "ERR_001"} + ) + assert result == "ERR_001" + + # From JSON string in reason + result = tasks.extract_error_code( + '{"error_code": "ERR_002"}', + parsed_error=None + ) + assert result == "ERR_002" + + # From nested detail + result = tasks.extract_error_code( + '{"detail": {"error_code": "ERR_003"}}', + parsed_error=None + ) + assert result == "ERR_003" + + # From regex pattern in raw string + result = tasks.extract_error_code( + 'Some error with "error_code": "ERR_004"', + parsed_error=None + ) + assert result == "ERR_004" + + # No error code found + result = tasks.extract_error_code("Plain error message", parsed_error=None) + assert result == "unknown_error" + + +def test_build_balanced_batches_various_sizes(monkeypatch): + """Test _build_balanced_batches with various input sizes.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # Empty input + result = tasks._build_balanced_batches([]) + assert result == [] + + # Single batch (below batch size) + chunks = [{"content": f"chunk_{i}"} for i in range(10)] + result = tasks._build_balanced_batches(chunks) + assert len(result) == 1 + + # Multiple batches + chunks = [{"content": f"chunk_{i}"} for i in range(200)] + result = tasks._build_balanced_batches(chunks) + assert len(result) > 1 + + # With image metadata chunks + chunks = [ + {"content": "text1", "process_source": "UniversalImageExtractor"}, + {"content": "text2"}, + {"content": "text3", "process_source": "UniversalImageExtractor"}, + {"content": "text4"}, + ] + result = tasks._build_balanced_batches(chunks, batch_size=2) + # Should distribute evenly + assert len(result) == 2 + + +def test_count_image_metadata_chunks(monkeypatch): + """Test _count_image_metadata_chunks counting.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # None input + result = tasks._count_image_metadata_chunks(None) + assert result == 0 + + # Empty list + result = tasks._count_image_metadata_chunks([]) + assert result == 0 + + # Mixed chunks + chunks = [ + {"content": "text1", "process_source": "UniversalImageExtractor"}, + {"content": "text2"}, + {"content": "text3", "metadata": {"process_source": "UniversalImageExtractor"}}, + {"content": "text4", "metadata": {"process_source": "Other"}}, + ] + result = tasks._count_image_metadata_chunks(chunks) + assert result == 2 + + +def test_compute_split_wait_timeout(monkeypatch): + """Test _compute_split_wait_timeout calculation.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + monkeypatch.setattr(tasks, "DP_REDIS_CHUNKS_WAIT_TIMEOUT_S", 30) + monkeypatch.setattr(tasks, "PER_WAVE_TIMEOUT", 60) + monkeypatch.setattr(tasks, "MAX_TIMEOUT", 300) + monkeypatch.setattr(tasks, "_estimate_parallel_parts", lambda: 2) + + # Single part (no waves) + result = tasks._compute_split_wait_timeout(1) + assert result == 30 + + # Multiple parts + result = tasks._compute_split_wait_timeout(10) + waves = math.ceil(10 / 2) + expected = min(300, 30 + max(0, waves - 1) * 60) + assert result == expected + + +def test_forward_context_creation(monkeypatch): + """Test _init_forward_context creates context correctly.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + ctx = tasks._init_forward_context( + task_id="task-1", + request_id="req-1", + start_time=1000.0, + source="/path/to/file.pdf", + index_name="test-index", + source_type="local", + original_filename="file.pdf", + ) + assert ctx.task_id == "task-1" + assert ctx.request_id == "req-1" + assert ctx.source == "/path/to/file.pdf" + assert ctx.index_name == "test-index" + assert ctx.original_filename == "file.pdf" + + +def test_is_forward_task_cancelled(monkeypatch): + """Test _is_forward_task_cancelled checks Redis flag.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + class MockRedisService: + def is_task_cancelled(self, task_id): + return task_id == "cancelled-task" + + monkeypatch.setattr(tasks, "get_redis_service", lambda: MockRedisService()) + + ctx = tasks._init_forward_context( + task_id="cancelled-task", + request_id="req", + start_time=1000.0, + source="s", + index_name="i", + source_type="local", + original_filename=None, + ) + + assert tasks._is_forward_task_cancelled(ctx) is True + + ctx2 = tasks._init_forward_context( + task_id="active-task", + request_id="req", + start_time=1000.0, + source="s", + index_name="i", + source_type="local", + original_filename=None, + ) + assert tasks._is_forward_task_cancelled(ctx2) is False + + +def test_build_forward_cancelled_result(monkeypatch): + """Test _build_forward_cancelled_result creates correct response.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + ctx = tasks._init_forward_context( + task_id="task-cancel", + request_id="req", + start_time=1000.0, + source="/path/to/file.pdf", + index_name="test-index", + source_type="local", + original_filename="file.pdf", + ) + + result = tasks._build_forward_cancelled_result(ctx) + assert result["task_id"] == "task-cancel" + assert result["source"] == "/path/to/file.pdf" + assert result["index_name"] == "test-index" + assert result["es_result"]["success"] is False + assert "cancelled" in result["es_result"]["message"] + + +def test_build_forward_error(monkeypatch): + """Test _build_forward_error creates exception with correct structure.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + exc = tasks._build_forward_error( + message="Test error", + index_name="test-index", + source="/path/to/file.pdf", + original_filename="file.pdf", + ) + + import json + error_dict = json.loads(str(exc)) + assert error_dict["message"] == "Test error" + assert error_dict["index_name"] == "test-index" + assert error_dict["source"] == "/path/to/file.pdf" + assert error_dict["original_filename"] == "file.pdf" + + +def test_parse_json_or_none(monkeypatch): + """Test _parse_json_or_none parses or returns None.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # Valid JSON dict + result = tasks._parse_json_or_none('{"key": "value"}') + assert result == {"key": "value"} + + # Valid JSON array (returns None) + result = tasks._parse_json_or_none('[1, 2, 3]') + assert result is None + + # Invalid JSON + result = tasks._parse_json_or_none("not json") + assert result is None + + # Empty string + result = tasks._parse_json_or_none("") + assert result is None + + +def test_global_ray_actor_pool_manager_ensure_pool(monkeypatch): + """Test GlobalRayActorPoolManager.ensure_pool logic.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + manager = tasks.GlobalRayActorPoolManager(warm_timeout_s=10.0) + assert manager.warm_timeout_s == 10.0 + assert len(manager.actors) == 0 + + # Note: _create_and_warm_actor requires a real Ray actor, + # so we just test the pool size calculation logic + result = manager.ensure_pool(desired=0, max_allowed=5) + assert result == 0 + + +def test_delete_source_file_via_http_sync(monkeypatch): + """Test _delete_source_file_via_http_sync makes correct HTTP call.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + captured = {} + + class FakeResponse: + status_code = 200 + text = '{"deleted": true}' + + def json(self): + return {"deleted": True} + + def mock_delete(url, params, headers, timeout): + captured["url"] = url + captured["params"] = params + captured["headers"] = headers + captured["timeout"] = timeout + return FakeResponse() + + monkeypatch.setattr(tasks.requests, "delete", mock_delete) + + result = tasks._delete_source_file_via_http_sync( + base_url="http://api", + index_name="test-index", + path_or_url="/path/to/file.pdf", + scope="source_only", + authorization="Bearer token123", + timeout_s=30.0, + ) + + assert result["http_status"] == 200 + assert result["response_json"] == {"deleted": True} + assert captured["url"] == "http://api/indices/test-index/documents" + assert captured["params"]["path_or_url"] == "/path/to/file.pdf" + assert captured["params"]["scope"] == "source_only" + assert captured["headers"]["Authorization"] == "Bearer token123" + + +def test_delete_source_file_via_http_sync_empty_base_url(monkeypatch): + """Test _delete_source_file_via_http_sync raises when base_url is empty.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + with pytest.raises(RuntimeError, match="not configured"): + tasks._delete_source_file_via_http_sync( + base_url="", + index_name="test-index", + path_or_url="/path/to/file.pdf", + scope="source_only", + ) + + +def test_submit_process_forward_chain(monkeypatch): + """Test submit_process_forward_chain creates correct chain.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + captured_chain = [] + + class MockChain: + def __init__(self, *steps): + self.steps = steps + + def set(self, queue=None): + return self + + def apply_async(self): + class Result: + id = "chain-id-123" + return Result() + + monkeypatch.setattr(tasks, "chain", lambda *args: MockChain(*args)) + + chain_id = tasks.submit_process_forward_chain( + source="/path/to/file.pdf", + source_type="local", + chunking_strategy="basic", + index_name="test-index", + original_filename="file.pdf", + authorization="Bearer token", + embedding_model_id=1, + tenant_id="tenant-1", + ) + + assert chain_id == "chain-id-123" + + +def test_aggregate_parts_empty_results(monkeypatch): + """Test aggregate_parts handles empty results.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + self = FakeSelf("agg-1") + result = tasks.aggregate_parts(self, parts_results=None) + assert result["chunks"] == [] + + result = tasks.aggregate_parts(self, parts_results=[]) + assert result["chunks"] == [] + + result = tasks.aggregate_parts(self, parts_results=[[], None, [{"content": "x"}]]) + assert result["chunks"] == [{"content": "x"}] + + +def test_process_sync_with_celery_context(monkeypatch, tmp_path): + """Test process_sync with Celery task context.""" + import_tasks_with_fake_ray(monkeypatch, initialized=True) + from backend.data_process import tasks + + f = tmp_path / "test.txt" + f.write_text("hello world") + + class FakeActor: + def __init__(self): + pass + + def process_file(self, *args, **kwargs): + class Ref: + pass + return Ref() + + fake_ray = sys.modules.get("ray") + fake_ray.get_returns = [{"content": "hello world", "metadata": {}}] + + monkeypatch.setattr(tasks, "get_ray_actor", lambda: FakeActor()) + + self = FakeSelf("sync-1") + result = tasks.process_sync( + self, + source=str(f), + source_type="local", + chunking_strategy="basic", + ) + + assert result["text"] == "hello world" + assert result["chunks_count"] == 1 + assert len(self.states) >= 1 + + +def test_process_and_forward_delegates_to_chain(monkeypatch): + """Test process_and_forward creates chain and returns ID.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + class MockResult: + id = "chain-456" + + captured = {} + + class MockChain: + def __init__(self, *steps): + captured["steps"] = len(steps) + + def set(self, queue=None): + return self + + def apply_async(self): + return MockResult() + + monkeypatch.setattr(tasks, "chain", lambda *args: MockChain(*args)) + + self = FakeSelf("paf-1") + result = tasks.process_and_forward( + self, + source="/path/to/file.pdf", + source_type="local", + chunking_strategy="basic", + index_name="test-index", + ) + + assert result == "chain-456" + assert captured["steps"] == 3 # process, forward, cleanup_source + + +def test_estimate_parallel_parts_edge_cases(monkeypatch): + """Test _estimate_parallel_parts handles edge cases.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + monkeypatch.setattr(tasks, "RAY_NUM_CPUS", 4) + monkeypatch.setattr(tasks, "RAY_ACTOR_NUM_CPUS", 2) + monkeypatch.setattr(tasks, "DP_PART_PROCESSOR_COUNT", 10) + + # Should respect MAX constraint + result = tasks._estimate_parallel_parts() + assert result <= 10 + + # With more reasonable settings + monkeypatch.setattr(tasks, "DP_PART_PROCESSOR_COUNT", 2) + result = tasks._estimate_parallel_parts() + assert result == 2 + + +def test_run_async_fallback_thread_executor(monkeypatch): + """Test run_async falls back to thread executor when nest_asyncio unavailable.""" + import_tasks_with_fake_ray(monkeypatch) + + class FakeLoop: + def is_running(self): + return True + + def run_until_complete(self, coro): + return "thread-result" + + async def sample_coro(): + return "async-result" + + # Remove nest_asyncio if present + if "nest_asyncio" in sys.modules: + del sys.modules["nest_asyncio"] + + monkeypatch.setattr(asyncio, "get_running_loop", lambda: FakeLoop()) + + from backend.data_process import tasks + result = tasks.run_async(sample_coro()) + assert result == "thread-result" + + +def test_extract_error_code_from_es_response(monkeypatch): + """Test _extract_error_code_from_es_response handles various ES responses.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # From parsed body with error_code + result = tasks._extract_error_code_from_es_response( + parsed_body={"error_code": "ES_ERR_001"}, + text='{"error": "some error"}', + ) + assert result == "ES_ERR_001" + + # From nested detail + result = tasks._extract_error_code_from_es_response( + parsed_body={"detail": {"error_code": "ES_ERR_002"}}, + text='{"detail": {"error_code": "ES_ERR_002"}}', + ) + assert result == "ES_ERR_002" + + # From regex in text + result = tasks._extract_error_code_from_es_response( + parsed_body={"error": "Some error"}, + text='{"error_code": "ES_ERR_003"}', + ) + assert result == "ES_ERR_003" + + # None when no error code + result = tasks._extract_error_code_from_es_response( + parsed_body={"error": "Generic error"}, + text='{"error": "no code here"}', + ) + assert result is None + + +def test_save_error_to_redis_empty_task_id(monkeypatch): + """Test save_error_to_redis handles empty task_id.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # Should not raise, just log warning + tasks.save_error_to_redis("", "Some error", 1000.0) + tasks.save_error_to_redis(None, "Some error", 1000.0) + + +def test_save_error_to_redis_empty_reason(monkeypatch): + """Test save_error_to_redis handles empty error_reason.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # Should not raise, just log warning + tasks.save_error_to_redis("task-1", "", 1000.0) + tasks.save_error_to_redis("task-1", None, 1000.0) + + +def test_distribute_chunks_round_robin(monkeypatch): + """Test _distribute_chunks_round_robin distributes evenly.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + # Create empty batches + batches = [[], [], []] + + # Distribute 10 chunks + chunks = [{"content": f"chunk_{i}"} for i in range(10)] + + tasks._distribute_chunks_round_robin( + batches=batches, + chunks=chunks, + batch_size=10, + error_context="test", + ) + + # Each batch should have some chunks + total = sum(len(b) for b in batches) + assert total == 10 + + +def test_prewarm_ray_actors(monkeypatch): + """Test prewarm_ray_actors calls pool manager.""" + import_tasks_with_fake_ray(monkeypatch) + from backend.data_process import tasks + + captured = {} + + class MockManager: + def __init__(self, warm_timeout_s): + pass + + def ensure_pool(self, desired, max_allowed): + captured["desired"] = desired + captured["max_allowed"] = max_allowed + return 3 + + monkeypatch.setattr(tasks, "_get_or_create_global_pool_manager", lambda: MockManager(60)) + monkeypatch.setattr(tasks, "_estimate_parallel_parts", lambda: 2) + + result = tasks.prewarm_ray_actors(target_size=5) + assert result == 3 + assert captured["desired"] == 5 + + +def test_get_split_actor(monkeypatch): + """Test _get_split_actor returns actor from pool.""" + import_tasks_with_fake_ray(monkeypatch, initialized=True) + from backend.data_process import tasks + + class MockActor: + pass + + class MockManager: + def get_actor(self): + return MockActor() + + captured_manager = [] + + def mock_get_manager(): + manager = MockManager() + captured_manager.append(manager) + return manager + + monkeypatch.setattr(tasks, "_get_or_create_global_pool_manager", mock_get_manager) + + actor = tasks._get_split_actor() + assert actor is MockActor + assert len(captured_manager) == 1 diff --git a/test/backend/data_process/test_worker.py b/test/backend/data_process/test_worker.py index f88fd588a3..2a1b1fc199 100644 --- a/test/backend/data_process/test_worker.py +++ b/test/backend/data_process/test_worker.py @@ -788,3 +788,429 @@ def test_worker_ready_handler_thread_schedule_failure(mocker): mocker.patch.object(worker_module.threading, "Thread", side_effect=RuntimeError("thread failed")) worker_module.worker_ready_handler() assert worker_module.worker_state["ready"] is True + + +def test_setup_worker_environment_sets_logging_level(mocker): + """Test setup_worker_environment sets celery worker strategy log level.""" + worker_module, _ = setup_mocks_for_worker(mocker, initialized=True) + + logger_capture = [] + + class FakeCeleryWorkerStrategyLogger: + def __init__(self): + pass + + def setLevel(self, level): + logger_capture.append(level) + + import logging + mocker.patch.dict(sys.modules, { + "celery.worker.strategy": types.SimpleNamespace( + Logger=FakeCeleryWorkerStrategyLogger + ) + }) + + worker_module.setup_worker_environment() + assert logging.WARNING in logger_capture + + +def test_setup_worker_environment_logs_environment_status(mocker): + """Test setup_worker_environment logs environment variables status.""" + worker_module, _ = setup_mocks_for_worker(mocker, initialized=True) + + # Mock logger to capture calls + import logging + log_calls = [] + + class MockLogger: + def debug(self, msg, *args): + log_calls.append(("debug", msg % args if args else msg)) + + def info(self, msg, *args): + log_calls.append(("info", msg % args if args else msg)) + + def error(self, msg, *args): + log_calls.append(("error", msg % args if args else msg)) + + def warning(self, msg, *args): + log_calls.append(("warning", msg % args if args else msg)) + + mocker.patch.object(worker_module.logger, "info", side_effect=lambda msg, *args: log_calls.append(("info", msg % args if args else msg))) + mocker.patch.object(worker_module.logger, "debug", side_effect=lambda msg, *args: log_calls.append(("debug", msg % args if args else msg))) + mocker.patch.object(worker_module.logger, "warning", side_effect=lambda msg, *args: log_calls.append(("warning", msg % args if args else msg))) + mocker.patch.object(worker_module.logger, "error", side_effect=lambda msg, *args: log_calls.append(("error", msg % args if args else msg))) + + worker_module.setup_worker_environment() + + # Verify that initialization happened + assert worker_module.worker_state['initialized'] is True + + +def test_validate_service_connections_handles_redis_exception(mocker): + """Test validate_service_connections handles Redis exception gracefully.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + # Mock validate_redis_connection to raise + mocker.patch.object(worker_module, "validate_redis_connection", side_effect=RuntimeError("Redis error")) + + result = worker_module.validate_service_connections() + assert result is False + + +def test_worker_state_keys_exist(): + """Test worker_state has all required keys.""" + # Test that worker module has worker_state with required structure + import backend.data_process.worker as worker_module + + assert "initialized" in worker_module.worker_state + assert "ready" in worker_module.worker_state + assert "start_time" in worker_module.worker_state + assert "process_id" in worker_module.worker_state + assert "tasks_completed" in worker_module.worker_state + assert "tasks_failed" in worker_module.worker_state + assert "environment_validated" in worker_module.worker_state + assert "services_validated" in worker_module.worker_state + + +def test_worker_ready_handler_with_process_part_queue(mocker): + """Test worker_ready_handler with process_part queue configures concurrency logger.""" + worker_module, _ = setup_mocks_for_worker(mocker) + worker_module.worker_state['start_time'] = 1000.0 + mocker.patch("backend.data_process.worker.time.time", return_value=1001.0) + mocker.patch("backend.data_process.worker.os.getpid", return_value=7) + + # Mock QUEUES to include process_part_q + if "consts.const" in sys.modules: + sys.modules["consts.const"].QUEUES = "process_part_q" + + calls = [] + + class FakeThread: + def __init__(self, target=None, daemon=None): + calls.append(target) + + def start(self): + pass + + mocker.patch.object(worker_module.threading, "Thread", FakeThread) + + # Need to reload to pick up new QUEUES value + import importlib + importlib.reload(worker_module) + + worker_module.worker_ready_handler() + # Should have started prewarm thread and potentially part concurrency thread + assert len(calls) >= 0 # Threads may or may not be started depending on queue config + + +def test_setup_worker_environment_with_env_var_check(mocker): + """Test setup_worker_environment checks sensitive environment variables.""" + worker_module, _ = setup_mocks_for_worker(mocker, initialized=True) + + debug_calls = [] + + class FakeLogger: + def debug(self, msg, *args): + debug_calls.append(msg % args if args else msg) + + def info(self, msg, *args): + pass + + def error(self, msg, *args): + pass + + def warning(self, msg, *args): + pass + + mocker.patch.object(worker_module.logger, "debug", side_effect=lambda msg, *args: debug_calls.append(msg % args if args else msg)) + + worker_module.setup_worker_environment() + + # Check that environment variable checks are logged + assert any("REDIS" in str(call) or "ELASTICSEARCH" in str(call) for call in debug_calls) + + +def test_validate_redis_connection_with_custom_timeout(mocker): + """Test validate_redis_connection uses correct socket_timeout.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + captured_timeout = [] + + class FakeRedisClient: + def ping(self): + return True + + class FakeRedis: + @staticmethod + def from_url(url, socket_timeout=5): + captured_timeout.append(socket_timeout) + return FakeRedisClient() + + fake_redis_module = types.SimpleNamespace(from_url=FakeRedis.from_url) + mocker.patch.dict(sys.modules, {"redis": fake_redis_module}) + + worker_module.validate_redis_connection() + assert captured_timeout[0] == 5 # Default timeout should be 5 + + +def test_start_worker_logs_configuration(mocker): + """Test start_worker logs Celery configuration.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + mocker.patch("backend.data_process.worker.os.getpid", return_value=12345) + + debug_calls = [] + + class FakeLogger: + def debug(self, msg, *args): + debug_calls.append(msg % args if args else msg) + + def info(self, msg, *args): + pass + + mocker.patch.object(worker_module.logger, "debug", side_effect=lambda msg, *args: debug_calls.append(msg % args if args else msg)) + + call_args = [] + + def mock_worker_main(args): + call_args.append(args) + + mocker.patch.object(worker_module.app, "worker_main", side_effect=mock_worker_main) + + worker_module.start_worker() + + # Verify configuration logging + assert any("broker_url" in str(call) or "result_backend" in str(call) for call in debug_calls) + + +def test_task_failure_handler_logs_exception_details(mocker): + """Test task_failure_handler logs exception type and message.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + error_calls = [] + + class FakeLogger: + def error(self, msg, *args): + error_calls.append(msg % args if args else msg) + + mocker.patch.object(worker_module.logger, "error", side_effect=lambda msg, *args: error_calls.append(msg % args if args else msg)) + + fake_sender = types.SimpleNamespace(name="test_task") + fake_exception = ValueError("Test error message") + + worker_module.task_failure_handler( + sender=fake_sender, + task_id="task-123", + exception=fake_exception + ) + + # Verify error logging + assert any("task-123" in str(call) for call in error_calls) + assert any("ValueError" in str(call) or "Test error" in str(call) for call in error_calls) + + +def test_task_failure_handler_updates_worker_state(mocker): + """Test task_failure_handler increments tasks_failed counter.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + initial_failed = worker_module.worker_state['tasks_failed'] + + fake_sender = types.SimpleNamespace(name="test_task") + fake_exception = RuntimeError("Task failed") + + # Mock logger to suppress output + mocker.patch.object(worker_module.logger, "error") + + worker_module.task_failure_handler( + sender=fake_sender, + task_id="task-456", + exception=fake_exception + ) + + assert worker_module.worker_state['tasks_failed'] == initial_failed + 1 + + +def test_worker_ready_handler_calculates_startup_time(mocker): + """Test worker_ready_handler calculates correct startup time.""" + worker_module, _ = setup_mocks_for_worker(mocker) + worker_module.worker_state['start_time'] = 1000.0 + + mocker.patch("backend.data_process.worker.os.getpid", return_value=12345) + mocker.patch("backend.data_process.worker.time.time", return_value=1005.5) + + # Mock logger to suppress output + mocker.patch.object(worker_module.logger, "debug") + mocker.patch.object(worker_module.logger, "info") + + worker_module.worker_ready_handler() + + assert worker_module.worker_state['ready'] is True + + +def test_start_worker_with_multiple_queues(mocker): + """Test start_worker handles multiple queue configuration.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + # Set multiple queues + if "consts.const" in sys.modules: + sys.modules["consts.const"].QUEUES = "process_q,forward_q,custom_q" + + mocker.patch("backend.data_process.worker.os.getpid", return_value=12345) + + call_args = [] + + def mock_worker_main(args): + call_args.append(args) + + mocker.patch.object(worker_module.app, "worker_main", side_effect=mock_worker_main) + + worker_module.start_worker() + + assert len(call_args) == 1 + assert "--queues=process_q,forward_q,custom_q" in call_args[0] + + +def test_worker_ready_handler_schedules_prewarm_for_process_queue(mocker): + """Test worker_ready_handler schedules Ray actor prewarm for process_q.""" + worker_module, _ = setup_mocks_for_worker(mocker) + worker_module.worker_state['start_time'] = 1000.0 + + mocker.patch("backend.data_process.worker.os.getpid", return_value=7) + mocker.patch("backend.data_process.worker.time.time", return_value=1001.0) + + # Set QUEUES to include process_q + if "consts.const" in sys.modules: + sys.modules["consts.const"].QUEUES = "process_q,forward_q" + + prewarm_calls = [] + + def mock_prewarm(target_size): + prewarm_calls.append(target_size) + return 2 + + class FakeThread: + def __init__(self, target=None, daemon=None): + self.target = target + + def start(self): + if self.target: + self.target() + + mocker.patch.object(worker_module.threading, "Thread", FakeThread) + + # Mock prewarm function + mocker.patch("backend.data_process.tasks.prewarm_ray_actors", mock_prewarm) + + worker_module.worker_ready_handler() + + # Give thread time to execute (in real scenario) + # The prewarm is started in a daemon thread + + +def test_setup_worker_process_resources_handles_monitoring_exception(mocker): + """Test setup_worker_process_resources handles monitoring initialization failure.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + mocker.patch("backend.data_process.worker.os.getpid", return_value=99999) + + # Mock validate_service_connections to succeed + mocker.patch.object(worker_module, "validate_service_connections", return_value=True) + + # Mock monitoring import to fail + import_original = __builtins__.__import__ if hasattr(__builtins__, '__import__') else builtins.__import__ + + def mock_import(name, *args, **kwargs): + if "monitoring" in name: + raise ImportError("No monitoring module") + return import_original(name, *args, **kwargs) + + mocker.patch("builtins.__import__", side_effect=mock_import) + + # Should not raise, just log warning + try: + worker_module.setup_worker_process_resources() + except ImportError: + pass # May or may not raise depending on import structure + + +def test_validate_service_connections_returns_true_on_success(mocker): + """Test validate_service_connections returns True when all checks pass.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + class FakeRedisClient: + def ping(self): + return True + + class FakeRedis: + @staticmethod + def from_url(url, socket_timeout=5): + return FakeRedisClient() + + fake_redis_module = types.SimpleNamespace(from_url=FakeRedis.from_url) + mocker.patch.dict(sys.modules, {"redis": fake_redis_module}) + + result = worker_module.validate_service_connections() + assert result is True + + +def test_worker_ready_handler_logs_worker_status(mocker): + """Test worker_ready_handler logs worker status summary.""" + worker_module, _ = setup_mocks_for_worker(mocker) + worker_module.worker_state['start_time'] = 1000.0 + + mocker.patch("backend.data_process.worker.os.getpid", return_value=12345) + mocker.patch("backend.data_process.worker.time.time", return_value=1002.0) + + debug_calls = [] + + class FakeLogger: + def debug(self, msg, *args): + debug_calls.append(msg % args if args else msg) + + def info(self, msg, *args): + pass + + mocker.patch.object(worker_module.logger, "debug", side_effect=lambda msg, *args: debug_calls.append(msg % args if args else msg)) + + worker_module.worker_ready_handler() + + # Should log status summary + assert any("status" in str(call).lower() or "summary" in str(call).lower() for call in debug_calls) + + +def test_task_postrun_handler_does_not_increment_on_failure_state(mocker): + """Test task_postrun_handler does not increment completed on FAILURE state.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + initial_completed = worker_module.worker_state['tasks_completed'] + + fake_task = types.SimpleNamespace(name="test_task") + worker_module.task_postrun_handler(task=fake_task, task_id="task-789", state="FAILURE") + + # Should not increment completed count + assert worker_module.worker_state['tasks_completed'] == initial_completed + + +def test_task_postrun_handler_does_not_increment_on_pending_state(mocker): + """Test task_postrun_handler does not increment completed on PENDING state.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + initial_completed = worker_module.worker_state['tasks_completed'] + + fake_task = types.SimpleNamespace(name="test_task") + worker_module.task_postrun_handler(task=fake_task, task_id="task-999", state="PENDING") + + # Should not increment completed count + assert worker_module.worker_state['tasks_completed'] == initial_completed + + +def test_task_postrun_handler_increments_on_success_state(mocker): + """Test task_postrun_handler increments completed on SUCCESS state.""" + worker_module, _ = setup_mocks_for_worker(mocker) + + initial_completed = worker_module.worker_state['tasks_completed'] + + fake_task = types.SimpleNamespace(name="test_task") + worker_module.task_postrun_handler(task=fake_task, task_id="task-success", state="SUCCESS") + + assert worker_module.worker_state['tasks_completed'] == initial_completed + 1 diff --git a/test/backend/test_data_process_service_entrypoint.py b/test/backend/test_data_process_service_entrypoint.py new file mode 100644 index 0000000000..f2eb74fbdc --- /dev/null +++ b/test/backend/test_data_process_service_entrypoint.py @@ -0,0 +1,269 @@ +"""Isolated unit tests for the data-process service entrypoint.""" + +import importlib.util +import signal +import sys +import types +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + + +@pytest.fixture() +def service_module(monkeypatch): + """Load the entrypoint with its runtime dependencies replaced by stubs.""" + uvicorn = types.ModuleType("uvicorn") + uvicorn.run = MagicMock() + ray = types.ModuleType("ray") + ray.is_initialized = MagicMock(return_value=False) + ray.shutdown = MagicMock() + ray.cluster_resources = MagicMock(return_value={}) + ray.get_runtime_context = MagicMock(return_value=types.SimpleNamespace(gcs_address="")) + + dotenv = types.ModuleType("dotenv") + dotenv.load_dotenv = MagicMock() + fastapi = types.ModuleType("fastapi") + + class FakeFastAPI: + def __init__(self, **kwargs): + self.kwargs = kwargs + self.routers = [] + + def include_router(self, router): + self.routers.append(router) + + fastapi.FastAPI = FakeFastAPI + ray_config = types.ModuleType("data_process.ray_config") + ray_config.RayConfig = types.SimpleNamespace(init_ray_for_service=MagicMock(return_value=True)) + logging_utils = types.ModuleType("utils.logging_utils") + logging_utils.configure_logging = MagicMock() + constants = types.ModuleType("consts.const") + constants.REDIS_URL = "redis://test:6379/0" + constants.REDIS_PORT = 6379 + constants.FLOWER_PORT = 5555 + constants.RAY_DASHBOARD_PORT = 8265 + constants.RAY_DASHBOARD_HOST = "127.0.0.1" + constants.RAY_ACTOR_NUM_CPUS = 1 + constants.RAY_NUM_CPUS = "2" + constants.DISABLE_RAY_DASHBOARD = False + constants.DISABLE_CELERY_FLOWER = False + constants.DOCKER_ENVIRONMENT = False + constants.RAY_OBJECT_STORE_MEMORY_GB = 1 + constants.RAY_preallocate_plasma = False + constants.RAY_TEMP_DIR = "/tmp/ray" + constants.DP_PART_PROCESSOR_COUNT = 2 + + monkeypatch.setitem(sys.modules, "uvicorn", uvicorn) + monkeypatch.setitem(sys.modules, "ray", ray) + monkeypatch.setitem(sys.modules, "dotenv", dotenv) + monkeypatch.setitem(sys.modules, "fastapi", fastapi) + monkeypatch.setitem(sys.modules, "data_process.ray_config", ray_config) + monkeypatch.setitem(sys.modules, "utils.logging_utils", logging_utils) + monkeypatch.setitem(sys.modules, "consts.const", constants) + + module_name = "backend.data_process_service" + spec = importlib.util.spec_from_file_location( + module_name, + Path(__file__).parents[2] / "backend" / "data_process_service.py", + ) + module = importlib.util.module_from_spec(spec) + with patch.object(signal, "signal"): + spec.loader.exec_module(module) + return module + + +def test_service_manager_merges_disable_flags(service_module, monkeypatch): + monkeypatch.setattr(service_module, "DISABLE_RAY_DASHBOARD", True) + monkeypatch.setattr(service_module, "DISABLE_CELERY_FLOWER", True) + + manager = service_module.ServiceManager( + {"disable_ray_dashboard": False, "disable_celery_flower": False} + ) + + assert manager.config["disable_ray_dashboard"] is True + assert manager.config["start_flower"] is False + assert manager.redis_port == 6379 + + +def test_start_ray_cluster_returns_when_disabled(service_module): + manager = service_module.ServiceManager({"start_ray": False}) + + assert manager.start_ray_cluster() is True + service_module.RayConfig.init_ray_for_service.assert_not_called() + + +def test_start_all_services_starts_enabled_services_in_order(service_module, monkeypatch): + scheduler = types.SimpleNamespace(start=MagicMock()) + scheduler_module = types.ModuleType("services.auto_summary_scheduler") + scheduler_module.auto_summary_scheduler = scheduler + monkeypatch.setitem(sys.modules, "services.auto_summary_scheduler", scheduler_module) + + manager = service_module.ServiceManager( + {"start_redis": True, "start_ray": True, "start_workers": False, "disable_celery_flower": True} + ) + started = [] + manager.start_redis = lambda: started.append("redis") or True + manager.start_ray_cluster = lambda: started.append("ray") or True + manager.log_service_info = MagicMock() + + assert manager.start_all_services() is True + assert started == ["redis", "ray"] + manager.log_service_info.assert_called_once() + scheduler.start.assert_called_once() + + +def test_start_all_services_reports_failure(service_module, monkeypatch): + scheduler_module = types.ModuleType("services.auto_summary_scheduler") + scheduler_module.auto_summary_scheduler = types.SimpleNamespace(start=MagicMock()) + monkeypatch.setitem(sys.modules, "services.auto_summary_scheduler", scheduler_module) + + manager = service_module.ServiceManager( + {"start_redis": True, "start_ray": False, "start_workers": False, "disable_celery_flower": True} + ) + manager.start_redis = MagicMock(return_value=False) + manager.log_service_info = MagicMock() + + assert manager.start_all_services() is False + manager.log_service_info.assert_not_called() + + +def test_stop_all_services_stops_workers_scheduler_and_redis(service_module, monkeypatch): + scheduler = types.SimpleNamespace(stop=MagicMock()) + scheduler_module = types.ModuleType("services.auto_summary_scheduler") + scheduler_module.auto_summary_scheduler = scheduler + monkeypatch.setitem(sys.modules, "services.auto_summary_scheduler", scheduler_module) + + worker = MagicMock() + worker.poll.return_value = None + redis_process = MagicMock() + service_module.service_processes.update( + {"workers": [{"process": worker, "name": "worker", "queue": "queue"}], "redis": redis_process, "flower": None} + ) + manager = service_module.ServiceManager({}) + + manager.stop_all_services() + + worker.terminate.assert_called_once() + worker.wait.assert_called_once_with(timeout=10) + redis_process.terminate.assert_called_once() + redis_process.wait.assert_called_once_with(timeout=5) + scheduler.stop.assert_called_once() + assert service_module.service_processes["workers"] == [] + assert service_module.service_processes["redis"] is None + manager.stop_all_services() + assert scheduler.stop.call_count == 1 + + +def test_create_app_registers_data_process_router(service_module, monkeypatch): + app_module = types.ModuleType("apps.data_process_app") + app_module.router = object() + monkeypatch.setitem(sys.modules, "apps.data_process_app", app_module) + + app = service_module.create_app() + + assert app.kwargs == {"root_path": "/api", "lifespan": service_module.lifespan} + assert app.routers == [app_module.router] + + +def test_check_redis_connection_handles_success_import_error_and_runtime_error(service_module, monkeypatch): + manager = service_module.ServiceManager({}) + redis_client = MagicMock() + redis_module = types.ModuleType("redis") + redis_module.from_url = MagicMock(return_value=redis_client) + redis_service = types.ModuleType("services.redis_service") + redis_service.get_redis_service = MagicMock(return_value=types.SimpleNamespace(cleanup_error_info_keys=lambda: {"removed": 1})) + monkeypatch.setitem(sys.modules, "redis", redis_module) + monkeypatch.setitem(sys.modules, "services.redis_service", redis_service) + + assert manager._check_redis_connection("redis://ignored") is True + redis_client.ping.assert_called_once() + + monkeypatch.delitem(sys.modules, "redis") + assert manager._check_redis_connection("redis://ignored") is False + + redis_module.from_url.side_effect = RuntimeError("unreachable") + monkeypatch.setitem(sys.modules, "redis", redis_module) + assert manager._check_redis_connection("redis://ignored") is False + + +def test_start_ray_cluster_tracks_new_cluster_and_ray_address(service_module, monkeypatch): + manager = service_module.ServiceManager({"start_ray": True}) + service_module.ray.is_initialized.return_value = False + service_module.ray.get_runtime_context.return_value = types.SimpleNamespace(gcs_address="ray://cluster") + + assert manager.start_ray_cluster() is True + + service_module.RayConfig.init_ray_for_service.assert_called_once_with( + num_cpus=2, + dashboard_port=8265, + try_connect_first=True, + include_dashboard=True, + ) + assert manager._ray_cluster_started is True + assert manager.config["ray_address"] == "ray://cluster" + assert service_module.service_processes["ray_cluster"] is True + monkeypatch.delenv("RAY_ADDRESS", raising=False) + + +def test_start_ray_cluster_uses_direct_fallback_when_helper_fails(service_module): + manager = service_module.ServiceManager({"start_ray": True, "disable_ray_dashboard": True}) + service_module.RayConfig.init_ray_for_service.return_value = False + service_module.ray.is_initialized.return_value = False + service_module.ray.init = MagicMock() + + assert manager.start_ray_cluster() is True + + service_module.ray.init.assert_called_once() + assert manager._ray_cluster_started is True + + +def test_parse_arguments_and_lifespan_shutdown(service_module, monkeypatch): + monkeypatch.setattr(sys, "argv", [ + "data_process_service.py", + "--no-workers", + "--no-ray", + "--disable-celery-flower", + "--disable-ray-dashboard", + "--redis-port", "6380", + "--api-port", "5013", + ]) + args = service_module.parse_arguments() + assert args.no_workers is True + assert args.no_ray is True + assert args.redis_port == 6380 + assert args.api_port == 5013 + + manager = MagicMock() + manager._shutdown_called = False + service_module.service_manager = manager + lifecycle = service_module.lifespan(object()) + + import asyncio + + async def run_lifespan(): + async with lifecycle: + pass + + asyncio.run(run_lifespan()) + manager.stop_all_services.assert_called_once() + + +def test_stop_all_services_kills_timed_out_worker_and_stops_ray(service_module, monkeypatch): + scheduler_module = types.ModuleType("services.auto_summary_scheduler") + scheduler_module.auto_summary_scheduler = types.SimpleNamespace(stop=MagicMock()) + monkeypatch.setitem(sys.modules, "services.auto_summary_scheduler", scheduler_module) + worker = MagicMock() + worker.poll.return_value = None + worker.wait.side_effect = [service_module.subprocess.TimeoutExpired("worker", 10), None] + service_module.ray.is_initialized.return_value = True + service_module.service_processes.update({"workers": [{"process": worker, "name": "worker", "queue": "queue"}], "redis": None, "flower": None}) + manager = service_module.ServiceManager({}) + manager._ray_cluster_started = True + monkeypatch.setattr(service_module.time, "sleep", lambda seconds: None) + + manager.stop_all_services() + + worker.kill.assert_called_once() + service_module.ray.shutdown.assert_called_once() + assert manager._ray_cluster_started is False diff --git a/test/backend/test_model_consts.py b/test/backend/test_model_consts.py index 31d49e0411..a8f93036e1 100644 --- a/test/backend/test_model_consts.py +++ b/test/backend/test_model_consts.py @@ -127,3 +127,1191 @@ def test_capacity_suggestion_response_has_required_fields(): f"ModelCapacitySuggestionResponse missing W11 fields: {missing}" ) + +def test_user_sign_up_request_validation(): + """Test UserSignUpRequest validation rules""" + # Valid signup request + req = model_consts.UserSignUpRequest( + email="test@example.com", + password="password123", + invite_code="INVITE123" + ) + assert req.email == "test@example.com" + assert req.password == "password123" + assert req.invite_code == "INVITE123" + assert req.auto_login is True + + # Invite code is stripped of whitespace + req = model_consts.UserSignUpRequest( + email="test@example.com", + password="password123", + invite_code=" CODE123 " + ) + assert req.invite_code == "CODE123" + + # Empty invite code raises + with pytest.raises(ValidationError): + model_consts.UserSignUpRequest( + email="test@example.com", + password="password123", + invite_code=" " + ) + + # Password min length validation + with pytest.raises(ValidationError): + model_consts.UserSignUpRequest( + email="test@example.com", + password="short", + invite_code="CODE123" + ) + + +def test_update_password_request(): + """Test UpdatePasswordRequest validation""" + req = model_consts.UpdatePasswordRequest( + old_password="oldpass123", + new_password="newpassword123" + ) + assert req.old_password == "oldpass123" + assert req.new_password == "newpassword123" + + # New password too short + with pytest.raises(ValidationError): + model_consts.UpdatePasswordRequest( + old_password="oldpass123", + new_password="short" + ) + + +def test_user_update_request(): + """Test UserUpdateRequest validation""" + # Valid with role + req = model_consts.UserUpdateRequest( + username="newname", + email="new@example.com", + role="ADMIN" + ) + assert req.username == "newname" + assert req.role == "ADMIN" + + # Invalid role pattern + with pytest.raises(ValidationError): + model_consts.UserUpdateRequest(role="INVALID_ROLE") + + # Optional fields can be None + req = model_consts.UserUpdateRequest() + assert req.username is None + assert req.role is None + + +def test_oauth_complete_request(): + """Test OAuthCompleteRequest""" + req = model_consts.OAuthCompleteRequest( + password="password123", + invite_code="CODE123" + ) + assert req.password == "password123" + assert req.invite_code == "CODE123" + assert req.email is None + + # With optional email + req = model_consts.OAuthCompleteRequest( + email="user@example.com", + password="password123", + invite_code="CODE123" + ) + assert req.email == "user@example.com" + + +def test_model_config_hierarchy(): + """Test ModelConfig, AppConfig, and GlobalConfig hierarchy""" + # Build a complete config + app_config = model_consts.AppConfig( + appName="TestApp", + appDescription="Test Description", + iconType="icon", + modelEngineEnabled=True + ) + assert app_config.appName == "TestApp" + assert app_config.modelEngineEnabled is True + + # Single model config + single_model = model_consts.SingleModelConfig( + modelName="gpt-4", + displayName="GPT-4", + dimension=1536 + ) + assert single_model.modelName == "gpt-4" + assert single_model.dimension == 1536 + + # STT model config + stt_model = model_consts.STTModelConfig( + modelName="whisper", + displayName="Whisper", + modelFactory="openai", + modelAppid="app123", + accessToken="token123" + ) + assert stt_model.modelAppid == "app123" + assert stt_model.accessToken == "token123" + + # TTS model config + tts_model = model_consts.TTSModelConfig( + modelName="tts-1", + displayName="TTS-1" + ) + assert tts_model.modelName == "tts-1" + + +def test_agent_request_validation(): + """Test AgentRequest model validation""" + # Basic agent request + req = model_consts.AgentRequest( + query="What is AI?" + ) + assert req.query == "What is AI?" + assert req.conversation_id is None + assert req.enable_plan is False + assert req.enable_automation_tool is True + + # With history + history = [ + model_consts.HistoryItem(role="user", content="Hello"), + model_consts.HistoryItem(role="assistant", content="Hi there") + ] + req = model_consts.AgentRequest( + query="Follow up", + history=history, + conversation_id=123 + ) + assert len(req.history) == 2 + assert req.conversation_id == 123 + + # With minio_files + req = model_consts.AgentRequest( + query="Analyze this", + minio_files=[ + {"filename": "doc.pdf", "path": "/path/to/doc.pdf"} + ] + ) + assert len(req.minio_files) == 1 + + # With tool_params + req = model_consts.AgentRequest( + query="Search", + tool_params=model_consts.ToolParamsRequest( + agents={ + "main": model_consts.AgentToolParamsRequest( + tools={"search": {"max_results": 10}} + ) + } + ) + ) + assert "main" in req.tool_params.agents + assert req.tool_params.agents["main"].tools["search"]["max_results"] == 10 + + +def test_message_unit_and_message_request(): + """Test MessageUnit and MessageRequest""" + # MessageUnit + msg = model_consts.MessageUnit(type="text", content="Hello world") + assert msg.type == "text" + assert msg.content == "Hello world" + assert msg.tool_call_id is None + + msg_with_tool = model_consts.MessageUnit( + type="tool_call", + content="Calling search", + tool_call_id="call_123" + ) + assert msg_with_tool.tool_call_id == "call_123" + + # MessageRequest + msg_req = model_consts.MessageRequest( + conversation_id=1, + message_idx=5, + role="user", + message=[msg] + ) + assert msg_req.conversation_id == 1 + assert msg_req.message_idx == 5 + + +def test_conversation_request_response(): + """Test ConversationRequest and ConversationResponse""" + # Default title is Chinese + req = model_consts.ConversationRequest() + assert req.title == "新对话" + + # Custom title + req = model_consts.ConversationRequest(title="My Chat") + assert req.title == "My Chat" + + # Response + resp = model_consts.ConversationResponse( + code=200, + message="Success", + data={"id": 123} + ) + assert resp.code == 200 + assert resp.data["id"] == 123 + + +def test_rename_request(): + """Test RenameRequest""" + req = model_consts.RenameRequest( + conversation_id=42, + name="New Name" + ) + assert req.conversation_id == 42 + assert req.name == "New Name" + + +def test_batch_task_request(): + """Test BatchTaskRequest""" + sources = [ + {"source": "file1.pdf", "source_type": "minio"}, + {"source": "file2.pdf", "source_type": "minio"} + ] + req = model_consts.BatchTaskRequest(sources=sources) + assert len(req.sources) == 2 + + +def test_indexing_response(): + """Test IndexingResponse""" + resp = model_consts.IndexingResponse( + success=True, + message="Indexed successfully", + total_indexed=100, + total_submitted=100 + ) + assert resp.success is True + assert resp.total_indexed == 100 + + +def test_chunk_create_update_requests(): + """Test ChunkCreateRequest and ChunkUpdateRequest""" + # Create request + chunk = model_consts.ChunkCreateRequest( + content="This is chunk content", + title="Chunk 1", + filename="doc.pdf", + path_or_url="/path/to/doc", + metadata={"source": "manual"} + ) + assert chunk.content == "This is chunk content" + assert chunk.metadata["source"] == "manual" + + # Update request + update = model_consts.ChunkUpdateRequest( + content="Updated content", + title="Updated Title" + ) + assert update.content == "Updated content" + assert update.title == "Updated Title" + + +def test_hybrid_search_request(): + """Test HybridSearchRequest validation""" + # Valid request + req = model_consts.HybridSearchRequest( + query="search term", + index_names=["index1", "index2"], + top_k=50, + weight_accurate=0.7 + ) + assert req.query == "search term" + assert len(req.index_names) == 2 + assert req.top_k == 50 + assert req.weight_accurate == 0.7 + + # Empty index_names raises + with pytest.raises(ValidationError): + model_consts.HybridSearchRequest( + query="search", + index_names=[] + ) + + # top_k out of range + with pytest.raises(ValidationError): + model_consts.HybridSearchRequest( + query="search", + index_names=["index1"], + top_k=200 + ) + + # weight_accurate out of range + with pytest.raises(ValidationError): + model_consts.HybridSearchRequest( + query="search", + index_names=["index1"], + weight_accurate=1.5 + ) + + +def test_convert_state_request(): + """Test ConvertStateRequest""" + req = model_consts.ConvertStateRequest( + process_state="SUCCESS", + forward_state="PENDING" + ) + assert req.process_state == "SUCCESS" + assert req.forward_state == "PENDING" + + # Empty values allowed + req = model_consts.ConvertStateRequest() + assert req.process_state == "" + assert req.forward_state == "" + + +def test_process_params(): + """Test ProcessParams model""" + params = model_consts.ProcessParams( + chunking_strategy="basic", + source_type="minio", + index_name="test-index", + authorization="Bearer token123" + ) + assert params.chunking_strategy == "basic" + assert params.source_type == "minio" + assert params.authorization == "Bearer token123" + + +def test_generate_prompt_request(): + """Test GeneratePromptRequest""" + req = model_consts.GeneratePromptRequest( + task_description="Create an agent", + agent_id=1, + model_id=2, + prompt_template_id=3, + tool_ids=[10, 20], + sub_agent_ids=[100] + ) + assert req.task_description == "Create an agent" + assert req.agent_id == 1 + assert len(req.tool_ids) == 2 + assert req.has_selected_resources is True + + +def test_optimize_prompt_section_request(): + """Test OptimizePromptSectionRequest""" + req = model_consts.OptimizePromptSectionRequest( + task_description="Improve prompt", + agent_id=1, + model_id=2, + section_type="duty", + section_title="System Prompt", + current_content="Old content", + feedback="Add more details", + mode="insert", + start_pos=10, + end_pos=20 + ) + assert req.section_type == "duty" + assert req.start_pos == 10 + assert req.end_pos == 20 + + +def test_bad_case_and_optimize_requests(): + """Test BadCaseItem and OptimizePromptBadCaseRequest""" + item = model_consts.BadCaseItem( + question="What is 2+2?", + answer="5", + label="wrong", + reason="Incorrect answer" + ) + assert item.label == "wrong" + + req = model_consts.OptimizePromptBadCaseRequest( + agent_id=1, + model_id=2, + current_content="Current prompt", + bad_cases=[item], + section_type="duty", + section_title="Math" + ) + assert len(req.bad_cases) == 1 + + +def test_generate_title_request(): + """Test GenerateTitleRequest""" + req = model_consts.GenerateTitleRequest( + conversation_id=42, + question="How do I learn Python?" + ) + assert req.conversation_id == 42 + assert "Python" in req.question + + +def test_agent_info_request(): + """Test AgentInfoRequest with various fields""" + req = model_consts.AgentInfoRequest( + agent_id=1, + name="test_agent", + display_name="Test Agent", + description="A test agent", + max_steps=10, + is_main_agent=True, + enabled=True, + version_no=1 + ) + assert req.agent_id == 1 + assert req.max_steps == 10 + assert req.is_main_agent is True + + +def test_tool_instance_requests(): + """Test ToolInstanceInfoRequest and ToolInstanceSearchRequest""" + tool_inst = model_consts.ToolInstanceInfoRequest( + tool_id=1, + agent_id=2, + params={"max_results": 5}, + enabled=True + ) + assert tool_inst.tool_id == 1 + assert tool_inst.params["max_results"] == 5 + + tool_search = model_consts.ToolInstanceSearchRequest( + tool_id=1, + agent_id=2 + ) + assert tool_search.tool_id == 1 + + +def test_skill_instance_info_request(): + """Test SkillInstanceInfoRequest""" + req = model_consts.SkillInstanceInfoRequest( + skill_id=1, + agent_id=2, + enabled=True, + config_values={"param1": "value1"} + ) + assert req.skill_id == 1 + assert req.config_values["param1"] == "value1" + + +def test_tool_source_enum(): + """Test ToolSourceEnum""" + assert model_consts.ToolSourceEnum.LOCAL.value == "local" + assert model_consts.ToolSourceEnum.MCP.value == "mcp" + assert model_consts.ToolSourceEnum.LANGCHAIN.value == "langchain" + assert model_consts.ToolSourceEnum.BUILTIN.value == "builtin" + + +def test_tool_info(): + """Test ToolInfo model""" + info = model_consts.ToolInfo( + name="search", + description="Search the web", + params=[{"name": "query", "type": "string"}], + source="local", + inputs="query: string", + output_type="text", + class_name="SearchTool", + usage="search(query)", + category="web", + labels=["search", "web"] + ) + assert info.name == "search" + assert info.category == "web" + assert len(info.labels) == 2 + + +def test_change_summary_request(): + """Test ChangeSummaryRequest""" + req = model_consts.ChangeSummaryRequest( + summary_result="The file is about AI" + ) + assert "AI" in req.summary_result + + +def test_message_id_request(): + """Test MessageIdRequest""" + req = model_consts.MessageIdRequest( + conversation_id=42, + message_index=5 + ) + assert req.conversation_id == 42 + assert req.message_index == 5 + + +def test_pagination_request(): + """Test PaginationRequest validation""" + req = model_consts.PaginationRequest(page=1, page_size=20) + assert req.page == 1 + assert req.page_size == 20 + + # Page must be >= 1 + with pytest.raises(ValidationError): + model_consts.PaginationRequest(page=0) + + # Page_size must be <= 100 + with pytest.raises(ValidationError): + model_consts.PaginationRequest(page=1, page_size=150) + + +def test_mcp_source_type_enum(): + """Test MCPSourceType enum""" + assert model_consts.MCPSourceType.LOCAL.value == "local" + assert model_consts.MCPSourceType.MCP_REGISTRY.value == "mcp_registry" + assert model_consts.MCPSourceType.COMMUNITY.value == "community" + + +def test_add_mcp_service_request(): + """Test AddMcpServiceRequest with strip validators""" + req = model_consts.AddMcpServiceRequest( + name=" my-mcp ", + server_url=" https://mcp.example.com ", + description=" A test MCP " + ) + assert req.name == "my-mcp" + assert req.server_url == "https://mcp.example.com" + assert req.description == "A test MCP" + + +def test_update_mcp_service_request(): + """Test UpdateMcpServiceRequest""" + req = model_consts.UpdateMcpServiceRequest( + mcp_id=1, + name="updated-mcp", + tags=["tag1", "tag2"], + group_ids="1,2,3" + ) + assert req.mcp_id == 1 + assert len(req.tags) == 2 + + +def test_community_list_request(): + """Test CommunityListRequest""" + req = model_consts.CommunityListRequest( + search=" search term ", + tag=" ai ", + transport_type="url", + limit=50 + ) + assert req.search == "search term" + assert req.tag == "ai" + assert req.limit == 50 + + +def test_registry_list_query(): + """Test RegistryListQuery with strip validators""" + req = model_consts.RegistryListQuery( + search=" test ", + version=" v1 ", + limit=25 + ) + assert req.search == "test" + assert req.version == "v1" + assert req.limit == 25 + assert req.include_deleted is False + + +def test_capacity_bare_model(): + """Test CapacityCoverageBareModel""" + model = model_consts.CapacityCoverageBareModel( + model_id=1, + model_name="gpt-4", + model_type="llm", + max_tokens=8192, + suggestion_available=True + ) + assert model.model_id == 1 + assert model.suggestion_available is True + + +def test_capacity_coverage_response(): + """Test CapacityCoverageResponse""" + resp = model_consts.CapacityCoverageResponse( + total_llm_vlm=10, + bare_count=3, + bare_models=[ + model_consts.CapacityCoverageBareModel( + model_id=1, + model_name="model1", + model_type="llm" + ) + ] + ) + assert resp.total_llm_vlm == 10 + assert len(resp.bare_models) == 1 + + +def test_memory_agent_share_mode(): + """Test MemoryAgentShareMode enum""" + assert model_consts.MemoryAgentShareMode.default() == model_consts.MemoryAgentShareMode.NEVER + assert model_consts.MemoryAgentShareMode.ALWAYS.value == "always" + assert model_consts.MemoryAgentShareMode.ASK.value == "ask" + + +def test_tenant_management_requests(): + """Test TenantCreateRequest and TenantUpdateRequest""" + create_req = model_consts.TenantCreateRequest( + tenant_name="New Tenant", + skill_ids=[1, 2, 3], + locale="zh" + ) + assert create_req.tenant_name == "New Tenant" + assert len(create_req.skill_ids) == 3 + assert create_req.locale == "zh" + + update_req = model_consts.TenantUpdateRequest( + tenant_name="Updated Tenant" + ) + assert update_req.tenant_name == "Updated Tenant" + + +def test_group_management_requests(): + """Test GroupCreateRequest, GroupUpdateRequest, GroupListRequest""" + create_req = model_consts.GroupCreateRequest( + tenant_id="tenant-1", + group_name="Admins", + group_description="Admin group" + ) + assert create_req.tenant_id == "tenant-1" + + update_req = model_consts.GroupUpdateRequest( + group_name="Super Admins" + ) + assert update_req.group_name == "Super Admins" + + list_req = model_consts.GroupListRequest( + tenant_id="tenant-1", + page=1, + page_size=50, + sort_by="created_at", + sort_order="asc" + ) + assert list_req.sort_order == "asc" + + +def test_user_list_request(): + """Test UserListRequest""" + req = model_consts.UserListRequest( + tenant_id="tenant-1", + page=2, + page_size=25 + ) + assert req.page == 2 + assert req.page_size == 25 + + +def test_invitation_requests(): + """Test invitation-related request models""" + create_req = model_consts.InvitationCreateRequest( + tenant_id="tenant-1", + code_type="ADMIN_INVITE", + capacity=5, + expiry_date="2025-12-31" + ) + assert create_req.code_type == "ADMIN_INVITE" + assert create_req.capacity == 5 + + update_req = model_consts.InvitationUpdateRequest( + capacity=10, + expiry_date="2026-06-30" + ) + assert update_req.capacity == 10 + + +def test_version_management_requests(): + """Test version management request/response models""" + publish_req = model_consts.VersionPublishRequest( + version_name="v1.0.0", + release_note="Initial release", + publish_as_a2a=True + ) + assert publish_req.version_name == "v1.0.0" + assert publish_req.publish_as_a2a is True + + rollback_req = model_consts.VersionRollbackRequest( + version_name="Rollback v1", + release_note="Rolling back" + ) + assert rollback_req.version_name == "Rollback v1" + + compare_req = model_consts.VersionCompareRequest( + version_no_a=1, + version_no_b=2 + ) + assert compare_req.version_no_a == 1 + + +def test_version_list_item_response(): + """Test VersionListItemResponse""" + resp = model_consts.VersionListItemResponse( + id=1, + version_no=1, + version_name="v1.0", + status="RELEASED", + is_a2a=False, + created_by="admin", + create_time="2025-01-01" + ) + assert resp.status == "RELEASED" + assert resp.created_by == "admin" + + +def test_current_version_response(): + """Test CurrentVersionResponse""" + resp = model_consts.CurrentVersionResponse( + version_no=5, + version_name="v1.5", + status="RELEASED", + source_type="NORMAL", + created_by="admin" + ) + assert resp.version_no == 5 + + +def test_skill_management_requests(): + """Test SkillCreateRequest and SkillUpdateRequest""" + create_req = model_consts.SkillCreateRequest( + name="my-skill", + description="A custom skill", + content="# SKILL\n\nThis is my skill", + tool_ids=[1, 2], + tags=["ai", "automation"] + ) + assert create_req.name == "my-skill" + assert len(create_req.tool_ids) == 2 + + file_data = model_consts.SkillFileData( + path="scripts/helper.py", + content="def help(): pass" + ) + update_req = model_consts.SkillUpdateRequest( + name="updated-skill", + files=[file_data] + ) + assert update_req.name == "updated-skill" + assert len(update_req.files) == 1 + + +def test_skill_repository_requests(): + """Test skill repository request models""" + install_req = model_consts.SkillRepositoryInstallRequest( + target_name="my-installed-skill" + ) + assert install_req.target_name == "my-installed-skill" + + listing_req = model_consts.SkillRepositoryListingDetailResponse( + skill_repository_id=1, + name="shared-skill", + status="approved", + tags=["featured"], + tool_ids=[1, 2] + ) + assert listing_req.status == "approved" + + +def test_manage_tenant_model_requests(): + """Test manage tenant model request models""" + list_req = model_consts.ManageTenantModelListRequest( + tenant_id="tenant-1", + model_type="llm", + page=1, + page_size=20 + ) + assert list_req.tenant_id == "tenant-1" + assert list_req.model_type == "llm" + + health_req = model_consts.ManageTenantModelHealthcheckRequest( + tenant_id="tenant-1", + display_name="GPT-4", + model_type="llm" + ) + assert health_req.tenant_id == "tenant-1" + + delete_req = model_consts.ManageTenantModelDeleteRequest( + tenant_id="tenant-1", + display_name="Old Model" + ) + assert delete_req.display_name == "Old Model" + + +def test_batch_create_models_request(): + """Test BatchCreateModelsRequest""" + req = model_consts.BatchCreateModelsRequest( + api_key="key123", + models=[{"name": "model1"}, {"name": "model2"}], + provider="openai", + type="llm" + ) + assert len(req.models) == 2 + + +def test_provider_model_requests(): + """Test provider model request models""" + list_req = model_consts.ManageProviderModelListRequest( + tenant_id="tenant-1", + provider="silicon", + model_type="llm" + ) + assert list_req.provider == "silicon" + + create_req = model_consts.ManageProviderModelCreateRequest( + tenant_id="tenant-1", + provider="openai", + model_type="llm", + api_key="key123", + base_url="https://api.openai.com" + ) + assert create_req.base_url == "https://api.openai.com" + + +def test_nl2_agent_skill_requests(): + """Test NL2AgentRunRequest and NL2SkillRunRequest""" + nl2_agent = model_consts.NL2AgentRunRequest( + query="Create a chatbot", + history=[], + minio_files=[] + ) + assert nl2_agent.query == "Create a chatbot" + + nl2_skill = model_consts.NL2SkillRunRequest( + query="Build an automation", + complexity="simple", + language="en" + ) + assert nl2_skill.complexity == "simple" + assert nl2_skill.language == "en" + + +def test_export_import_requests(): + """Test export and import request models""" + agent_info = model_consts.ExportAndImportAgentInfo( + agent_id=1, + tenant_id="tenant-1", + name="exported-agent", + display_name="Exported Agent", + description="An exported agent", + business_description="Business desc", + max_steps=10, + is_main_agent=True, + provide_run_summary=True, + enabled=True, + tools=[], + managed_agents=[] + ) + assert agent_info.agent_id == 1 + + mcp_info = model_consts.MCPInfo( + mcp_server_name="test-mcp", + mcp_url="https://mcp.test.com" + ) + assert mcp_info.mcp_server_name == "test-mcp" + + +def test_agent_repository_snapshot(): + """Test AgentRepositorySnapshot""" + snapshot = model_consts.AgentRepositorySnapshot( + agent_id=1, + agent_info={}, + mcp_info=[], + skills=[ + model_consts.SkillZipEntry( + skill_name="my-skill", + skill_zip_base64="base64data==" + ) + ] + ) + assert len(snapshot.skills) == 1 + + +def test_repository_import_requests(): + """Test repository import request models""" + req = model_consts.RepositoryImportPrecheckResponse( + agent_repository_id=1, + display_name="Test Repo", + total_count=5, + available_count=3, + percent=60, + has_abnormal=True, + items=[ + model_consts.RepositoryImportRequirementItem( + type="model", + key="gpt-4", + name="GPT-4", + available=True + ) + ] + ) + assert req.has_abnormal is True + assert len(req.items) == 1 + + +def test_agent_name_batch_requests(): + """Test agent name batch request models""" + batch_regen = model_consts.AgentNameBatchRegenerateRequest( + items=[ + model_consts.AgentNameBatchRegenerateItem( + name="old-name", + display_name="Old Name", + task_description="Redo the name" + ) + ] + ) + assert len(batch_regen.items) == 1 + + batch_check = model_consts.AgentNameBatchCheckRequest( + items=[ + model_consts.AgentNameBatchCheckItem( + name="check-name", + agent_id=1 + ) + ] + ) + assert len(batch_check.items) == 1 + + +def test_nl2_skill_run_with_complexity(): + """Test NL2SkillRunRequest with different complexity modes""" + req_simple = model_consts.NL2SkillRunRequest( + query="Simple task", + complexity="simple" + ) + assert req_simple.complexity == "simple" + + req_complicated = model_consts.NL2SkillRunRequest( + query="Complex task", + complexity="complicated" + ) + assert req_complicated.complexity == "complicated" + + +def test_model_api_config(): + """Test ModelApiConfig""" + config = model_consts.ModelApiConfig( + apiKey="secret-key", + modelUrl="https://api.example.com" + ) + assert config.apiKey == "secret-key" + + +def test_add_container_mcp_service_request(): + """Test AddContainerMcpServiceRequest""" + mcp_config = model_consts.MCPConfigRequest( + mcpServers={ + "server1": model_consts.MCPServerConfig( + command="npx", + args=["-y", "server1"], + port=5020 + ) + } + ) + req = model_consts.AddContainerMcpServiceRequest( + name="container-mcp", + description="A container MCP", + port=5020, + mcp_config=mcp_config + ) + assert req.port == 5020 + + +def test_list_mcp_services_query(): + """Test ListMcpServicesQuery with strip validation""" + req = model_consts.ListMcpServicesQuery( + tag=" ai " + ) + assert req.tag == "ai" + + +def test_port_conflict_check_request(): + """Test PortConflictCheckRequest""" + req = model_consts.PortConflictCheckRequest( + port=8080 + ) + assert req.port == 8080 + + # Invalid port range + with pytest.raises(ValidationError): + model_consts.PortConflictCheckRequest(port=0) + + with pytest.raises(ValidationError): + model_consts.PortConflictCheckRequest(port=70000) + + +def test_community_review_requests(): + """Test community review request models""" + review_list = model_consts.CommunityReviewListRequest( + status=" pending " + ) + assert review_list.status == "pending" + + review_action = model_consts.CommunityReviewActionRequest( + review_id=1, + content=" Approved " + ) + assert review_action.content == "Approved" + + +def test_group_members_update_request(): + """Test GroupMembersUpdateRequest""" + req = model_consts.GroupMembersUpdateRequest( + user_ids=["user1", "user2", "user3"] + ) + assert len(req.user_ids) == 3 + + +def test_set_default_group_request(): + """Test SetDefaultGroupRequest""" + req = model_consts.SetDefaultGroupRequest( + default_group_id=5 + ) + assert req.default_group_id == 5 + + # Invalid group_id + with pytest.raises(ValidationError): + model_consts.SetDefaultGroupRequest(default_group_id=0) + + +def test_voice_connectivity_models(): + """Test VoiceConnectivityRequest and VoiceConnectivityResponse""" + req = model_consts.VoiceConnectivityRequest( + model_type="stt" + ) + assert req.model_type == "stt" + + resp = model_consts.VoiceConnectivityResponse( + connected=True, + model_type="tts", + message="Service available" + ) + assert resp.connected is True + + +def test_tool_validate_request(): + """Test ToolValidateRequest""" + req = model_consts.ToolValidateRequest( + name="search", + source="local", + usage="search(query)", + inputs={"query": {"type": "string"}}, + params={"max_results": 10} + ) + assert req.name == "search" + assert req.params["max_results"] == 10 + + +def test_update_knowledge_list_request(): + """Test UpdateKnowledgeListRequest""" + req = model_consts.UpdateKnowledgeListRequest( + nexent=["index1", "index2"], + datamate=["dm-index1"] + ) + assert len(req.nexent) == 2 + assert req.datamate == ["dm-index1"] + + +def test_mcp_update_request(): + """Test MCPUpdateRequest""" + req = model_consts.MCPUpdateRequest( + current_service_name="old-name", + current_mcp_url="https://old.example.com", + new_service_name="new-name", + new_mcp_url="https://new.example.com", + new_authorization_token="Bearer new-token" + ) + assert req.new_service_name == "new-name" + + +def test_invitation_list_request(): + """Test InvitationListRequest""" + req = model_consts.InvitationListRequest( + tenant_id="tenant-1", + page=2, + page_size=50 + ) + assert req.page == 2 + assert req.page_size == 50 + + +def test_invitation_response(): + """Test InvitationResponse""" + resp = model_consts.InvitationResponse( + invitation_id=1, + invitation_code="INV123", + code_type="ADMIN_INVITE", + group_ids=[1, 2], + capacity=10, + status="active", + created_at="2025-01-01" + ) + assert resp.status == "active" + assert len(resp.group_ids) == 2 + + +def test_manage_tenant_model_list_response(): + """Test ManageTenantModelListResponse""" + resp = model_consts.ManageTenantModelListResponse( + tenant_id="tenant-1", + tenant_name="Test Tenant", + models=[{"name": "model1"}, {"name": "model2"}], + total=2, + page=1, + page_size=20, + total_pages=1 + ) + assert resp.total == 2 + assert resp.total_pages == 1 + + +def test_agent_repository_listing_requests(): + """Test AgentRepositoryListingCreateRequest""" + req = model_consts.AgentRepositoryListingCreateRequest( + icon="🚀", + downloads=100, + tags=["ai", "automation"], + tool_count=10, + content="This is a great agent" + ) + assert req.icon == "🚀" + assert req.downloads == 100 + + +def test_agent_repository_listing_detail_response(): + """Test AgentRepositoryListingDetailResponse""" + resp = model_consts.AgentRepositoryListingDetailResponse( + agent_repository_id=1, + name="great-agent", + status="approved", + downloads=500, + tools=["search", "write"] + ) + assert resp.downloads == 500 + assert len(resp.tools) == 2 + + +def test_community_publish_update_requests(): + """Test CommunityPublishRequest and CommunityUpdateRequest""" + publish = model_consts.CommunityPublishRequest( + mcp_id=1, + name="new-mcp", + tags=["featured"] + ) + assert publish.mcp_id == 1 + + update = model_consts.CommunityUpdateRequest( + market_id=1, + description="Updated description" + ) + assert update.description == "Updated description" + + +def test_community_status_update_request(): + """Test CommunityStatusUpdateRequest""" + req = model_consts.CommunityStatusUpdateRequest( + status="shared", + content="Approved for community" + ) + assert req.status == "shared" + + +def test_delete_mcp_service_request(): + """Test DeleteMcpServiceRequest""" + req = model_consts.DeleteMcpServiceRequest( + mcp_id=42 + ) + assert req.mcp_id == 42 + diff --git a/test/backend/utils/test_file_management_utils.py b/test/backend/utils/test_file_management_utils.py index ce15596a09..5a580ca08f 100644 --- a/test/backend/utils/test_file_management_utils.py +++ b/test/backend/utils/test_file_management_utils.py @@ -121,6 +121,24 @@ async def read(self) -> bytes: assert ok is False +@pytest.mark.asyncio +async def test_trigger_data_process_continues_when_tenant_lookup_fails(fmu, monkeypatch): + fake_client = _FakeAsyncClient(_Resp(201, {"task_id": "t1"})) + fake_httpx = types.SimpleNamespace(AsyncClient=lambda: fake_client, RequestError=_FakeRequestError) + monkeypatch.setattr(fmu, "httpx", fake_httpx) + monkeypatch.setattr(fmu, "get_current_user_id", lambda authorization: (_ for _ in ()).throw(ValueError("bad token"))) + monkeypatch.setattr(fmu, "inject_trace_context", lambda: {"traceparent": "00-test"}) + + result = await fmu.trigger_data_process( + [{"path_or_url": "/data/a.txt", "filename": "a.txt"}], + _ProcessParams("tok", "local", "basic", "idx"), + ) + + assert result == {"task_id": "t1"} + assert fake_client.last_post["headers"]["traceparent"] == "00-test" + assert fake_client.last_post["json"]["tenant_id"] is None + + # -------------------- trigger_data_process -------------------- diff --git a/test/backend/utils/test_knowledge_telemetry.py b/test/backend/utils/test_knowledge_telemetry.py index 0b8cfe61ee..cf7581f5fe 100644 --- a/test/backend/utils/test_knowledge_telemetry.py +++ b/test/backend/utils/test_knowledge_telemetry.py @@ -1,3 +1,4 @@ +import types from unittest.mock import MagicMock, patch import pytest @@ -32,32 +33,199 @@ def test_safe_attributes_maps_part_diagnostics(): } -def test_knowledge_span_marks_celery_retry_without_error(): - Retry = type("Retry", (Exception,), {"__module__": "celery.exceptions"}) +def _install_otel_stubs(monkeypatch): span = MagicMock() span_cm = MagicMock() span_cm.__enter__.return_value = span + tracer = MagicMock() + tracer.start_as_current_span.return_value = span_cm + trace = MagicMock() + trace.get_tracer.return_value = tracer + propagate = MagicMock() + context = MagicMock() + status_code = MagicMock() + status_code.OK = "OK" + status_code.ERROR = "ERROR" + status = MagicMock(side_effect=lambda code, description=None: (code, description)) + metrics = MagicMock() + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", True) + monkeypatch.setattr(knowledge_telemetry, "trace", trace, raising=False) + monkeypatch.setattr(knowledge_telemetry, "propagate", propagate, raising=False) + monkeypatch.setattr(knowledge_telemetry, "otel_context", context, raising=False) + monkeypatch.setattr(knowledge_telemetry, "Status", status, raising=False) + monkeypatch.setattr(knowledge_telemetry, "StatusCode", status_code, raising=False) + monkeypatch.setattr(knowledge_telemetry, "metrics", metrics, raising=False) + monkeypatch.setattr(knowledge_telemetry, "_resource_snapshot", lambda: {}) + monkeypatch.setattr(knowledge_telemetry, "_record_metrics", MagicMock()) + return span, span_cm, trace, propagate, context, status_code + + +def test_knowledge_span_marks_celery_retry_without_error(monkeypatch): + Retry = type("Retry", (Exception,), {"__module__": "celery.exceptions"}) + span, _, _, _, _, status_code = _install_otel_stubs(monkeypatch) - with ( - patch.object(knowledge_telemetry, "OTEL_AVAILABLE", True), - patch.object(knowledge_telemetry.trace, "get_tracer") as get_tracer, - patch.object(knowledge_telemetry, "_resource_snapshot", return_value={}), - patch.object(knowledge_telemetry, "_record_metrics"), - ): - get_tracer.return_value.start_as_current_span.return_value = span_cm - with pytest.raises(Retry): - with knowledge_telemetry.knowledge_span( - "knowledge.forward.redis_read", - "forward.redis_read", - retry_attempt=2, - retry_delay_seconds=5, - ): - raise Retry() + with pytest.raises(Retry): + with knowledge_telemetry.knowledge_span( + "knowledge.forward.redis_read", + "forward.redis_read", + retry_attempt=2, + retry_delay_seconds=5, + ): + raise Retry() span.record_exception.assert_not_called() span.set_attribute.assert_any_call("ingestion.status", "retry") span.set_attribute.assert_any_call("retry.attempt", 2) span.set_attribute.assert_any_call("retry.delay_seconds", 5.0) - span.set_status.assert_called_with( - knowledge_telemetry.Status(knowledge_telemetry.StatusCode.OK) + span.set_status.assert_called_with((status_code.OK, None)) + + +def test_safe_attributes_hashes_ids_scales_bytes_and_ignores_invalid_sizes(): + attrs = knowledge_telemetry._safe_attributes({ + "task_id": "task-1", + "tenant_id": "tenant-1", + "index_name": "knowledge-base-1", + "file_size_bytes": 2 * knowledge_telemetry.BYTES_PER_MB, + "original_filename": "REPORT.PDF", + "chunk_count": "invalid", + }) + + assert attrs["task.id"] == "task-1" + assert attrs["tenant.id_hash"] == knowledge_telemetry._safe_hash("tenant-1") + assert attrs["knowledge_base.id"] == knowledge_telemetry._safe_hash("knowledge-base-1") + assert attrs["file.size_mb"] == 2.0 + assert attrs["file.extension"] == ".pdf" + assert attrs["chunk.count"] == "invalid" + + +def test_inject_trace_context_degrades_when_propagation_fails(monkeypatch): + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", True) + propagate = MagicMock() + propagate.inject.side_effect = RuntimeError("collector unavailable") + monkeypatch.setattr(knowledge_telemetry, "propagate", propagate, raising=False) + + assert knowledge_telemetry.inject_trace_context() == {} + + +def test_knowledge_span_records_regular_exceptions(monkeypatch): + span, _, _, _, _, status_code = _install_otel_stubs(monkeypatch) + + with pytest.raises(ValueError, match="invalid"): + with knowledge_telemetry.knowledge_span("knowledge.process", "process"): + raise ValueError("invalid") + + span.record_exception.assert_called_once() + span.set_status.assert_any_call((status_code.ERROR, "ValueError")) + span.set_attribute.assert_any_call("error.type", "ValueError") + + +def test_trace_knowledge_operation_supports_sync_async_and_task_context(monkeypatch): + captured = [] + + class _Span: + def __enter__(self): + return None + + def __exit__(self, *args): + return False + + def fake_span(name, stage, **attributes): + captured.append((name, stage, attributes)) + return _Span() + + monkeypatch.setattr(knowledge_telemetry, "knowledge_span", fake_span) + + class _Task: + request = type("Request", (), {"id": "task-7", "retries": 1})() + + @knowledge_telemetry.trace_knowledge_operation("knowledge.process", "process") + def run(self, params): + return params["value"] + + @knowledge_telemetry.trace_knowledge_operation("knowledge.forward", "forward") + async def forward(params): + return params["value"] + + assert _Task().run({"value": 3, "telemetry_context": {"traceparent": "abc"}}) == 3 + assert __import__("asyncio").run(forward({"value": 4})) == 4 + assert captured[0][2]["task_id"] == "task-7" + assert captured[0][2]["retry_attempt"] == 2 + assert captured[0][2]["telemetry_context"] == {"traceparent": "abc"} + assert captured[1][2]["telemetry_context"] is None + + +def test_set_span_attributes_only_updates_recording_span(monkeypatch): + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", True) + span = MagicMock() + span.is_recording.return_value = True + trace = MagicMock() + trace.get_current_span.return_value = span + monkeypatch.setattr(knowledge_telemetry, "trace", trace, raising=False) + + knowledge_telemetry.set_span_attributes(task_id="task-1", original_filename="notes.txt") + + span.set_attributes.assert_called_once_with({"task.id": "task-1", "file.extension": ".txt"}) + + +def test_resource_snapshot_collects_process_host_and_cgroup_metrics(monkeypatch): + child = MagicMock() + child.is_running.return_value = True + child.memory_info.return_value = types.SimpleNamespace(rss=knowledge_telemetry.BYTES_PER_MB) + process = MagicMock() + process.children.return_value = [child] + process.memory_info.return_value = types.SimpleNamespace(rss=2 * knowledge_telemetry.BYTES_PER_MB) + process.cpu_percent.return_value = 12.3456 + process.num_threads.return_value = 7 + psutil = types.SimpleNamespace( + Process=lambda: process, + virtual_memory=lambda: types.SimpleNamespace(percent=42.1234, available=3 * knowledge_telemetry.BYTES_PER_MB), + cpu_percent=lambda interval: 24.5678, + ) + real_import = __import__ + + def fake_import(name, *args, **kwargs): + if name == "psutil": + return psutil + return real_import(name, *args, **kwargs) + + monkeypatch.setattr("builtins.__import__", fake_import) + monkeypatch.setattr("builtins.open", MagicMock(side_effect=[ + MagicMock(__enter__=lambda self: self, __exit__=lambda *args: None, read=lambda: str(4 * knowledge_telemetry.BYTES_PER_MB)), + MagicMock(__enter__=lambda self: self, __exit__=lambda *args: None, read=lambda: "max"), + ])) + + snapshot = knowledge_telemetry._resource_snapshot() + + assert snapshot["process.rss_memory_mb"] == 2.0 + assert snapshot["process_tree.rss_memory_mb"] == 3.0 + assert snapshot["container.used_memory_mb"] == 4.0 + assert "container.memory_limit_mb" not in snapshot + + +def test_record_metrics_records_available_resource_measurements(monkeypatch): + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", True) + histogram = MagicMock() + meter = MagicMock() + meter.create_histogram.return_value = histogram + metrics = MagicMock() + metrics.get_meter.return_value = meter + monkeypatch.setattr(knowledge_telemetry, "metrics", metrics, raising=False) + + knowledge_telemetry._record_metrics( + "forward", + 10.5, + {"process.rss_memory_mb": 3.0, "process.cpu_percent": 12.0}, ) + + assert histogram.record.call_count == 3 + + +def test_knowledge_span_degrades_when_setup_fails(monkeypatch): + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", True) + trace = MagicMock() + trace.get_tracer.side_effect = RuntimeError("tracing unavailable") + monkeypatch.setattr(knowledge_telemetry, "trace", trace, raising=False) + monkeypatch.setattr(knowledge_telemetry, "_resource_snapshot", lambda: {}) + + with knowledge_telemetry.knowledge_span("knowledge.process", "process") as span: + assert span is None diff --git a/test/backend/utils/test_monitoring.py b/test/backend/utils/test_monitoring.py index d94e20518c..10ae778932 100644 --- a/test/backend/utils/test_monitoring.py +++ b/test/backend/utils/test_monitoring.py @@ -4,8 +4,40 @@ Tests the actual functionality and integration of the OTLP monitoring system. """ -import pytest +import sys +import types from unittest.mock import MagicMock + +import pytest + +fake_consts = types.ModuleType("consts.const") +for name, value in { + "ENABLE_TELEMETRY": False, + "MONITORING_PROVIDER": "", + "MONITORING_PROJECT_NAME": "", + "OTEL_SERVICE_NAME": "test-service", + "OTEL_EXPORTER_OTLP_ENDPOINT": "http://localhost:4318", + "OTEL_EXPORTER_OTLP_TRACES_ENDPOINT": "", + "OTEL_EXPORTER_OTLP_METRICS_ENDPOINT": "", + "OTEL_EXPORTER_OTLP_PROTOCOL": "http", + "OTEL_EXPORTER_OTLP_METRICS_ENABLED": True, + "MONITORING_INSTRUMENT_REQUESTS": False, + "MONITORING_FASTAPI_INCLUDED_URLS": "", + "MONITORING_FASTAPI_EXCLUDED_URLS": "", + "MONITORING_FASTAPI_EXCLUDE_SPANS": "receive,send", + "MONITORING_TRACE_CONTENT_MODE": "summary", + "MONITORING_TRACE_MAX_CHARS": "4000", + "MONITORING_TRACE_MAX_ITEMS": "20", + "OTLP_HEADERS": {}, + "TELEMETRY_SAMPLE_RATE": 1.0, +}.items(): + setattr(fake_consts, name, value) + +fake_consts_package = types.ModuleType("consts") +fake_consts_package.const = fake_consts +sys.modules.setdefault("consts", fake_consts_package) +sys.modules.setdefault("consts.const", fake_consts) + from backend.utils.monitoring import monitoring_manager @@ -214,3 +246,41 @@ def test_get_current_span(self): def test_get_tracer(self): """Test getting tracer property.""" tracer = monitoring_manager.tracer + + def test_module_uses_backend_constants_when_top_level_import_is_unavailable(self, monkeypatch): + """The module should support imports from the backend package namespace.""" + import importlib + import backend.utils.monitoring as monitoring_module + + backend_consts = types.ModuleType("backend.consts.const") + for name, value in { + "ENABLE_TELEMETRY": False, + "MONITORING_PROVIDER": "", + "MONITORING_PROJECT_NAME": "", + "OTEL_SERVICE_NAME": "fallback-service", + "OTEL_EXPORTER_OTLP_ENDPOINT": "http://localhost:4318", + "OTEL_EXPORTER_OTLP_TRACES_ENDPOINT": "", + "OTEL_EXPORTER_OTLP_METRICS_ENDPOINT": "", + "OTEL_EXPORTER_OTLP_PROTOCOL": "http", + "OTEL_EXPORTER_OTLP_METRICS_ENABLED": True, + "MONITORING_INSTRUMENT_REQUESTS": False, + "MONITORING_FASTAPI_INCLUDED_URLS": "", + "MONITORING_FASTAPI_EXCLUDED_URLS": "", + "MONITORING_FASTAPI_EXCLUDE_SPANS": "receive,send", + "MONITORING_TRACE_CONTENT_MODE": "summary", + "MONITORING_TRACE_MAX_CHARS": "4000", + "MONITORING_TRACE_MAX_ITEMS": "20", + "OTLP_HEADERS": {}, + "TELEMETRY_SAMPLE_RATE": 1.0, + }.items(): + setattr(backend_consts, name, value) + + monkeypatch.delitem(sys.modules, "consts", raising=False) + monkeypatch.delitem(sys.modules, "consts.const", raising=False) + monkeypatch.setitem(sys.modules, "consts", types.ModuleType("consts")) + monkeypatch.setitem(sys.modules, "backend.consts.const", backend_consts) + + reloaded_module = importlib.reload(monitoring_module) + + assert reloaded_module.monitoring_manager is not None + assert reloaded_module.__all__ == ["monitoring_manager"] diff --git a/test/sdk/data_process/test_core.py b/test/sdk/data_process/test_core.py index 5576afe48e..dda22795ad 100644 --- a/test/sdk/data_process/test_core.py +++ b/test/sdk/data_process/test_core.py @@ -14,10 +14,17 @@ fake_logger.logger = types.SimpleNamespace(info=lambda *a, **k: None, warning=lambda *a, **k: None, error=lambda *a, **k: None) fake_models.tables = fake_tables fake_unstructured.models = fake_models +fake_partition = types.ModuleType("unstructured.partition") +fake_partition_auto = types.ModuleType("unstructured.partition.auto") +fake_partition_auto.partition = lambda *args, **kwargs: [] +fake_partition.auto = fake_partition_auto sys.modules.setdefault("unstructured_inference", fake_unstructured) sys.modules.setdefault("unstructured_inference.models", fake_models) sys.modules.setdefault("unstructured_inference.models.tables", fake_tables) sys.modules.setdefault("unstructured_inference.logger", fake_logger) +sys.modules.setdefault("unstructured", types.ModuleType("unstructured")) +sys.modules.setdefault("unstructured.partition", fake_partition) +sys.modules.setdefault("unstructured.partition.auto", fake_partition_auto) from sdk.nexent.data_process.core import DataProcessCore @@ -451,3 +458,12 @@ def test_file_split_splitter_exception_falls_back(self, core): assert len(parts) == 1 assert parts[0].getvalue() == data + + def test_file_split_unknown_splitter_falls_back(self, core): + """A requested splitter that is unavailable should retain the input bytes.""" + data = b"hello" + + parts = core.file_split(data, "data.txt", splitter="MissingSplitter") + + assert len(parts) == 1 + assert parts[0].getvalue() == data diff --git a/test/sdk/data_process/test_file_splitter.py b/test/sdk/data_process/test_file_splitter.py index e645bf10e0..b83ffeafaa 100644 --- a/test/sdk/data_process/test_file_splitter.py +++ b/test/sdk/data_process/test_file_splitter.py @@ -17,10 +17,17 @@ fake_logger.logger = types.SimpleNamespace(info=lambda *a, **k: None, warning=lambda *a, **k: None, error=lambda *a, **k: None) fake_models.tables = fake_tables fake_unstructured.models = fake_models +fake_partition = types.ModuleType("unstructured.partition") +fake_partition_auto = types.ModuleType("unstructured.partition.auto") +fake_partition_auto.partition = lambda *args, **kwargs: [] +fake_partition.auto = fake_partition_auto sys.modules.setdefault("unstructured_inference", fake_unstructured) sys.modules.setdefault("unstructured_inference.models", fake_models) sys.modules.setdefault("unstructured_inference.models.tables", fake_tables) sys.modules.setdefault("unstructured_inference.logger", fake_logger) +sys.modules.setdefault("unstructured", types.ModuleType("unstructured")) +sys.modules.setdefault("unstructured.partition", fake_partition) +sys.modules.setdefault("unstructured.partition.auto", fake_partition_auto) from sdk.nexent.data_process.file_splitter import FileSplitter @@ -382,3 +389,88 @@ def __exit__(self, *a): monkeypatch.setattr("sdk.nexent.data_process.file_splitter.subprocess.run", lambda *a, **k: None) with pytest.raises(RuntimeError, match="produced no output"): splitter._convert_bytes_with_libreoffice(b"doc", ".docx", ".pdf") + + +@pytest.mark.parametrize( + ("filename", "method_name"), + [ + ("book.epub", "split_epub_by_size"), + ("book.xlsx", "split_excel"), + ("items.json", "split_json_stream"), + ("notes.md", "split_markdown"), + ("report.pdf", "split_pdf_by_size"), + ("notes.txt", "split_txt_by_size"), + ("data.xml", "split_xml_by_size"), + ], +) +def test_file_process_routes_supported_extensions(monkeypatch, filename, method_name): + splitter = FileSplitter() + expected = [BytesIO(b"part")] + captured = {} + + def route(file_data, *args, **kwargs): + captured["file_data"] = file_data + captured["kwargs"] = kwargs + return expected + + monkeypatch.setattr(splitter, method_name, route) + + result = splitter.file_process(b"source", filename, max_size=8, encoding="latin-1") + + assert result == expected + assert captured["file_data"] == b"source" + + +def test_resolve_max_size_uses_target_parts_and_default(): + splitter = FileSplitter() + + assert splitter._resolve_max_size(b"123456789", target_parts=4) == 3 + assert splitter._resolve_max_size(b"x", max_size=0) == 5 * 1024 * 1024 + + +def test_copy_images_safe_handles_missing_images_and_anchor_copy_failure(monkeypatch): + splitter = FileSplitter() + added = [] + + class Source: + _images = [] + + class ImageSource: + anchor = object() + + def _data(self): + return b"image" + + class Destination: + def add_image(self, image, anchor): + added.append((image, anchor)) + + monkeypatch.setattr("openpyxl.drawing.image.Image", lambda _bio: object()) + monkeypatch.setattr("sdk.nexent.data_process.file_splitter.copy", lambda _anchor: (_ for _ in ()).throw(ValueError())) + + splitter.copy_images_safe(Source(), Destination()) + splitter.copy_images_safe(type("WithImage", (), {"_images": [ImageSource()]})(), Destination()) + + assert added[0][1] is ImageSource.anchor + + +def test_convert_bytes_with_libreoffice_wraps_command_failure(monkeypatch, tmp_path): + splitter = FileSplitter() + work = tmp_path / "w3" + work.mkdir() + + class TDir: + def __enter__(self): + return str(work) + + def __exit__(self, *args): + return False + + monkeypatch.setattr("sdk.nexent.data_process.file_splitter.tempfile.TemporaryDirectory", lambda: TDir()) + monkeypatch.setattr( + "sdk.nexent.data_process.file_splitter.subprocess.run", + lambda *args, **kwargs: (_ for _ in ()).throw(OSError("missing soffice")), + ) + + with pytest.raises(RuntimeError, match="LibreOffice conversion failed"): + splitter._convert_bytes_with_libreoffice(b"doc", ".docx", ".pdf") diff --git a/test/sdk/data_process/test_file_splitter_coverage.py b/test/sdk/data_process/test_file_splitter_coverage.py new file mode 100644 index 0000000000..6d7bff02ab --- /dev/null +++ b/test/sdk/data_process/test_file_splitter_coverage.py @@ -0,0 +1,181 @@ +from io import BytesIO +import sys +import types + +import pytest + +pytest.importorskip("ijson") +pytest.importorskip("openpyxl") +pytest.importorskip("pypdf") + +fake_unstructured = types.ModuleType("unstructured_inference") +fake_models = types.ModuleType("unstructured_inference.models") +fake_tables = types.ModuleType("unstructured_inference.models.tables") +fake_tables.tables_agent = types.SimpleNamespace(model=None) +fake_logger = types.ModuleType("unstructured_inference.logger") +fake_logger.logger = types.SimpleNamespace( + info=lambda *args, **kwargs: None, + warning=lambda *args, **kwargs: None, + error=lambda *args, **kwargs: None, +) +fake_models.tables = fake_tables +fake_unstructured.models = fake_models +fake_partition = types.ModuleType("unstructured.partition") +fake_partition_auto = types.ModuleType("unstructured.partition.auto") +fake_partition_auto.partition = lambda *args, **kwargs: [] +fake_partition.auto = fake_partition_auto +sys.modules.setdefault("unstructured_inference", fake_unstructured) +sys.modules.setdefault("unstructured_inference.models", fake_models) +sys.modules.setdefault("unstructured_inference.models.tables", fake_tables) +sys.modules.setdefault("unstructured_inference.logger", fake_logger) +sys.modules.setdefault("unstructured", types.ModuleType("unstructured")) +sys.modules.setdefault("unstructured.partition", fake_partition) +sys.modules.setdefault("unstructured.partition.auto", fake_partition_auto) + +from sdk.nexent.data_process.file_splitter import FileSplitter + + +def test_split_csv_recursively_preserves_header_and_rows(): + splitter = FileSplitter() + source = b"name,value\nalpha,111111\nbeta,222222\ngamma,333333\n" + + parts = splitter.split_csv_by_size(source, max_size=25) + + assert len(parts) == 3 + assert all(part.getvalue().startswith(b"name,value") for part in parts) + assert b"alpha,111111" in parts[0].getvalue() + assert b"gamma,333333" in parts[-1].getvalue() + + +def test_copy_images_safe_ignores_image_construction_failure(monkeypatch): + splitter = FileSplitter() + + class ImageSource: + anchor = "A1" + + def _data(self): + return b"invalid-image" + + destination = types.SimpleNamespace(add_image=lambda *args: pytest.fail("image should not be added")) + monkeypatch.setattr( + "openpyxl.drawing.image.Image", + lambda _buffer: (_ for _ in ()).throw(ValueError("invalid image")), + ) + + splitter.copy_images_safe(types.SimpleNamespace(_images=[ImageSource()]), destination) + + +def test_split_excel_skips_blank_header_sheet(monkeypatch): + splitter = FileSplitter() + + class Worksheet: + def iter_rows(self, values_only=True): + return iter([(None, None)]) + + class Workbook: + sheetnames = ["blank"] + + def __getitem__(self, _name): + return Worksheet() + + monkeypatch.setattr("openpyxl.load_workbook", lambda *args, **kwargs: Workbook()) + + assert splitter.split_excel(b"x" * 20, max_size=5) == [] + + +def test_split_markdown_without_headers_recurses_to_terminal_level(monkeypatch): + splitter = FileSplitter() + + class Document: + page_content = "plain content" + metadata = {} + + class MarkdownSplitter: + def __init__(self, headers_to_split_on): + self.headers_to_split_on = headers_to_split_on + + def split_text(self, _content): + return [Document()] + + monkeypatch.setattr("langchain_text_splitters.MarkdownHeaderTextSplitter", MarkdownSplitter) + + parts = splitter.split_markdown(b"plain content", max_size=3) + + assert [part.getvalue() for part in parts] == [b"plain content"] + + +def test_split_markdown_rebuilds_parent_header(monkeypatch): + splitter = FileSplitter() + + class Document: + def __init__(self, content, metadata): + self.page_content = content + self.metadata = metadata + + class MarkdownSplitter: + def __init__(self, headers_to_split_on): + self.level = len(headers_to_split_on[0][0]) + + def split_text(self, content): + if self.level == 2: + return [ + Document("first", {"h2": "One"}), + Document("second", {"h2": "Two"}), + ] + return [Document(content, {})] + + monkeypatch.setattr("langchain_text_splitters.MarkdownHeaderTextSplitter", MarkdownSplitter) + + parts = splitter.split_markdown(b"## One\nfirst\n## Two\nsecond", max_size=8) + + assert parts[0].getvalue().startswith(b"## One\n") + assert parts[1].getvalue().startswith(b"## Two\n") + + +def test_split_pdf_by_parts_returns_empty_for_document_without_pages(monkeypatch): + splitter = FileSplitter() + monkeypatch.setattr("pypdf.PdfReader", lambda _buffer: types.SimpleNamespace(pages=[])) + + assert splitter.split_pdf_by_parts(b"%PDF", target_parts=2) == [] + + +def test_convert_bytes_with_libreoffice_uses_discovered_output(monkeypatch, tmp_path): + splitter = FileSplitter() + output_file = tmp_path / "converted.PDF" + + class TemporaryDirectory: + def __enter__(self): + return str(tmp_path) + + def __exit__(self, *args): + return False + + def run_conversion(*args, **kwargs): + output_file.write_bytes(b"converted") + + monkeypatch.setattr( + "sdk.nexent.data_process.file_splitter.tempfile.TemporaryDirectory", + lambda: TemporaryDirectory(), + ) + monkeypatch.setattr("sdk.nexent.data_process.file_splitter.subprocess.run", run_conversion) + + result = splitter._convert_bytes_with_libreoffice(b"source", ".docx", ".pdf") + + assert result == b"converted" + + +def test_file_process_pdf_with_target_parts_uses_part_splitter(monkeypatch): + splitter = FileSplitter() + expected = [BytesIO(b"one"), BytesIO(b"two")] + captured = {} + + def split_by_parts(file_data, target_parts): + captured.update(file_data=file_data, target_parts=target_parts) + return expected + + monkeypatch.setattr(splitter, "split_pdf_by_parts", split_by_parts) + + result = splitter.file_process(b"%PDF", "report.pdf", target_parts=2) + + assert result == expected + assert captured == {"file_data": b"%PDF", "target_parts": 2} From d81e45beba3fdb62c8b1fa5fd99053ad65870d77 Mon Sep 17 00:00:00 2001 From: Jasonxia007 Date: Fri, 14 Aug 2026 12:39:34 +0800 Subject: [PATCH 3/5] =?UTF-8?q?=F0=9F=A7=AA=20Add=20test=20files?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/consts/const.py | 4 +- test/backend/data_process/test_ray_actors.py | 31 ++++++++----- test/backend/data_process/test_tasks.py | 45 +++++++------------ test/backend/data_process/test_worker.py | 40 +++++++---------- test/backend/services/test_agent_service.py | 6 ++- .../services/test_data_process_service.py | 2 + .../test_knowledge_storage_service.py | 5 +++ .../services/test_vectordatabase_service.py | 6 +++ 8 files changed, 71 insertions(+), 68 deletions(-) diff --git a/backend/consts/const.py b/backend/consts/const.py index 33b0219ab3..545102d98d 100644 --- a/backend/consts/const.py +++ b/backend/consts/const.py @@ -259,6 +259,8 @@ class VectorDatabaseType(str, Enum): # Ray Configuration +DP_PART_PROCESSOR_COUNT = int(os.getenv("DP_PART_PROCESSOR_COUNT", "3")) +DP_FILE_SPLIT_SIZE_MB = int(os.getenv("DP_FILE_SPLIT_SIZE_MB", "5")) RAY_ACTOR_NUM_CPUS = int(os.getenv("RAY_ACTOR_NUM_CPUS", "2")) RAY_DASHBOARD_PORT = int(os.getenv("RAY_DASHBOARD_PORT", "8265")) RAY_DASHBOARD_HOST = os.getenv("RAY_DASHBOARD_HOST", "0.0.0.0") @@ -298,8 +300,6 @@ class VectorDatabaseType(str, Enum): # Will be dynamically set based on PID if not provided WORKER_NAME = os.getenv("WORKER_NAME") WORKER_CONCURRENCY = DP_PART_PROCESSOR_COUNT + 1 -DP_PART_PROCESSOR_COUNT = int(os.getenv("DP_PART_PROCESSOR_COUNT", "3")) -DP_FILE_SPLIT_SIZE_MB = int(os.getenv("DP_FILE_SPLIT_SIZE_MB", "5")) RAY_WARM_ACTOR_POOL_SIZE_PART = int( os.getenv("RAY_WARM_ACTOR_POOL_SIZE_PART", "2")) RAY_WARM_ACTOR_POOL_SIZE_PROCESS = int( diff --git a/test/backend/data_process/test_ray_actors.py b/test/backend/data_process/test_ray_actors.py index 0a2579930a..1d02b08032 100644 --- a/test/backend/data_process/test_ray_actors.py +++ b/test/backend/data_process/test_ray_actors.py @@ -6,6 +6,14 @@ import pytest +class _NoopKnowledgeSpan: + def __enter__(self): + return self + + def __exit__(self, *args): + return None + + def make_fake_ray_module_identity_decorator(): fake_ray = types.ModuleType("ray") @@ -88,6 +96,10 @@ def import_module(monkeypatch): # Stub DataProcessCore and get_file_stream monkeypatch.setitem(sys.modules, "nexent.data_process", types.SimpleNamespace(DataProcessCore=FakeDataProcessCore)) + telemetry_module = types.SimpleNamespace( + knowledge_span=lambda *args, **kwargs: _NoopKnowledgeSpan() + ) + monkeypatch.setitem(sys.modules, "utils.knowledge_telemetry", telemetry_module) # Provide a full stub module for database.attachment_db to avoid importing real Minio client fake_attachment_db_mod = types.ModuleType("database.attachment_db") @@ -792,8 +804,8 @@ def file_process(self, *a, **k): return ( [{"content": "text", "metadata": {}}], [ - {"not": "dict"}, # Not a dict - {"image_bytes": b"img"}, # Missing image_format + "not-a-dict", + {"image_format": "png"}, # Missing image_bytes ], ) @@ -812,8 +824,8 @@ def file_process(self, *a, **k): actor = ray_actors.DataProcessorRayActor() chunks = [{"content": "text", "metadata": {}}] images = [ - {"not": "dict"}, - {"image_bytes": b"img"}, + "not-a-dict", + {"image_format": "png"}, ] actor._append_image_chunks("source.pdf", chunks, images) # Only valid text chunk should remain, no image chunks added @@ -1010,6 +1022,7 @@ def __init__(self, name, operation, **kwargs): self.name = name self.operation = operation self.kwargs = kwargs + captured_spans.append(self) def __enter__(self): return self @@ -1026,11 +1039,8 @@ def file_process(self, file_data, filename, chunking_strategy, **params): return [{"content": "test content", "metadata": {"creation_date": "2024-01-01"}}] monkeypatch.setattr(ray_actors, "DataProcessCore", RecordingCore) - monkeypatch.setattr( - ray_actors, - "knowledge_span", - MockKnowledgeSpan, - ) + telemetry_module = types.SimpleNamespace(knowledge_span=MockKnowledgeSpan) + monkeypatch.setitem(sys.modules, "utils.knowledge_telemetry", telemetry_module) actor = ray_actors.DataProcessorRayActor() result = actor._run_file_process( @@ -1046,4 +1056,5 @@ def file_process(self, file_data, filename, chunking_strategy, **params): assert len(result) == 1 assert result[0]["content"] == "test content" - + assert len(captured_spans) == 1 + assert captured_spans[0].kwargs["telemetry_context"] == {"trace_id": "abc123"} diff --git a/test/backend/data_process/test_tasks.py b/test/backend/data_process/test_tasks.py index 1716265acc..f69b1746d1 100644 --- a/test/backend/data_process/test_tasks.py +++ b/test/backend/data_process/test_tasks.py @@ -1,5 +1,6 @@ import asyncio import io +import math import sys import types import json @@ -696,12 +697,6 @@ def _unbound_run(task_obj): if not hasattr(tasks, "DataProcessorRayActor") or not hasattr(getattr(tasks, "DataProcessorRayActor"), "remote"): tasks.DataProcessorRayActor = types.SimpleNamespace( remote=lambda: default_actor) - # Keep split path stable across tests even when get_ray_actor is monkeypatched. - tasks._get_split_actor = lambda: types.SimpleNamespace( - split_file=types.SimpleNamespace( - remote=lambda *a, **k: "__split_parts__") - ) - # Preprocess for forward: drop empty/whitespace-only chunks before calling real run def _forward_preprocess(args, kwargs): pd = kwargs.get("processed_data") @@ -3496,15 +3491,14 @@ def test_process_sync_with_celery_context(monkeypatch, tmp_path): class FakeActor: def __init__(self): - pass - - def process_file(self, *args, **kwargs): - class Ref: - pass - return Ref() + self.process_file = types.SimpleNamespace( + remote=lambda *args, **kwargs: "__process_ref__" + ) fake_ray = sys.modules.get("ray") - fake_ray.get_returns = [{"content": "hello world", "metadata": {}}] + fake_ray.get_returns = { + "__process_ref__": [{"content": "hello world", "metadata": {}}] + } monkeypatch.setattr(tasks, "get_ray_actor", lambda: FakeActor()) @@ -3686,15 +3680,17 @@ def test_prewarm_ray_actors(monkeypatch): class MockManager: def __init__(self, warm_timeout_s): - pass + self.ensure_pool = types.SimpleNamespace(remote=self._ensure_pool) - def ensure_pool(self, desired, max_allowed): + @staticmethod + def _ensure_pool(desired, max_allowed): captured["desired"] = desired captured["max_allowed"] = max_allowed - return 3 + return "__pool_ref__" monkeypatch.setattr(tasks, "_get_or_create_global_pool_manager", lambda: MockManager(60)) monkeypatch.setattr(tasks, "_estimate_parallel_parts", lambda: 2) + sys.modules["ray"].get_returns = {"__pool_ref__": 3} result = tasks.prewarm_ray_actors(target_size=5) assert result == 3 @@ -3709,19 +3705,8 @@ def test_get_split_actor(monkeypatch): class MockActor: pass - class MockManager: - def get_actor(self): - return MockActor() - - captured_manager = [] - - def mock_get_manager(): - manager = MockManager() - captured_manager.append(manager) - return manager - - monkeypatch.setattr(tasks, "_get_or_create_global_pool_manager", mock_get_manager) + expected_actor = MockActor() + monkeypatch.setattr(tasks, "get_ray_actor", lambda: expected_actor) actor = tasks._get_split_actor() - assert actor is MockActor - assert len(captured_manager) == 1 + assert actor is expected_actor diff --git a/test/backend/data_process/test_worker.py b/test/backend/data_process/test_worker.py index 2a1b1fc199..830cc1fe03 100644 --- a/test/backend/data_process/test_worker.py +++ b/test/backend/data_process/test_worker.py @@ -1,3 +1,4 @@ +import builtins import sys import types import importlib @@ -69,6 +70,7 @@ def setup_mocks_for_worker(mocker, initialized=False): const_mod.RAY_ACTOR_NUM_CPUS = 1 const_mod.RAY_NUM_CPUS = 4 const_mod.DP_PART_PROCESSOR_COUNT = 3 + const_mod.DP_FILE_SPLIT_SIZE_MB = 5 const_mod.PER_WAVE_TIMEOUT = 300 const_mod.MAX_TIMEOUT = 3600 const_mod.RAY_ACTOR_WARM_TIMEOUT_S = 60 @@ -467,11 +469,7 @@ def test_start_worker_with_custom_name(mocker): worker_module, _ = setup_mocks_for_worker(mocker) # Set custom worker name - if "consts.const" in sys.modules: - sys.modules["consts.const"].WORKER_NAME = "custom-worker" - - # Reload to pick up new constant value - importlib.reload(worker_module) + worker_module.WORKER_NAME = "custom-worker" call_args = [] @@ -797,18 +795,19 @@ def test_setup_worker_environment_sets_logging_level(mocker): logger_capture = [] class FakeCeleryWorkerStrategyLogger: - def __init__(self): - pass - def setLevel(self, level): logger_capture.append(level) import logging - mocker.patch.dict(sys.modules, { - "celery.worker.strategy": types.SimpleNamespace( - Logger=FakeCeleryWorkerStrategyLogger - ) - }) + strategy_logger = FakeCeleryWorkerStrategyLogger() + real_get_logger = logging.getLogger + mocker.patch.object( + worker_module.logging, + "getLogger", + side_effect=lambda name=None: strategy_logger + if name == "celery.worker.strategy" + else real_get_logger(name), + ) worker_module.setup_worker_environment() assert logging.WARNING in logger_capture @@ -880,8 +879,7 @@ def test_worker_ready_handler_with_process_part_queue(mocker): mocker.patch("backend.data_process.worker.os.getpid", return_value=7) # Mock QUEUES to include process_part_q - if "consts.const" in sys.modules: - sys.modules["consts.const"].QUEUES = "process_part_q" + worker_module.QUEUES = "process_part_q" calls = [] @@ -894,10 +892,6 @@ def start(self): mocker.patch.object(worker_module.threading, "Thread", FakeThread) - # Need to reload to pick up new QUEUES value - import importlib - importlib.reload(worker_module) - worker_module.worker_ready_handler() # Should have started prewarm thread and potentially part concurrency thread assert len(calls) >= 0 # Threads may or may not be started depending on queue config @@ -980,7 +974,7 @@ def mock_worker_main(args): worker_module.start_worker() # Verify configuration logging - assert any("broker_url" in str(call) or "result_backend" in str(call) for call in debug_calls) + assert any("Broker URL" in str(call) or "Backend URL" in str(call) for call in debug_calls) def test_task_failure_handler_logs_exception_details(mocker): @@ -1052,8 +1046,7 @@ def test_start_worker_with_multiple_queues(mocker): worker_module, _ = setup_mocks_for_worker(mocker) # Set multiple queues - if "consts.const" in sys.modules: - sys.modules["consts.const"].QUEUES = "process_q,forward_q,custom_q" + worker_module.QUEUES = "process_q,forward_q,custom_q" mocker.patch("backend.data_process.worker.os.getpid", return_value=12345) @@ -1079,8 +1072,7 @@ def test_worker_ready_handler_schedules_prewarm_for_process_queue(mocker): mocker.patch("backend.data_process.worker.time.time", return_value=1001.0) # Set QUEUES to include process_q - if "consts.const" in sys.modules: - sys.modules["consts.const"].QUEUES = "process_q,forward_q" + worker_module.QUEUES = "process_q,forward_q" prewarm_calls = [] diff --git a/test/backend/services/test_agent_service.py b/test/backend/services/test_agent_service.py index 649786d9ae..94eec89a8a 100644 --- a/test/backend/services/test_agent_service.py +++ b/test/backend/services/test_agent_service.py @@ -188,10 +188,13 @@ async def close(self, *args, **kwargs): sys.modules['agents.create_agent_info'].create_agent_info = mock_create_agent_info # Mock utils submodules -sys.modules['utils'] = MagicMock() +utils_module = types.ModuleType("utils") +utils_module.__path__ = [] +sys.modules['utils'] = utils_module sys.modules['utils.auth_utils'] = MagicMock() sys.modules['utils.thread_utils'] = MagicMock() sys.modules['utils.context_utils'] = MagicMock() +sys.modules['utils.knowledge_telemetry'] = MagicMock() sys.modules['utils.context_utils'].build_authorized_context_input = ( lambda agent_run_info, historical_context=None: MockContextInput( items=tuple(agent_run_info.agent_config.context_items or ()) @@ -17002,4 +17005,3 @@ def test_inject_user_timezone_time_with_invalid_timezone(): request.headers = {"x-user-timezone": "Invalid/Timezone"} result = _inject_user_timezone_time("What time is it?", request) assert result == "What time is it?" - diff --git a/test/backend/services/test_data_process_service.py b/test/backend/services/test_data_process_service.py index 98d8bf2f8b..8312e790de 100644 --- a/test/backend/services/test_data_process_service.py +++ b/test/backend/services/test_data_process_service.py @@ -1707,6 +1707,7 @@ async def async_test_create_batch_tasks_impl_success(self, mock_submit_chain): 'authorization': 'Bearer test_token', 'embedding_model_id': None, 'tenant_id': None, + 'telemetry_context': {}, }, { 'source': 'http://example.com/doc2.pdf', @@ -1717,6 +1718,7 @@ async def async_test_create_batch_tasks_impl_success(self, mock_submit_chain): 'authorization': 'Bearer test_token', 'embedding_model_id': None, 'tenant_id': None, + 'telemetry_context': {}, }, ] actual_calls = [kwargs for args, kwargs in mock_submit_chain.call_args_list] diff --git a/test/backend/services/test_knowledge_storage_service.py b/test/backend/services/test_knowledge_storage_service.py index 5e1ba418d0..20544eaf8a 100644 --- a/test/backend/services/test_knowledge_storage_service.py +++ b/test/backend/services/test_knowledge_storage_service.py @@ -16,6 +16,11 @@ resolve_storage_reference = storage_service.resolve_storage_reference +@pytest.fixture(autouse=True) +def default_bucket(monkeypatch): + monkeypatch.setattr(storage_service, "MINIO_DEFAULT_BUCKET", "test-bucket") + + @pytest.fixture def storage_context(): return KnowledgeStorageContext( diff --git a/test/backend/services/test_vectordatabase_service.py b/test/backend/services/test_vectordatabase_service.py index 57ff793dec..ab1f93715b 100644 --- a/test/backend/services/test_vectordatabase_service.py +++ b/test/backend/services/test_vectordatabase_service.py @@ -362,6 +362,12 @@ async def _mock_get_all_files_status(index_name): setattr(sys.modules['utils'], 'config_utils', config_utils_mock) setattr(sys.modules['backend.utils'], 'config_utils', config_utils_mock) +knowledge_telemetry_mock = types.ModuleType('utils.knowledge_telemetry') +knowledge_telemetry_mock.set_span_attributes = MagicMock() +knowledge_telemetry_mock.trace_knowledge_operation = MagicMock() +sys.modules['utils.knowledge_telemetry'] = knowledge_telemetry_mock +setattr(sys.modules['utils'], 'knowledge_telemetry', knowledge_telemetry_mock) + # Shared mock instances for MinIO storage_client_mock = MagicMock() storage_client_mock.delete_file.return_value = (True, None) From 829cdb945d5b8d95b96a70961a54ed3a4fd95511 Mon Sep 17 00:00:00 2001 From: Jasonxia007 Date: Fri, 14 Aug 2026 15:03:07 +0800 Subject: [PATCH 4/5] =?UTF-8?q?=F0=9F=A7=AA=20Add=20test=20files?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- test/backend/data_process/test_worker.py | 65 ++++++++++++------------ 1 file changed, 32 insertions(+), 33 deletions(-) diff --git a/test/backend/data_process/test_worker.py b/test/backend/data_process/test_worker.py index 830cc1fe03..acf896b013 100644 --- a/test/backend/data_process/test_worker.py +++ b/test/backend/data_process/test_worker.py @@ -44,39 +44,39 @@ def setup_mocks_for_worker(mocker, initialized=False): # Mock ray module mocker.patch.dict(sys.modules, {"ray": fake_ray}) - # Stub consts.const module + # Stub consts.const with deterministic values even when another test has + # already imported the real configuration module. if "consts" not in sys.modules: sys.modules["consts"] = types.ModuleType("consts") setattr(sys.modules["consts"], "__path__", []) - if "consts.const" not in sys.modules: - const_mod = types.ModuleType("consts.const") - const_mod.CELERY_TASK_TIME_LIMIT = 3600 - const_mod.CELERY_WORKER_PREFETCH_MULTIPLIER = 1 - const_mod.ELASTICSEARCH_SERVICE = "http://elasticsearch:9200" - const_mod.QUEUES = "process_q,process_part_q,forward_q" - const_mod.RAY_ADDRESS = "auto" - const_mod.RAY_preallocate_plasma = False - const_mod.REDIS_URL = "redis://localhost:6379" - const_mod.REDIS_BACKEND_URL = "redis://localhost:6379" - const_mod.WORKER_CONCURRENCY = 4 - const_mod.WORKER_NAME = None - const_mod.FORWARD_REDIS_RETRY_DELAY_S = 0 - const_mod.FORWARD_REDIS_RETRY_MAX = 1 - const_mod.DISABLE_RAY_DASHBOARD = False - const_mod.DATA_PROCESS_SERVICE = "http://data-process" - const_mod.ROOT_DIR = "/mock/root" - const_mod.DP_REDIS_CHUNKS_WAIT_TIMEOUT_S = 30 - const_mod.DP_REDIS_CHUNKS_POLL_INTERVAL_MS = 100 - const_mod.RAY_ACTOR_NUM_CPUS = 1 - const_mod.RAY_NUM_CPUS = 4 - const_mod.DP_PART_PROCESSOR_COUNT = 3 - const_mod.DP_FILE_SPLIT_SIZE_MB = 5 - const_mod.PER_WAVE_TIMEOUT = 300 - const_mod.MAX_TIMEOUT = 3600 - const_mod.RAY_ACTOR_WARM_TIMEOUT_S = 60 - const_mod.RAY_GLOBAL_ACTOR_POOL_NAME = "global_actor_pool" - const_mod.RAY_GLOBAL_ACTOR_POOL_NAMESPACE = "nexent" - sys.modules["consts.const"] = const_mod + const_mod = types.ModuleType("consts.const") + const_mod.CELERY_TASK_TIME_LIMIT = 3600 + const_mod.CELERY_WORKER_PREFETCH_MULTIPLIER = 1 + const_mod.ELASTICSEARCH_SERVICE = "http://elasticsearch:9200" + const_mod.QUEUES = "process_q,process_part_q,forward_q" + const_mod.RAY_ADDRESS = "auto" + const_mod.RAY_preallocate_plasma = False + const_mod.REDIS_URL = "redis://localhost:6379" + const_mod.REDIS_BACKEND_URL = "redis://localhost:6379" + const_mod.WORKER_CONCURRENCY = 4 + const_mod.WORKER_NAME = None + const_mod.FORWARD_REDIS_RETRY_DELAY_S = 0 + const_mod.FORWARD_REDIS_RETRY_MAX = 1 + const_mod.DISABLE_RAY_DASHBOARD = False + const_mod.DATA_PROCESS_SERVICE = "http://data-process" + const_mod.ROOT_DIR = "/mock/root" + const_mod.DP_REDIS_CHUNKS_WAIT_TIMEOUT_S = 30 + const_mod.DP_REDIS_CHUNKS_POLL_INTERVAL_MS = 100 + const_mod.RAY_ACTOR_NUM_CPUS = 1 + const_mod.RAY_NUM_CPUS = 4 + const_mod.DP_PART_PROCESSOR_COUNT = 3 + const_mod.DP_FILE_SPLIT_SIZE_MB = 5 + const_mod.PER_WAVE_TIMEOUT = 300 + const_mod.MAX_TIMEOUT = 3600 + const_mod.RAY_ACTOR_WARM_TIMEOUT_S = 60 + const_mod.RAY_GLOBAL_ACTOR_POOL_NAME = "global_actor_pool" + const_mod.RAY_GLOBAL_ACTOR_POOL_NAMESPACE = "nexent" + mocker.patch.dict(sys.modules, {"consts.const": const_mod}) # Stub celery module and submodules (required by tasks.py imported via __init__.py) if "celery.backends.base" not in sys.modules: @@ -856,10 +856,9 @@ def test_validate_service_connections_handles_redis_exception(mocker): assert result is False -def test_worker_state_keys_exist(): +def test_worker_state_keys_exist(mocker): """Test worker_state has all required keys.""" - # Test that worker module has worker_state with required structure - import backend.data_process.worker as worker_module + worker_module, _ = setup_mocks_for_worker(mocker) assert "initialized" in worker_module.worker_state assert "ready" in worker_module.worker_state From bacdd3e76fbe0c3e9c32699d3db57b1fe21cf541 Mon Sep 17 00:00:00 2001 From: Jasonxia007 Date: Fri, 14 Aug 2026 15:59:33 +0800 Subject: [PATCH 5/5] =?UTF-8?q?=F0=9F=A7=AA=20Add=20test=20files?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- test/backend/data_process/test_ray_actors.py | 64 +++++++ test/backend/data_process/test_tasks.py | 180 ++++++++++++++++++ test/backend/data_process/test_worker.py | 149 +++++++++++++++ .../backend/utils/test_knowledge_telemetry.py | 123 ++++++++++++ 4 files changed, 516 insertions(+) diff --git a/test/backend/data_process/test_ray_actors.py b/test/backend/data_process/test_ray_actors.py index 1d02b08032..e3267a5e54 100644 --- a/test/backend/data_process/test_ray_actors.py +++ b/test/backend/data_process/test_ray_actors.py @@ -1058,3 +1058,67 @@ def file_process(self, file_data, filename, chunking_strategy, **params): assert result[0]["content"] == "test content" assert len(captured_spans) == 1 assert captured_spans[0].kwargs["telemetry_context"] == {"trace_id": "abc123"} + + +def test_actor_initializes_monitoring_when_available(monkeypatch): + ray_actors = import_module(monkeypatch) + manager = types.SimpleNamespace(is_enabled=True) + monitoring_module = types.SimpleNamespace(monitoring_manager=manager) + monkeypatch.setitem(sys.modules, "utils.monitoring", monitoring_module) + + actor = ray_actors.DataProcessorRayActor() + + assert actor._monitoring_manager is manager + + +def test_actor_degrades_when_monitoring_status_fails(monkeypatch): + ray_actors = import_module(monkeypatch) + + class BrokenManager: + @property + def is_enabled(self): + raise RuntimeError("monitoring unavailable") + + monitoring_module = types.SimpleNamespace(monitoring_manager=BrokenManager()) + monkeypatch.setitem(sys.modules, "utils.monitoring", monitoring_module) + + actor = ray_actors.DataProcessorRayActor() + + assert actor._monitoring_manager is None + + +def test_split_file_fetches_stream_and_model_type_is_optional(monkeypatch): + ray_actors = import_module(monkeypatch) + + class Part: + def getvalue(self): + return b"part" + + class RecordingCore(FakeDataProcessCore): + captured_file_data = None + + def file_split(self, file_data, **kwargs): + RecordingCore.captured_file_data = file_data + return [Part()] + + monkeypatch.setattr(ray_actors, "DataProcessCore", RecordingCore) + monkeypatch.setattr(ray_actors, "get_file_stream", lambda source: io.BytesIO(b"stream-data")) + monkeypatch.setattr( + ray_actors, + "get_model_by_model_id", + lambda model_id, tenant_id=None: { + "expected_chunk_size": 100, + "maximum_chunk_size": 200, + "display_name": "model-without-type", + "model_type": None, + }, + ) + + actor = ray_actors.DataProcessorRayActor() + params = {} + actor._apply_model_chunk_sizes(model_id=1, tenant_id="tenant", params=params) + parts = actor.split_file("s3://bucket/source.pdf", "minio") + + assert "model_type" not in params + assert parts == [b"part"] + assert RecordingCore.captured_file_data == b"stream-data" diff --git a/test/backend/data_process/test_tasks.py b/test/backend/data_process/test_tasks.py index f69b1746d1..e8cdf1115e 100644 --- a/test/backend/data_process/test_tasks.py +++ b/test/backend/data_process/test_tasks.py @@ -3710,3 +3710,183 @@ class MockActor: actor = tasks._get_split_actor() assert actor is expected_actor + + +def test_fetch_minio_source_success_and_missing_stream(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + attributes = [] + monkeypatch.setattr(tasks, "set_span_attributes", lambda **kwargs: attributes.append(kwargs)) + monkeypatch.setattr(tasks, "get_file_stream", lambda source: io.BytesIO(b"payload")) + + assert tasks._fetch_minio_source("s3://bucket/file") == b"payload" + assert attributes == [{"file_size_bytes": 7, "stage": "minio.fetch"}] + + monkeypatch.setattr(tasks, "get_file_stream", lambda source: None) + with pytest.raises(FileNotFoundError, match="Unable to fetch"): + tasks._fetch_minio_source("s3://bucket/missing") + + +def test_wait_for_split_ready_handles_missing_and_non_list_payloads(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + monkeypatch.setattr(tasks, "REDIS_BACKEND_URL", "redis://x") + + class Client: + def __init__(self, cached): + self.cached = cached + + def get(self, key): + return "1" if key.endswith(":ready") else self.cached + + current = {"client": Client(None)} + monkeypatch.setitem( + sys.modules, + "redis", + types.SimpleNamespace( + Redis=types.SimpleNamespace(from_url=lambda *args, **kwargs: current["client"]) + ), + ) + + assert tasks._wait_for_split_ready("dp:key", 1, 1) == 0 + current["client"] = Client('{"chunk": 1}') + assert tasks._wait_for_split_ready("dp:key", 1, 1) == 0 + + +def test_distribute_chunks_reports_capacity_context(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + + with pytest.raises(RuntimeError, match="while distributing text chunks"): + tasks._distribute_chunks_round_robin( + batches=[[{"content": "full"}]], + chunks=[{"content": "extra"}], + batch_size=1, + error_context="text chunks", + ) + + +def test_extract_error_code_handles_regex_failure(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + monkeypatch.setattr(tasks.re, "search", lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("regex"))) + + assert tasks.extract_error_code("plain error") == "unknown_error" + assert tasks._extract_error_code_from_es_response(None, "plain error") is None + + +def test_delete_source_file_handles_non_json_response(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + + class Response: + status_code = 503 + text = "service unavailable" + + def json(self): + raise ValueError("not json") + + monkeypatch.setattr(tasks.requests, "delete", lambda *args, **kwargs: Response()) + + result = tasks._delete_source_file_via_http_sync( + base_url="http://api/", + index_name="index", + path_or_url="s3://bucket/file", + scope="source_only", + ) + + assert result == { + "http_status": 503, + "response_json": None, + "response_text": "service unavailable", + } + + +def test_global_pool_manager_tolerates_actor_kill_failures(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + + class Actor: + ping = types.SimpleNamespace(remote=lambda: "ping-ref") + + monkeypatch.setattr( + tasks, + "DataProcessorRayActor", + types.SimpleNamespace(remote=lambda: Actor()), + ) + monkeypatch.setattr(tasks.ray, "get", lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("warm"))) + monkeypatch.setattr(tasks.ray, "kill", lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("kill")), raising=False) + manager = tasks.GlobalRayActorPoolManager(warm_timeout_s=1) + + assert manager._create_and_warm_actor() is None + manager.actors = [Actor()] + assert manager.ensure_pool(desired=0, max_allowed=1) == 0 + + +def test_get_or_create_pool_manager_creates_and_recovers_from_name_race(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + monkeypatch.setattr(tasks, "init_ray_in_worker", lambda: None) + + class Options: + def __init__(self, remote_result=None, remote_error=None): + self.remote_result = remote_result + self.remote_error = remote_error + + def remote(self, timeout): + if self.remote_error: + raise self.remote_error + return self.remote_result + + class ManagerFactory: + def __init__(self, options): + self.options_result = options + + def options(self, **kwargs): + if kwargs.get("get_if_exists"): + raise TypeError("unsupported") + return self.options_result + + calls = {"count": 0} + + def get_actor(*args, **kwargs): + calls["count"] += 1 + if calls["count"] == 1: + raise RuntimeError("missing") + return "raced-manager" + + monkeypatch.setattr(tasks.ray, "get_actor", get_actor, raising=False) + monkeypatch.setattr(tasks, "GlobalRayActorPoolManager", ManagerFactory(Options(remote_result="new-manager"))) + assert tasks._get_or_create_global_pool_manager() == "new-manager" + + calls["count"] = 0 + monkeypatch.setattr( + tasks, + "GlobalRayActorPoolManager", + ManagerFactory(Options(remote_error=RuntimeError("name race"))), + ) + assert tasks._get_or_create_global_pool_manager() == "raced-manager" + + +def test_logging_task_delegates_lifecycle_hooks(monkeypatch): + tasks, _ = import_tasks_with_fake_ray(monkeypatch) + monkeypatch.setattr(tasks.Task, "on_success", lambda self, *args: "success", raising=False) + monkeypatch.setattr(tasks.Task, "on_failure", lambda self, *args: "failure", raising=False) + monkeypatch.setattr(tasks.Task, "on_retry", lambda self, *args: "retry", raising=False) + task = tasks.LoggingTask() + task.name = "logging-task" + + assert task.on_success({}, "task-1", (), {}) == "success" + assert task.on_failure(ValueError("bad"), "task-1", (), {}, None) == "failure" + assert task.on_retry(RuntimeError("later"), "task-1", (), {}, None) == "retry" + + +def test_process_sync_without_celery_id_skips_state_updates(monkeypatch, tmp_path): + tasks, fake_ray = import_tasks_with_fake_ray(monkeypatch, initialized=True) + source = tmp_path / "sync.txt" + source.write_text("text", encoding="utf-8") + + class Actor: + process_file = types.SimpleNamespace(remote=lambda *args, **kwargs: "chunks-ref") + + fake_ray.get_returns = {"chunks-ref": [{"content": "one"}, {"content": "two"}]} + monkeypatch.setattr(tasks, "get_ray_actor", lambda: Actor()) + self = FakeSelf(None) + + result = tasks.process_sync(self, str(source), "local") + + assert result["text"] == "one\n\ntwo" + assert self.states == [] diff --git a/test/backend/data_process/test_worker.py b/test/backend/data_process/test_worker.py index acf896b013..8c8c25c4c2 100644 --- a/test/backend/data_process/test_worker.py +++ b/test/backend/data_process/test_worker.py @@ -1124,6 +1124,155 @@ def mock_import(name, *args, **kwargs): pass # May or may not raise depending on import structure +def test_setup_worker_environment_logs_missing_sensitive_variables(mocker): + worker_module, _ = setup_mocks_for_worker(mocker, initialized=True) + worker_module.REDIS_URL = "" + worker_module.ELASTICSEARCH_SERVICE = "" + error_calls = [] + mocker.patch.object( + worker_module.logger, + "error", + side_effect=lambda msg, *args, **kwargs: error_calls.append(msg % args if args else msg), + ) + + worker_module.setup_worker_environment() + + assert any("REDIS_URL: NOT SET" in message for message in error_calls) + assert any("ELASTICSEARCH_SERVICE: NOT SET" in message for message in error_calls) + + +def test_setup_worker_process_resources_logs_monitoring_status(mocker): + worker_module, _ = setup_mocks_for_worker(mocker) + monitoring_module = types.ModuleType("utils.monitoring") + monitoring_module.monitoring_manager = types.SimpleNamespace(is_enabled=True) + mocker.patch.dict(sys.modules, {"utils.monitoring": monitoring_module}) + mocker.patch.object(worker_module, "validate_service_connections", return_value=True) + info = mocker.patch.object(worker_module.logger, "info") + + worker_module.setup_worker_process_resources() + + assert any( + call.args[:2] == ( + "Knowledge telemetry initialized in worker process: enabled=%s", + True, + ) + for call in info.call_args_list + ) + assert worker_module.worker_state["services_validated"] is True + + +@pytest.mark.parametrize("should_fail", [False, True]) +def test_worker_ready_handler_runs_prewarm_background(mocker, should_fail): + worker_module, _ = setup_mocks_for_worker(mocker) + worker_module.QUEUES = "process_q" + prewarm_calls = [] + + def prewarm(target_size): + prewarm_calls.append(target_size) + if should_fail: + raise RuntimeError("warm failed") + return 2 + + tasks_module = types.ModuleType("data_process.tasks") + tasks_module.prewarm_ray_actors = prewarm + mocker.patch.dict(sys.modules, {"data_process.tasks": tasks_module}) + + class ImmediateThread: + def __init__(self, target=None, daemon=None): + self.target = target + + def start(self): + self.target() + + mocker.patch.object(worker_module.threading, "Thread", ImmediateThread) + warning = mocker.patch.object(worker_module.logger, "warning") + + worker_module.worker_ready_handler() + + assert prewarm_calls == [worker_module.DP_PART_PROCESSOR_COUNT] + if should_fail: + assert any("Background prewarm failed" in call.args[0] for call in warning.call_args_list) + + +def test_worker_ready_handler_collects_part_concurrency_once(mocker): + worker_module, fake_ray = setup_mocks_for_worker(mocker, initialized=True) + worker_module.QUEUES = "process_part_q" + tasks_module = types.ModuleType("data_process.tasks") + tasks_module.prewarm_ray_actors = lambda target_size: 1 + mocker.patch.dict(sys.modules, {"data_process.tasks": tasks_module}) + + targets = [] + + class CapturingThread: + def __init__(self, target=None, daemon=None): + targets.append(target) + + def start(self): + pass + + inspector = types.SimpleNamespace( + active=lambda: { + "worker-1": [ + {"name": "data_process.tasks.process_part"}, + {"name": "another.task"}, + ], + "worker-2": None, + } + ) + worker_module.app.control = types.SimpleNamespace( + inspect=lambda timeout: inspector + ) + fake_ray.available_resources = lambda: {"CPU": 3.5} + mocker.patch.object(worker_module.threading, "Thread", CapturingThread) + mocker.patch.object(worker_module.time, "sleep", side_effect=StopIteration) + info = mocker.patch.object(worker_module.logger, "info") + + worker_module.worker_ready_handler() + assert len(targets) == 2 + with pytest.raises(StopIteration): + targets[1]() + + assert any( + "active=1, ray_available_cpu=3.5" in call.args[0] + for call in info.call_args_list + ) + + +def test_worker_ready_handler_logs_thread_schedule_failures(mocker): + worker_module, _ = setup_mocks_for_worker(mocker) + worker_module.QUEUES = "process_part_q" + tasks_module = types.ModuleType("data_process.tasks") + tasks_module.prewarm_ray_actors = lambda target_size: 1 + mocker.patch.dict(sys.modules, {"data_process.tasks": tasks_module}) + mocker.patch.object(worker_module.threading, "Thread", side_effect=RuntimeError("thread failed")) + warning = mocker.patch.object(worker_module.logger, "warning") + + worker_module.worker_ready_handler() + + messages = [call.args[0] for call in warning.call_args_list] + assert any("Failed to schedule Ray actor prewarm" in message for message in messages) + assert any("Failed to start process_part concurrency logger" in message for message in messages) + + +def test_start_worker_logs_successful_monitoring_initialization(mocker): + worker_module, _ = setup_mocks_for_worker(mocker) + monitoring_module = types.ModuleType("utils.monitoring") + monitoring_module.monitoring_manager = types.SimpleNamespace(is_enabled=True) + mocker.patch.dict(sys.modules, {"utils.monitoring": monitoring_module}) + mocker.patch.object(worker_module.app, "worker_main") + info = mocker.patch.object(worker_module.logger, "info") + + worker_module.start_worker() + + assert any( + call.args[:2] == ( + "Knowledge telemetry initialized before worker start: enabled=%s", + True, + ) + for call in info.call_args_list + ) + + def test_validate_service_connections_returns_true_on_success(mocker): """Test validate_service_connections returns True when all checks pass.""" worker_module, _ = setup_mocks_for_worker(mocker) diff --git a/test/backend/utils/test_knowledge_telemetry.py b/test/backend/utils/test_knowledge_telemetry.py index cf7581f5fe..a2b56f9ac1 100644 --- a/test/backend/utils/test_knowledge_telemetry.py +++ b/test/backend/utils/test_knowledge_telemetry.py @@ -229,3 +229,126 @@ def test_knowledge_span_degrades_when_setup_fails(monkeypatch): with knowledge_telemetry.knowledge_span("knowledge.process", "process") as span: assert span is None + + +def test_helpers_degrade_when_otel_is_unavailable(monkeypatch): + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", False) + + assert knowledge_telemetry._safe_hash("") == "" + assert knowledge_telemetry.inject_trace_context() == {} + assert knowledge_telemetry._safe_attributes({"file_size": "invalid"}) == {} + knowledge_telemetry._record_metrics("process", 1.0, {}) + knowledge_telemetry.set_span_attributes(task_id="task-1") + with knowledge_telemetry.knowledge_span("knowledge.process", "process") as span: + assert span is None + + +def test_resource_snapshot_returns_empty_when_psutil_fails(monkeypatch): + real_import = __import__ + + def fake_import(name, *args, **kwargs): + if name == "psutil": + raise RuntimeError("psutil unavailable") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr("builtins.__import__", fake_import) + + assert knowledge_telemetry._resource_snapshot() == {} + + +def test_record_metrics_swallows_meter_failures(monkeypatch): + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", True) + metrics = MagicMock() + metrics.get_meter.side_effect = RuntimeError("meter unavailable") + monkeypatch.setattr(knowledge_telemetry, "metrics", metrics, raising=False) + + knowledge_telemetry._record_metrics("forward", 2.5, {}) + + metrics.get_meter.assert_called_once_with("nexent.knowledge_ingestion") + + +def test_knowledge_span_success_continues_and_detaches_remote_context(monkeypatch): + span, span_cm, _, propagate, context, status_code = _install_otel_stubs(monkeypatch) + context.attach.return_value = "context-token" + monkeypatch.setattr( + knowledge_telemetry, + "_resource_snapshot", + MagicMock(side_effect=[{"process.rss_memory_mb": 2.0}, {"process.cpu_percent": 3.0}]), + ) + + carrier = {"traceparent": "00-test"} + with knowledge_telemetry.knowledge_span( + "knowledge.process", + "process", + telemetry_context=carrier, + filename="report.PDF", + ) as active_span: + assert active_span is span + + propagate.extract.assert_called_once_with(carrier) + context.attach.assert_called_once_with(propagate.extract.return_value) + context.detach.assert_called_once_with("context-token") + span.set_status.assert_called_with((status_code.OK, None)) + span.set_attribute.assert_any_call("ingestion.status", "success") + assert span_cm.__exit__.call_count == 1 + + +def test_knowledge_span_cleans_up_when_enter_fails(monkeypatch): + _, span_cm, trace, propagate, context, _ = _install_otel_stubs(monkeypatch) + context.attach.return_value = "context-token" + span_cm.__enter__.side_effect = RuntimeError("enter failed") + span_cm.__exit__.side_effect = RuntimeError("close failed") + context.detach.side_effect = RuntimeError("detach failed") + trace.get_tracer.return_value.start_as_current_span.return_value = span_cm + + with knowledge_telemetry.knowledge_span( + "knowledge.process", + "process", + telemetry_context={"traceparent": "00-test"}, + ) as span: + assert span is None + + propagate.extract.assert_called_once() + span_cm.__exit__.assert_called_once_with(None, None, None) + context.detach.assert_called_once_with("context-token") + + +def test_knowledge_span_swallows_reporting_and_cleanup_failures(monkeypatch): + span, span_cm, _, _, context, _ = _install_otel_stubs(monkeypatch) + context.attach.return_value = "context-token" + span.set_status.side_effect = RuntimeError("status failed") + span.set_attribute.side_effect = RuntimeError("attribute failed") + span_cm.__exit__.side_effect = RuntimeError("close failed") + context.detach.side_effect = RuntimeError("detach failed") + + with knowledge_telemetry.knowledge_span( + "knowledge.process", + "process", + telemetry_context={"traceparent": "00-test"}, + ): + pass + + +def test_knowledge_span_preserves_error_when_failure_reporting_fails(monkeypatch): + span, _, _, _, _, _ = _install_otel_stubs(monkeypatch) + span.record_exception.side_effect = RuntimeError("record failed") + + with pytest.raises(ValueError, match="original"), knowledge_telemetry.knowledge_span( + "knowledge.process", "process" + ): + raise ValueError("original") + + +def test_set_span_attributes_handles_inactive_and_failing_spans(monkeypatch): + monkeypatch.setattr(knowledge_telemetry, "OTEL_AVAILABLE", True) + trace = MagicMock() + inactive_span = MagicMock() + inactive_span.is_recording.return_value = False + trace.get_current_span.return_value = inactive_span + monkeypatch.setattr(knowledge_telemetry, "trace", trace, raising=False) + + knowledge_telemetry.set_span_attributes(task_id="task-1") + inactive_span.set_attributes.assert_not_called() + + trace.get_current_span.side_effect = RuntimeError("span unavailable") + knowledge_telemetry.set_span_attributes(task_id="task-2")