diff --git a/.github/workflows/test-knowledge-runtime-postgres.yml b/.github/workflows/test-knowledge-runtime-postgres.yml index 660215d2b..a0538dad8 100644 --- a/.github/workflows/test-knowledge-runtime-postgres.yml +++ b/.github/workflows/test-knowledge-runtime-postgres.yml @@ -24,9 +24,14 @@ on: - "apps/shared/tests/domain/test_knowledge_runtime_candidates.py" - "apps/shared/tests/services/test_knowledge_permission_runtime_bulk.py" - "apps/workflow_engine/adapters/knowledge_runtime_candidates.py" + - "apps/workflow_engine/adapters/rag_retrieval_connection_acquirer.py" + - "apps/workflow_engine/adapters/rag_retrieval_executor.py" + - "apps/workflow_engine/adapters/rag_retrieval_session.py" + - "apps/workflow_engine/application/rag_retrieval_fanout.py" - "apps/workflow_engine/application/runtime_retrieval/**" - "apps/workflow_engine/composition/runtime_retrieval.py" - "apps/workflow_engine/tests/adapters/test_postgres_knowledge_runtime_candidate_adapter.py" + - "apps/workflow_engine/tests/adapters/test_rag_retrieval_session_postgres.py" - ".github/workflows/test-knowledge-runtime-postgres.yml" permissions: @@ -95,4 +100,5 @@ jobs: apps/workflow_engine/.venv/bin/python -m pytest apps/shared/tests/db/test_knowledge_runtime_snapshot_disposable_postgres.py apps/workflow_engine/tests/adapters/test_postgres_knowledge_runtime_candidate_adapter.py + apps/workflow_engine/tests/adapters/test_rag_retrieval_session_postgres.py -q diff --git a/apps/gateway/tests/architecture/test_workflow_node_import_boundary.py b/apps/gateway/tests/architecture/test_workflow_node_import_boundary.py new file mode 100644 index 000000000..d222be444 --- /dev/null +++ b/apps/gateway/tests/architecture/test_workflow_node_import_boundary.py @@ -0,0 +1,83 @@ +"""Gateway와 Workflow worker 패키지의 import 경계를 검증한다.""" + +from __future__ import annotations + +import os +from pathlib import Path +import subprocess +import sys + + +REPOSITORY_ROOT = Path(__file__).resolve().parents[4] + + +def test_llm_entity_import_does_not_load_worker_runtime() -> None: + """Data-only schema import must not require worker-only dependencies.""" + + script = """ +import sys + +from apps.workflow_engine.workflow.nodes.llm.entities import LLMNodeData + +assert LLMNodeData.__name__ == "LLMNodeData" +assert "apps.workflow_engine.workflow.nodes.llm.llm_node" not in sys.modules +""" + environment = os.environ.copy() + existing_pythonpath = environment.get("PYTHONPATH") + environment["PYTHONPATH"] = os.pathsep.join( + value + for value in (str(REPOSITORY_ROOT), existing_pythonpath) + if value + ) + + completed = subprocess.run( + [sys.executable, "-c", script], + cwd=REPOSITORY_ROOT, + env=environment, + capture_output=True, + text=True, + check=False, + ) + + assert completed.returncode == 0, completed.stderr + + +def test_node_factory_import_does_not_require_gevent_runtime() -> None: + """Gateway-side graph validation must not require worker-only gevent.""" + + script = """ +import sys + + +class BlockGeventImport: + def find_spec(self, fullname, path=None, target=None): + if fullname == "gevent" or fullname.startswith("gevent."): + raise ModuleNotFoundError("blocked worker-only gevent dependency") + return None + + +sys.meta_path.insert(0, BlockGeventImport()) + +from apps.workflow_engine.workflow.core.workflow_node_factory import NodeFactory + +assert "llmNode" in NodeFactory.NODE_REGISTRY +assert "gevent" not in sys.modules +""" + environment = os.environ.copy() + existing_pythonpath = environment.get("PYTHONPATH") + environment["PYTHONPATH"] = os.pathsep.join( + value + for value in (str(REPOSITORY_ROOT), existing_pythonpath) + if value + ) + + completed = subprocess.run( + [sys.executable, "-c", script], + cwd=REPOSITORY_ROOT, + env=environment, + capture_output=True, + text=True, + check=False, + ) + + assert completed.returncode == 0, completed.stderr diff --git a/apps/shared/services/tracing/metadata.py b/apps/shared/services/tracing/metadata.py index 55d513057..4994f09f6 100644 --- a/apps/shared/services/tracing/metadata.py +++ b/apps/shared/services/tracing/metadata.py @@ -246,9 +246,11 @@ "total_tokens", }, "rag": { + "candidate_resolution_latency_ms", "citation_ids", "context_token_estimate", "document_ids", + "evidence_policy_latency_ms", "evidence_sufficient", "fanout_concurrency", "fanout_timeout_seconds", @@ -262,10 +264,12 @@ "permission_filter_applied", "latency_ms", "partial_result", + "query_embedding_latency_ms", "query_rewrite_applied", "query_rewrite_strategy", "raw_content_returned", "rag_mode", + "retrieval_fanout_latency_ms", "retrieval_payload_id", "retrieval_strategy", "retrieved_chunk_summary_truncated", @@ -275,6 +279,7 @@ "score_summary", "selected_kb_count", "selected_kb_count_bucket", + "slowest_search_latency_ms", "source_tier_policy", "source_tier_used", "stored_result_count", @@ -303,6 +308,14 @@ "latency_ms", }, } +RAG_STAGE_LATENCY_FIELDS = { + "candidate_resolution_latency_ms", + "query_embedding_latency_ms", + "retrieval_fanout_latency_ms", + "slowest_search_latency_ms", + "evidence_policy_latency_ms", +} +MAX_RAG_STAGE_LATENCY_MS = 300_000 RAG_RESULT_FIELDS = { "document_id", "chunk_id", @@ -780,6 +793,14 @@ def _sanitize_rag_section(cls, value: Any) -> dict[str, Any]: sanitized: dict[str, Any] = {} for key in SPAN_SECTION_FIELDS["rag"]: if key in safe_value: + if key in RAG_STAGE_LATENCY_FIELDS: + latency = safe_value[key] + if ( + type(latency) is int + and 0 <= latency <= MAX_RAG_STAGE_LATENCY_MS + ): + sanitized[key] = latency + continue sanitized_value = cls._sanitize_allowed_value(safe_value[key]) if sanitized_value is not None: sanitized[key] = sanitized_value diff --git a/apps/shared/tests/services/test_tracing_metadata.py b/apps/shared/tests/services/test_tracing_metadata.py index c941a6580..1c26a8698 100644 --- a/apps/shared/tests/services/test_tracing_metadata.py +++ b/apps/shared/tests/services/test_tracing_metadata.py @@ -582,6 +582,43 @@ def test_rag_span_metadata_preserves_evidence_summary_fields_only(): assert "raw_rewritten_query" not in metadata["rag"] +def test_rag_span_metadata_allows_only_bounded_integer_stage_latencies(): + metadata = TraceMetadataSanitizer.sanitize_span_metadata( + "llmNode", + { + "rag": { + "candidate_resolution_latency_ms": 0, + "query_embedding_latency_ms": 12, + "retrieval_fanout_latency_ms": 34, + "slowest_search_latency_ms": 56, + "evidence_policy_latency_ms": 300_000, + "per_kb_latency_ms": {"hidden-resource": 99}, + } + }, + ) + + assert metadata["rag"] == { + "candidate_resolution_latency_ms": 0, + "query_embedding_latency_ms": 12, + "retrieval_fanout_latency_ms": 34, + "slowest_search_latency_ms": 56, + "evidence_policy_latency_ms": 300_000, + } + + +@pytest.mark.parametrize( + "invalid_value", + [True, -1, 1.5, "12", float("nan"), float("inf"), 300_001], +) +def test_rag_span_metadata_drops_invalid_stage_latency(invalid_value): + metadata = TraceMetadataSanitizer.sanitize_span_metadata( + "llmNode", + {"rag": {"retrieval_fanout_latency_ms": invalid_value}}, + ) + + assert metadata.get("rag", {}) == {} + + def test_trace_detail_metadata_view_hides_error_message_and_sanitizes_metadata(): run = SimpleNamespace( id=uuid.uuid4(), diff --git a/apps/workflow_engine/adapters/rag_retrieval_connection_acquirer.py b/apps/workflow_engine/adapters/rag_retrieval_connection_acquirer.py new file mode 100644 index 000000000..f87731cc2 --- /dev/null +++ b/apps/workflow_engine/adapters/rag_retrieval_connection_acquirer.py @@ -0,0 +1,243 @@ +"""Bounded native-thread connection acquisition for RAG retrieval sessions.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Callable, Protocol + +from gevent.monkey import get_original + +from apps.workflow_engine.application.rag_retrieval_fanout import ( + DEFAULT_RAG_FANOUT_MAX_WORKERS, + RAGRetrievalCancellation, +) + + +_allocate_native_lock = get_original("_thread", "allocate_lock") +_start_native_thread = get_original("_thread", "start_new_thread") +_NANOSECONDS_PER_SECOND = 1_000_000_000 +_CANCELLATION_POLL_SECONDS = 0.01 + + +class RAGRetrievalConnectionAcquisitionError(RuntimeError): + """Redaction-safe connection acquisition failure.""" + + code = "rag.retrieval_connection_acquisition_failed" + + def __init__(self) -> None: + super().__init__("RAG retrieval connection acquisition failed.") + + +class RAGRetrievalConnectionAcquisitionTimeout(TimeoutError): + """Redaction-safe connection acquisition timeout.""" + + code = "rag.retrieval_timed_out" + + def __init__(self) -> None: + super().__init__("RAG retrieval timed out.") + + +@dataclass(frozen=True, slots=True) +class AcquiredRAGRetrievalConnection: + session: object = field(repr=False) + connection: object = field(repr=False) + + +class RAGRetrievalConnectionAcquirer(Protocol): + def acquire( + self, + *, + session_factory: Callable[[], object], + cancellation: RAGRetrievalCancellation, + deadline_ns: int, + monotonic_ns: Callable[[], int], + ) -> AcquiredRAGRetrievalConnection: ... + + +class _ConnectionAcquisitionState: + def __init__(self) -> None: + self._lock = _allocate_native_lock() + self._signal = _allocate_native_lock() + self._signal.acquire() + self._status = "pending" + self._acquired: AcquiredRAGRetrievalConnection | None = None + + @property + def status(self) -> str: + with self._lock: + return self._status + + def abandon(self) -> bool: + with self._lock: + if self._status != "pending": + return False + self._status = "abandoned" + self._signal.release() + return True + + def publish_success( + self, + acquired: AcquiredRAGRetrievalConnection, + ) -> bool: + with self._lock: + if self._status != "pending": + return False + self._status = "succeeded" + self._acquired = acquired + self._signal.release() + return True + + def publish_failure(self) -> bool: + with self._lock: + if self._status != "pending": + return False + self._status = "failed" + self._signal.release() + return True + + def wait(self, timeout_seconds: float) -> bool: + return self._signal.acquire(timeout=max(0.0, timeout_seconds)) + + def result(self) -> AcquiredRAGRetrievalConnection: + with self._lock: + if self._status == "succeeded" and self._acquired is not None: + return self._acquired + if self._status == "failed": + raise RAGRetrievalConnectionAcquisitionError() + raise RAGRetrievalConnectionAcquisitionTimeout() + + +class NativeThreadRAGRetrievalConnectionAcquirer: + """Bound checkout stalls without changing the shared SQLAlchemy pool.""" + + def __init__(self, *, max_workers: int) -> None: + if ( + isinstance(max_workers, bool) + or not isinstance(max_workers, int) + or not 1 <= max_workers <= 20 + ): + raise RAGRetrievalConnectionAcquisitionError() + self._max_workers = max_workers + self._lock = _allocate_native_lock() + self._active_workers = 0 + + def acquire( + self, + *, + session_factory: Callable[[], object], + cancellation: RAGRetrievalCancellation, + deadline_ns: int, + monotonic_ns: Callable[[], int], + ) -> AcquiredRAGRetrievalConnection: + if ( + not callable(session_factory) + or not isinstance(cancellation, RAGRetrievalCancellation) + or isinstance(deadline_ns, bool) + or not isinstance(deadline_ns, int) + or not callable(monotonic_ns) + ): + raise RAGRetrievalConnectionAcquisitionError() + if cancellation.cancelled or deadline_ns <= monotonic_ns(): + raise RAGRetrievalConnectionAcquisitionTimeout() + if not self._reserve_worker(): + raise RAGRetrievalConnectionAcquisitionTimeout() + + state = _ConnectionAcquisitionState() + try: + _start_native_thread( + self._acquire_owned, + (state, session_factory), + ) + except BaseException: + self._release_worker() + raise RAGRetrievalConnectionAcquisitionError() from None + + while state.status == "pending": + if cancellation.cancelled: + if state.abandon(): + raise RAGRetrievalConnectionAcquisitionTimeout() + break + remaining_ns = deadline_ns - monotonic_ns() + if remaining_ns <= 0: + if state.abandon(): + raise RAGRetrievalConnectionAcquisitionTimeout() + break + state.wait( + min( + remaining_ns / _NANOSECONDS_PER_SECOND, + _CANCELLATION_POLL_SECONDS, + ) + ) + + return state.result() + + def _reserve_worker(self) -> bool: + with self._lock: + if self._active_workers >= self._max_workers: + return False + self._active_workers += 1 + return True + + def _release_worker(self) -> None: + with self._lock: + self._active_workers -= 1 + + def _acquire_owned( + self, + state: _ConnectionAcquisitionState, + session_factory: Callable[[], object], + ) -> None: + session = None + try: + if state.status == "abandoned": + return + session = session_factory() + if state.status == "abandoned": + return + connection = session.connection() + if state.publish_success( + AcquiredRAGRetrievalConnection( + session=session, + connection=connection, + ) + ): + session = None + except BaseException: + state.publish_failure() + finally: + try: + if session is not None: + self._cleanup_session(session) + finally: + self._release_worker() + + @staticmethod + def _cleanup_session(session: object) -> None: + try: + session.rollback() + except BaseException: + pass + try: + session.close() + except BaseException: + pass + + +_process_connection_acquirer = NativeThreadRAGRetrievalConnectionAcquirer( + max_workers=DEFAULT_RAG_FANOUT_MAX_WORKERS +) + + +def get_process_rag_retrieval_connection_acquirer( +) -> NativeThreadRAGRetrievalConnectionAcquirer: + return _process_connection_acquirer + + +__all__ = [ + "AcquiredRAGRetrievalConnection", + "NativeThreadRAGRetrievalConnectionAcquirer", + "RAGRetrievalConnectionAcquirer", + "RAGRetrievalConnectionAcquisitionError", + "RAGRetrievalConnectionAcquisitionTimeout", + "get_process_rag_retrieval_connection_acquirer", +] diff --git a/apps/workflow_engine/adapters/rag_retrieval_executor.py b/apps/workflow_engine/adapters/rag_retrieval_executor.py new file mode 100644 index 000000000..5f9bb4bcd --- /dev/null +++ b/apps/workflow_engine/adapters/rag_retrieval_executor.py @@ -0,0 +1,260 @@ +"""Native-thread executor adapter for a gevent-monkey-patched Workflow worker.""" + +from __future__ import annotations + +from typing import Callable, Generic, TypeVar + +import gevent +from gevent.monkey import get_original +from gevent.threadpool import ThreadPool + +from apps.workflow_engine.application.rag_retrieval_fanout import ( + DEFAULT_RAG_FANOUT_MAX_WORKERS, + RAGRetrievalFanoutConfigurationError, + RAGRetrievalFanoutError, + RAGRetrievalJob, +) + + +T = TypeVar("T") +_allocate_native_lock = get_original("_thread", "allocate_lock") +_native_get_ident = get_original("_thread", "get_ident") + +PROCESS_RAG_RETRIEVAL_MAX_WORKERS = DEFAULT_RAG_FANOUT_MAX_WORKERS +PROCESS_RAG_RETRIEVAL_CONTROL_MAX_WORKERS = 2 +_IDLE_TASK_TIMEOUT_SECONDS = 0.1 + +_process_pool_lock = _allocate_native_lock() +_process_data_admission_lock = _allocate_native_lock() +_process_data_pool: ThreadPool | None = None +_process_control_pool: ThreadPool | None = None +_process_admitted_data_jobs = 0 + + +def _reserve_process_data_job() -> bool: + global _process_admitted_data_jobs + with _process_data_admission_lock: + if _process_admitted_data_jobs >= PROCESS_RAG_RETRIEVAL_MAX_WORKERS: + return False + _process_admitted_data_jobs += 1 + return True + + +def _release_process_data_job() -> None: + global _process_admitted_data_jobs + with _process_data_admission_lock: + _process_admitted_data_jobs -= 1 + + +def _get_process_data_pool() -> ThreadPool: + global _process_data_pool + with _process_pool_lock: + if _process_data_pool is None: + _process_data_pool = ThreadPool( + PROCESS_RAG_RETRIEVAL_MAX_WORKERS, + idle_task_timeout=_IDLE_TASK_TIMEOUT_SECONDS, + ) + return _process_data_pool + + +def _get_process_control_pool() -> ThreadPool: + global _process_control_pool + with _process_pool_lock: + if _process_control_pool is None: + _process_control_pool = ThreadPool( + PROCESS_RAG_RETRIEVAL_CONTROL_MAX_WORKERS, + idle_task_timeout=_IDLE_TASK_TIMEOUT_SECONDS, + ) + return _process_control_pool + + +class _RegisteredCancellationCallback: + """Keep an asynchronously dispatched callback inside its resource lifetime.""" + + def __init__(self, callback: Callable[[], None]) -> None: + self._lock = _allocate_native_lock() + self._completion = _allocate_native_lock() + self._completion.acquire() + self._callback: Callable[[], None] | None = callback + self._active = True + self._running = False + + def invoke(self) -> None: + with self._lock: + callback = self._callback + if not self._active or callback is None: + return + self._active = False + self._running = True + self._callback = None + try: + callback() + finally: + with self._lock: + self._running = False + self._completion.release() + + def deactivate(self) -> None: + with self._lock: + self._active = False + self._callback = None + running = self._running + if running: + self._completion.acquire() + self._completion.release() + + +class NativeThreadRAGRetrievalCancellation: + """Cancellation registry safe between the gevent hub and native workers.""" + + def __init__(self) -> None: + self._lock = _allocate_native_lock() + self._cancelled = False + self._next_token = 0 + self._callbacks: dict[int, _RegisteredCancellationCallback] = {} + self._owner_thread_id = _native_get_ident() + self._control_pool = _get_process_control_pool() + + @property + def cancelled(self) -> bool: + with self._lock: + return self._cancelled + + def register(self, callback: Callable[[], None]) -> Callable[[], None]: + if not callable(callback): + raise RAGRetrievalFanoutConfigurationError() + registered = _RegisteredCancellationCallback(callback) + with self._lock: + if self._cancelled: + token = None + else: + token = self._next_token + self._next_token += 1 + self._callbacks[token] = registered + if token is None: + self._invoke_or_dispatch(registered.invoke) + + def unregister() -> None: + if token is not None: + with self._lock: + self._callbacks.pop(token, None) + registered.deactivate() + + return unregister + + def cancel(self) -> None: + with self._lock: + if self._cancelled: + return + self._cancelled = True + callbacks = tuple(self._callbacks.values()) + self._callbacks.clear() + for callback in callbacks: + self._invoke_or_dispatch(callback.invoke) + + def _invoke_or_dispatch(self, callback: Callable[[], None]) -> None: + if _native_get_ident() != self._owner_thread_id: + self._invoke_safely(callback) + return + try: + self._control_pool.apply_async( + self._invoke_safely, + args=(callback,), + ) + except Exception: + return + + @staticmethod + def _invoke_safely(callback: Callable[[], None]) -> None: + try: + callback() + except Exception: + return + + +class GeventNativeThreadJob(Generic[T]): + def __init__(self, result) -> None: + self._result = result + + def ready(self) -> bool: + return self._result.ready() + + def result(self) -> T: + return self._result.get() + + +class GeventNativeThreadRAGRetrievalExecutor: + """Coordinate one invocation through the process-wide native data pool.""" + + def __init__(self, max_workers: int) -> None: + if ( + isinstance(max_workers, bool) + or not isinstance(max_workers, int) + or not 1 <= max_workers <= 20 + ): + raise RAGRetrievalFanoutConfigurationError() + self._pool = _get_process_data_pool() + self._max_workers = max_workers + self._state_lock = _allocate_native_lock() + self._active_jobs = 0 + self._closed = False + + def submit(self, callback: Callable[[], T]) -> RAGRetrievalJob[T]: + if not callable(callback): + raise RAGRetrievalFanoutConfigurationError() + self._reserve_invocation_job() + if not _reserve_process_data_job(): + self._release_invocation_job() + raise RAGRetrievalFanoutError() + try: + return GeventNativeThreadJob(gevent.spawn(self._run_admitted, callback)) + except BaseException: + _release_process_data_job() + self._release_invocation_job() + raise + + def _run_admitted(self, callback: Callable[[], T]) -> T: + try: + return self._pool.apply(callback) + finally: + _release_process_data_job() + self._release_invocation_job() + + def _reserve_invocation_job(self) -> None: + with self._state_lock: + if self._closed: + raise RAGRetrievalFanoutConfigurationError() + if self._active_jobs >= self._max_workers: + raise RAGRetrievalFanoutError() + self._active_jobs += 1 + + def _release_invocation_job(self) -> None: + with self._state_lock: + self._active_jobs -= 1 + + def wait( + self, + jobs: tuple[RAGRetrievalJob[object], ...], + *, + timeout_seconds: float, + ) -> None: + if self._closed: + return + raw_results = tuple( + job._result for job in jobs if isinstance(job, GeventNativeThreadJob) + ) + if raw_results: + gevent.wait(raw_results, timeout=max(0.0, timeout_seconds), count=1) + + def close(self) -> None: + with self._state_lock: + self._closed = True + + +__all__ = [ + "PROCESS_RAG_RETRIEVAL_CONTROL_MAX_WORKERS", + "PROCESS_RAG_RETRIEVAL_MAX_WORKERS", + "GeventNativeThreadJob", + "GeventNativeThreadRAGRetrievalExecutor", + "NativeThreadRAGRetrievalCancellation", +] diff --git a/apps/workflow_engine/adapters/rag_retrieval_session.py b/apps/workflow_engine/adapters/rag_retrieval_session.py new file mode 100644 index 000000000..65112fb02 --- /dev/null +++ b/apps/workflow_engine/adapters/rag_retrieval_session.py @@ -0,0 +1,250 @@ +"""SQLAlchemy session boundary for one blocking RAG retrieval task.""" + +from __future__ import annotations + +import time +from typing import Callable, TypeVar + +from sqlalchemy import event, text +from sqlalchemy.engine import Connection + +from apps.workflow_engine.adapters.rag_retrieval_connection_acquirer import ( + RAGRetrievalConnectionAcquirer, + get_process_rag_retrieval_connection_acquirer, +) +from apps.workflow_engine.application.rag_retrieval_fanout import ( + RAGRetrievalCancellation, +) + + +T = TypeVar("T") + + +def _noop() -> None: + return None + + +class RAGRetrievalSessionError(RuntimeError): + """Redaction-safe retrieval session failure.""" + + code = "rag.retrieval_session_failed" + + def __init__(self) -> None: + super().__init__("RAG retrieval session failed.") + + +class RAGRetrievalSessionTimeout(TimeoutError): + """Redaction-safe timeout raised after PostgreSQL query cancellation.""" + + code = "rag.retrieval_timed_out" + + def __init__(self) -> None: + super().__init__("RAG retrieval timed out.") + + +class _StatementDeadlineGuard: + _NANOSECONDS_PER_MILLISECOND = 1_000_000 + + def __init__( + self, + *, + cancellation: RAGRetrievalCancellation, + deadline_ns: int, + monotonic_ns: Callable[[], int], + ) -> None: + self._cancellation = cancellation + self._deadline_ns = deadline_ns + self._monotonic_ns = monotonic_ns + + def remaining_timeout_ms(self) -> int: + if self._cancellation.cancelled: + raise RAGRetrievalSessionTimeout() + remaining_ns = self._deadline_ns - self._monotonic_ns() + remaining_ms = remaining_ns // self._NANOSECONDS_PER_MILLISECOND + if remaining_ms < 1: + raise RAGRetrievalSessionTimeout() + return remaining_ms + + def __call__( + self, + connection, + cursor, + _statement, + _parameters, + _context, + _executemany, + ) -> None: + remaining_timeout_ms = self.remaining_timeout_ms() + dialect = getattr(getattr(connection, "dialect", None), "name", None) + if dialect == "postgresql": + cursor.execute( + "SELECT set_config('statement_timeout', %s, true)", + (f"{remaining_timeout_ms}ms",), + ) + self.remaining_timeout_ms() + + +class RAGRetrievalSessionRunner: + MAX_STATEMENT_TIMEOUT_MS = 30_000 + + def __init__( + self, + *, + session_factory: Callable[[], object], + monotonic_ns: Callable[[], int] | None = None, + connection_acquirer: RAGRetrievalConnectionAcquirer | None = None, + ) -> None: + if not callable(session_factory): + raise RAGRetrievalSessionError() + self._session_factory = session_factory + self._monotonic_ns = monotonic_ns or time.monotonic_ns + self._connection_acquirer = ( + connection_acquirer or get_process_rag_retrieval_connection_acquirer() + ) + + def run( + self, + *, + timeout_ms: int, + cancellation: RAGRetrievalCancellation, + operation: Callable[[object], T], + ) -> T: + if ( + isinstance(timeout_ms, bool) + or not isinstance(timeout_ms, int) + or not 1 <= timeout_ms <= self.MAX_STATEMENT_TIMEOUT_MS + or not isinstance(cancellation, RAGRetrievalCancellation) + or not callable(operation) + ): + raise RAGRetrievalSessionError() + + deadline_ns = ( + self._monotonic_ns() + + timeout_ms * _StatementDeadlineGuard._NANOSECONDS_PER_MILLISECOND + ) + deadline_guard = _StatementDeadlineGuard( + cancellation=cancellation, + deadline_ns=deadline_ns, + monotonic_ns=self._monotonic_ns, + ) + + session = None + unregister: Callable[[], None] = _noop + unregister_statement_guard: Callable[[], None] = _noop + result: T | None = None + failure: Exception | None = None + cleanup_failed = False + try: + acquired = self._connection_acquirer.acquire( + session_factory=self._session_factory, + cancellation=cancellation, + deadline_ns=deadline_ns, + monotonic_ns=self._monotonic_ns, + ) + session = acquired.session + connection = acquired.connection + cancel = self._driver_cancel_callback(connection) + if cancel is not None: + unregister = cancellation.register(cancel) + deadline_guard.remaining_timeout_ms() + session.execute(text("SET TRANSACTION READ ONLY")) + remaining_timeout_ms = deadline_guard.remaining_timeout_ms() + session.execute( + text( + "SELECT set_config('statement_timeout', :statement_timeout, true)" + ), + {"statement_timeout": f"{remaining_timeout_ms}ms"}, + ) + unregister_statement_guard = self._install_statement_guard( + connection, + deadline_guard, + ) + deadline_guard.remaining_timeout_ms() + result = operation(session) + deadline_guard.remaining_timeout_ms() + except Exception as exc: + failure = self._safe_failure(exc, cancellation=cancellation) + finally: + try: + unregister() + except Exception: + cleanup_failed = True + try: + unregister_statement_guard() + except Exception: + cleanup_failed = True + if session is not None: + try: + session.rollback() + except Exception: + cleanup_failed = True + try: + session.close() + except Exception: + cleanup_failed = True + + if failure is not None: + raise failure from None + if cleanup_failed: + raise RAGRetrievalSessionError() from None + return result # type: ignore[return-value] + + @staticmethod + def _install_statement_guard( + connection, + guard: _StatementDeadlineGuard, + ) -> Callable[[], None]: + if not isinstance(connection, Connection): + return _noop + event.listen(connection, "before_cursor_execute", guard) + + def unregister() -> None: + event.remove(connection, "before_cursor_execute", guard) + + return unregister + + @staticmethod + def _driver_cancel_callback(connection) -> Callable[[], None] | None: + proxy = getattr(connection, "connection", None) + driver = getattr(proxy, "driver_connection", None) + if driver is None: + driver = getattr(proxy, "connection", None) + cancel = getattr(driver, "cancel", None) + return cancel if callable(cancel) else None + + @classmethod + def _safe_failure( + cls, + exc: Exception, + *, + cancellation: RAGRetrievalCancellation, + ) -> Exception: + if ( + isinstance(exc, (RAGRetrievalSessionTimeout, TimeoutError)) + or cancellation.cancelled + or cls._is_postgres_query_cancel(exc) + ): + return RAGRetrievalSessionTimeout() + return RAGRetrievalSessionError() + + @staticmethod + def _is_postgres_query_cancel(exc: Exception) -> bool: + current: object = exc + for _index in range(3): + sqlstate = getattr(current, "sqlstate", None) or getattr( + current, "pgcode", None + ) + if sqlstate == "57014": + return True + nested = getattr(current, "orig", None) + if nested is None or nested is current: + break + current = nested + return False + + +__all__ = [ + "RAGRetrievalSessionError", + "RAGRetrievalSessionRunner", + "RAGRetrievalSessionTimeout", +] diff --git a/apps/workflow_engine/application/rag_retrieval_fanout.py b/apps/workflow_engine/application/rag_retrieval_fanout.py new file mode 100644 index 000000000..cf03c226a --- /dev/null +++ b/apps/workflow_engine/application/rag_retrieval_fanout.py @@ -0,0 +1,477 @@ +"""Bounded application scheduler for blocking per-KB retrieval work.""" + +from __future__ import annotations + +import math +import time +from dataclasses import dataclass, field +from typing import Callable, Generic, Protocol, TypeVar, runtime_checkable + + +T = TypeVar("T") +DEFAULT_RAG_FANOUT_MAX_WORKERS = 5 + + +class RAGRetrievalFanoutConfigurationError(ValueError): + """Redaction-safe configuration failure for the retrieval scheduler.""" + + code = "rag.retrieval_fanout_configuration_invalid" + + +class RAGRetrievalFanoutError(RuntimeError): + """Redaction-safe terminal failure for fail-fast retrieval.""" + + code = "rag.retrieval_fanout_failed" + + def __init__(self) -> None: + super().__init__("RAG retrieval fan-out failed.") + + +@dataclass(frozen=True, slots=True) +class RAGRetrievalFanoutTask: + ordinal: int + resource_ref: str = field(repr=False) + + def __post_init__(self) -> None: + if ( + isinstance(self.ordinal, bool) + or not isinstance(self.ordinal, int) + or self.ordinal < 0 + or not isinstance(self.resource_ref, str) + or not self.resource_ref + ): + raise RAGRetrievalFanoutConfigurationError() + + +@runtime_checkable +class RAGRetrievalCancellation(Protocol): + @property + def cancelled(self) -> bool: ... + + def register(self, callback: Callable[[], None]) -> Callable[[], None]: ... + + def cancel(self) -> None: ... + + +class RAGRetrievalJob(Protocol, Generic[T]): + def ready(self) -> bool: ... + + def result(self) -> T: ... + + +class RAGRetrievalExecutor(Protocol): + def submit(self, callback: Callable[[], T]) -> RAGRetrievalJob[T]: ... + + def wait( + self, + jobs: tuple[RAGRetrievalJob[object], ...], + *, + timeout_seconds: float, + ) -> None: ... + + def close(self) -> None: ... + + +RAGRetrievalExecutorFactory = Callable[[int], RAGRetrievalExecutor] +RAGRetrievalCancellationFactory = Callable[[], RAGRetrievalCancellation] +RAGRetrievalWorker = Callable[ + [RAGRetrievalFanoutTask, RAGRetrievalCancellation, int], + T, +] + + +@dataclass(frozen=True, slots=True) +class RAGRetrievalFanoutResult(Generic[T]): + results: tuple[tuple[RAGRetrievalFanoutTask, T], ...] = field(repr=False) + failed_count: int + timeout_count: int + slowest_search_latency_ms: int | None + + +@dataclass(slots=True) +class _TaskState: + task: RAGRetrievalFanoutTask + cancellation: RAGRetrievalCancellation + started_at: float | None = None + finished_at: float | None = None + budget_seconds: float | None = None + + +@dataclass(frozen=True, slots=True) +class _TerminalOutcome(Generic[T]): + task: RAGRetrievalFanoutTask + value: T | None = field(default=None, repr=False) + failed: bool = False + timed_out: bool = False + elapsed_ms: int | None = None + + +class _TaskStartDeadlineExceeded(TimeoutError): + pass + + +class RAGRetrievalFanoutScheduler: + """Coordinate authorized blocking searches through an injected executor port.""" + + MAX_TASKS = 20 + _POLL_INTERVAL_SECONDS = 0.01 + + def __init__( + self, + *, + executor_factory: RAGRetrievalExecutorFactory, + cancellation_factory: RAGRetrievalCancellationFactory, + max_workers: int = DEFAULT_RAG_FANOUT_MAX_WORKERS, + per_task_timeout_seconds: float = 10.0, + aggregate_timeout_seconds: float = 30.0, + cleanup_reserve_seconds: float = 1.0, + minimum_start_budget_ms: int = 250, + clock: Callable[[], float] = time.monotonic, + ) -> None: + if ( + not callable(executor_factory) + or not callable(cancellation_factory) + or isinstance(max_workers, bool) + or not isinstance(max_workers, int) + or not 1 <= max_workers <= self.MAX_TASKS + or not self._valid_positive_number(per_task_timeout_seconds) + or not self._valid_positive_number(aggregate_timeout_seconds) + or not self._valid_positive_number(cleanup_reserve_seconds) + or cleanup_reserve_seconds >= aggregate_timeout_seconds + or isinstance(minimum_start_budget_ms, bool) + or not isinstance(minimum_start_budget_ms, int) + or minimum_start_budget_ms < 1 + or minimum_start_budget_ms + >= (aggregate_timeout_seconds - cleanup_reserve_seconds) * 1000 + or not callable(clock) + ): + raise RAGRetrievalFanoutConfigurationError() + self._executor_factory = executor_factory + self._cancellation_factory = cancellation_factory + self._max_workers = max_workers + self._per_task_timeout_seconds = float(per_task_timeout_seconds) + self._aggregate_timeout_seconds = float(aggregate_timeout_seconds) + self._cleanup_reserve_seconds = float(cleanup_reserve_seconds) + self._minimum_start_budget_ms = minimum_start_budget_ms + self._clock = clock + + def execute( + self, + *, + tasks: tuple[RAGRetrievalFanoutTask, ...], + worker: RAGRetrievalWorker[T], + fail_fast: bool = False, + ) -> RAGRetrievalFanoutResult[T]: + self._validate_request(tasks=tasks, worker=worker, fail_fast=fail_fast) + if not tasks: + return RAGRetrievalFanoutResult((), 0, 0, None) + + invocation_started_at = self._clock() + hard_deadline = invocation_started_at + self._aggregate_timeout_seconds + search_deadline = hard_deadline - self._cleanup_reserve_seconds + try: + stop_signal = self._cancellation_factory() + queued = [ + _TaskState( + task=task, + cancellation=self._cancellation_factory(), + ) + for task in tasks + ] + except Exception: + raise RAGRetrievalFanoutError() from None + if not isinstance(stop_signal, RAGRetrievalCancellation) or any( + not isinstance(state.cancellation, RAGRetrievalCancellation) + for state in queued + ): + raise RAGRetrievalFanoutError() from None + running: list[tuple[RAGRetrievalJob[T], _TaskState]] = [] + terminal: dict[int, _TerminalOutcome[T]] = {} + fail_fast_triggered = False + + try: + executor = self._executor_factory(min(self._max_workers, len(tasks))) + except Exception: + raise RAGRetrievalFanoutError() from None + + def invoke(state: _TaskState) -> T: + if stop_signal.cancelled: + raise _TaskStartDeadlineExceeded() + now = self._clock() + remaining_seconds = search_deadline - now + if remaining_seconds * 1000 <= self._minimum_start_budget_ms: + raise _TaskStartDeadlineExceeded() + budget_seconds = min(self._per_task_timeout_seconds, remaining_seconds) + if stop_signal.cancelled: + raise _TaskStartDeadlineExceeded() + state.started_at = now + state.budget_seconds = budget_seconds + try: + return worker( + state.task, + state.cancellation, + max(1, int(budget_seconds * 1000)), + ) + finally: + state.finished_at = self._clock() + + try: + while (queued or running) and not fail_fast_triggered: + fail_fast_triggered = self._consume_ready( + running=running, + terminal=terminal, + fail_fast=fail_fast, + ) + if fail_fast_triggered: + break + + now = self._clock() + for _job, state in tuple(running): + if state.task.ordinal in terminal: + continue + if ( + state.started_at is not None + and state.budget_seconds is not None + and now >= state.started_at + state.budget_seconds + ): + state.cancellation.cancel() + terminal[state.task.ordinal] = _TerminalOutcome( + task=state.task, + failed=True, + timed_out=True, + elapsed_ms=self._elapsed_ms(state, now), + ) + if fail_fast: + fail_fast_triggered = True + break + if fail_fast_triggered: + break + + now = self._clock() + if now >= search_deadline: + break + while queued and len(running) < self._max_workers: + remaining_ms = int((search_deadline - self._clock()) * 1000) + if remaining_ms <= self._minimum_start_budget_ms: + break + state = queued.pop(0) + try: + job = executor.submit(lambda state=state: invoke(state)) + except Exception: + terminal[state.task.ordinal] = _TerminalOutcome( + task=state.task, + failed=True, + ) + if fail_fast: + fail_fast_triggered = True + break + continue + running.append((job, state)) + if fail_fast_triggered: + break + if not running: + break + executor.wait( + tuple(job for job, _state in running), + timeout_seconds=min( + self._POLL_INTERVAL_SECONDS, + max(0.0, search_deadline - self._clock()), + ), + ) + + if fail_fast_triggered: + stop_signal.cancel() + for _job, state in running: + if state.task.ordinal not in terminal: + state.cancellation.cancel() + else: + stop_signal.cancel() + now = self._clock() + for state in queued: + terminal[state.task.ordinal] = _TerminalOutcome( + task=state.task, + failed=True, + timed_out=True, + ) + queued.clear() + for _job, state in running: + if state.task.ordinal in terminal: + continue + state.cancellation.cancel() + terminal[state.task.ordinal] = _TerminalOutcome( + task=state.task, + failed=True, + timed_out=True, + elapsed_ms=self._elapsed_ms(state, now), + ) + + while running and self._clock() < hard_deadline: + self._discard_ready(running) + if not running: + break + executor.wait( + tuple(job for job, _state in running), + timeout_seconds=min( + self._POLL_INTERVAL_SECONDS, + max(0.0, hard_deadline - self._clock()), + ), + ) + except Exception: + self._cancel_safely(stop_signal) + for _job, state in running: + self._cancel_safely(state.cancellation) + raise RAGRetrievalFanoutError() from None + finally: + self._cancel_safely(stop_signal) + for _job, state in running: + self._cancel_safely(state.cancellation) + try: + executor.close() + except Exception: + pass + + if fail_fast_triggered: + raise RAGRetrievalFanoutError() + + ordered_outcomes = tuple(terminal[index] for index in sorted(terminal)) + results = tuple( + (outcome.task, outcome.value) + for outcome in ordered_outcomes + if not outcome.failed and not outcome.timed_out + ) + elapsed_values = tuple( + outcome.elapsed_ms + for outcome in ordered_outcomes + if outcome.elapsed_ms is not None + ) + return RAGRetrievalFanoutResult( + results=results, + failed_count=sum(outcome.failed for outcome in ordered_outcomes), + timeout_count=sum(outcome.timed_out for outcome in ordered_outcomes), + slowest_search_latency_ms=max(elapsed_values, default=None), + ) + + def _consume_ready( + self, + *, + running: list[tuple[RAGRetrievalJob[T], _TaskState]], + terminal: dict[int, _TerminalOutcome[T]], + fail_fast: bool, + ) -> bool: + for job, state in tuple(running): + if not job.ready(): + continue + running.remove((job, state)) + if state.task.ordinal in terminal: + self._consume_ignored_job(job) + continue + elapsed_ms = self._elapsed_ms(state, self._clock()) + try: + value = job.result() + except (_TaskStartDeadlineExceeded, TimeoutError): + terminal[state.task.ordinal] = _TerminalOutcome( + task=state.task, + failed=True, + timed_out=True, + elapsed_ms=elapsed_ms, + ) + if fail_fast: + return True + except Exception: + terminal[state.task.ordinal] = _TerminalOutcome( + task=state.task, + failed=True, + elapsed_ms=elapsed_ms, + ) + if fail_fast: + return True + else: + overrun = ( + state.started_at is not None + and state.finished_at is not None + and state.budget_seconds is not None + and state.finished_at > state.started_at + state.budget_seconds + ) + terminal[state.task.ordinal] = _TerminalOutcome( + task=state.task, + value=None if overrun else value, + failed=overrun, + timed_out=overrun, + elapsed_ms=elapsed_ms, + ) + if overrun and fail_fast: + return True + return False + + @classmethod + def _discard_ready( + cls, + running: list[tuple[RAGRetrievalJob[T], _TaskState]], + ) -> None: + for job, state in tuple(running): + if not job.ready(): + continue + running.remove((job, state)) + cls._consume_ignored_job(job) + + @staticmethod + def _consume_ignored_job(job: RAGRetrievalJob[object]) -> None: + try: + job.result() + except Exception: + return + + @staticmethod + def _cancel_safely(cancellation: RAGRetrievalCancellation) -> None: + try: + cancellation.cancel() + except Exception: + return + + @staticmethod + def _elapsed_ms(state: _TaskState, fallback_end: float) -> int | None: + if state.started_at is None: + return None + end = state.finished_at if state.finished_at is not None else fallback_end + return max(0, int((end - state.started_at) * 1000)) + + def _validate_request( + self, + *, + tasks: tuple[RAGRetrievalFanoutTask, ...], + worker: RAGRetrievalWorker[T], + fail_fast: bool, + ) -> None: + if ( + not isinstance(tasks, tuple) + or len(tasks) > self.MAX_TASKS + or any(not isinstance(task, RAGRetrievalFanoutTask) for task in tasks) + or len({task.ordinal for task in tasks}) != len(tasks) + or not callable(worker) + or type(fail_fast) is not bool + ): + raise RAGRetrievalFanoutConfigurationError() + + @staticmethod + def _valid_positive_number(value: object) -> bool: + return ( + not isinstance(value, bool) + and isinstance(value, (int, float)) + and math.isfinite(value) + and value > 0 + ) + + +__all__ = [ + "DEFAULT_RAG_FANOUT_MAX_WORKERS", + "RAGRetrievalCancellation", + "RAGRetrievalCancellationFactory", + "RAGRetrievalExecutor", + "RAGRetrievalExecutorFactory", + "RAGRetrievalFanoutConfigurationError", + "RAGRetrievalFanoutError", + "RAGRetrievalFanoutResult", + "RAGRetrievalFanoutScheduler", + "RAGRetrievalFanoutTask", + "RAGRetrievalJob", +] diff --git a/apps/workflow_engine/services/retrieval.py b/apps/workflow_engine/services/retrieval.py index 30fd1f4dc..c5b7fe53e 100644 --- a/apps/workflow_engine/services/retrieval.py +++ b/apps/workflow_engine/services/retrieval.py @@ -2,6 +2,7 @@ import os import re +from gevent.monkey import get_original from sqlalchemy import and_, bindparam, or_, select from sqlalchemy.orm import Session, aliased @@ -46,6 +47,7 @@ RAG_RERANK_ENABLED_ENV = "RAG_CROSS_ENCODER_RERANK_ENABLED" RAG_RERANK_MODEL_ENV = "RAG_CROSS_ENCODER_MODEL" DEFAULT_RAG_RERANK_MODEL = "cross-encoder/ms-marco-MiniLM-L-12-v2" +_allocate_native_lock = get_original("_thread", "allocate_lock") def _env_flag(name: str, default: bool = False) -> bool: @@ -58,6 +60,7 @@ def _env_flag(name: str, default: bool = False) -> bool: class RetrievalService: _cross_encoder_model = None _cross_encoder_model_name = None + _cross_encoder_model_lock = _allocate_native_lock() def __init__(self, db: Session, user_id, organization_id=None): self.db = db @@ -78,11 +81,19 @@ def _get_cross_encoder_model(cls): model_name = DEFAULT_RAG_RERANK_MODEL if ( - cls._cross_encoder_model is None - or cls._cross_encoder_model_name != model_name + cls._cross_encoder_model is not None + and cls._cross_encoder_model_name == model_name ): - cls._cross_encoder_model = CrossEncoder(model_name, max_length=512) - cls._cross_encoder_model_name = model_name + return cls._cross_encoder_model + + with cls._cross_encoder_model_lock: + if ( + cls._cross_encoder_model is None + or cls._cross_encoder_model_name != model_name + ): + model = CrossEncoder(model_name, max_length=512) + cls._cross_encoder_model = model + cls._cross_encoder_model_name = model_name return cls._cross_encoder_model def _get_efficient_rewrite_model(self) -> str: @@ -812,9 +823,9 @@ async def search_documents( if item["score"] > all_candidates[chunk_id]["score"]: all_candidates[chunk_id] = item - except Exception as e: - logger.error(f"Search Failed: {e}") - raise e + except Exception as exc: + logger.error("Search failed: error_type=%s", type(exc).__name__) + raise final_list = [] merged_candidates = sorted( @@ -1037,6 +1048,9 @@ def search_documents_sync( if not knowledge_base_id: logger.error("Missing knowledge_base_id") return [] + if self.organization_id is None: + logger.error("Missing organization_id for synchronous retrieval") + return [] hierarchy_mode = normalize_hierarchy_mode(hierarchy_mode) source_tier_policy = normalize_source_tier_policy(source_tier_policy) @@ -1045,7 +1059,10 @@ def search_documents_sync( try: kb = ( self.db.query(KnowledgeBase) - .filter(KnowledgeBase.id == knowledge_base_id) + .filter( + KnowledgeBase.id == knowledge_base_id, + KnowledgeBase.organization_id == self.organization_id, + ) .first() ) if not kb or not kb.embedding_model: @@ -1222,9 +1239,9 @@ def search_documents_sync( if item["score"] > all_candidates[chunk_id]["score"]: all_candidates[chunk_id] = item - except Exception as e: - logger.error(f"Search Failed: {e}") - raise e + except Exception as exc: + logger.error("Search failed: error_type=%s", type(exc).__name__) + raise final_list = [] merged_candidates = sorted( diff --git a/apps/workflow_engine/tests/adapters/test_rag_retrieval_connection_acquirer.py b/apps/workflow_engine/tests/adapters/test_rag_retrieval_connection_acquirer.py new file mode 100644 index 000000000..e90ca7113 --- /dev/null +++ b/apps/workflow_engine/tests/adapters/test_rag_retrieval_connection_acquirer.py @@ -0,0 +1,204 @@ +import threading +import time + +import pytest + +from apps.workflow_engine.adapters.rag_retrieval_connection_acquirer import ( + NativeThreadRAGRetrievalConnectionAcquirer, + RAGRetrievalConnectionAcquisitionTimeout, +) +from apps.workflow_engine.adapters.rag_retrieval_executor import ( + NativeThreadRAGRetrievalCancellation, +) + + +class _Session: + def __init__(self) -> None: + self.connection_value = object() + self.rollback_count = 0 + self.close_count = 0 + self.closed = threading.Event() + + def connection(self): + return self.connection_value + + def rollback(self) -> None: + self.rollback_count += 1 + + def close(self) -> None: + self.close_count += 1 + self.closed.set() + + +def test_acquirer_returns_session_and_connection_within_deadline() -> None: + session = _Session() + acquirer = NativeThreadRAGRetrievalConnectionAcquirer(max_workers=1) + + acquired = acquirer.acquire( + session_factory=lambda: session, + cancellation=NativeThreadRAGRetrievalCancellation(), + deadline_ns=time.monotonic_ns() + 1_000_000_000, + monotonic_ns=time.monotonic_ns, + ) + + assert acquired.session is session + assert acquired.connection is session.connection_value + assert session.rollback_count == 0 + assert session.close_count == 0 + + +def test_acquirer_returns_at_deadline_and_late_session_is_cleaned_by_owner() -> None: + factory_started = threading.Event() + release_factory = threading.Event() + deadline_expired = threading.Event() + session = _Session() + acquirer = NativeThreadRAGRetrievalConnectionAcquirer(max_workers=1) + + def delayed_factory(): + factory_started.set() + release_factory.wait(timeout=1) + return session + + def expire_deadline(): + assert factory_started.wait(timeout=1) + deadline_expired.set() + + expiry = threading.Thread(target=expire_deadline) + expiry.start() + with pytest.raises(RAGRetrievalConnectionAcquisitionTimeout): + acquirer.acquire( + session_factory=delayed_factory, + cancellation=NativeThreadRAGRetrievalCancellation(), + deadline_ns=1_000_000_000, + monotonic_ns=lambda: ( + 2_000_000_000 if deadline_expired.is_set() else 0 + ), + ) + expiry.join(timeout=1) + + assert factory_started.is_set() + assert session.close_count == 0 + + release_factory.set() + assert session.closed.wait(timeout=1) + assert session.rollback_count == 1 + assert session.close_count == 1 + + +def test_acquirer_rejects_new_checkout_when_process_slots_are_occupied() -> None: + factory_started = threading.Event() + release_factory = threading.Event() + deadline_expired = threading.Event() + first_session = _Session() + second_factory_calls = 0 + acquirer = NativeThreadRAGRetrievalConnectionAcquirer(max_workers=1) + + def delayed_factory(): + factory_started.set() + release_factory.wait(timeout=1) + return first_session + + def expire_deadline(): + assert factory_started.wait(timeout=1) + deadline_expired.set() + + expiry = threading.Thread(target=expire_deadline) + expiry.start() + with pytest.raises(RAGRetrievalConnectionAcquisitionTimeout): + acquirer.acquire( + session_factory=delayed_factory, + cancellation=NativeThreadRAGRetrievalCancellation(), + deadline_ns=1_000_000_000, + monotonic_ns=lambda: ( + 2_000_000_000 if deadline_expired.is_set() else 0 + ), + ) + expiry.join(timeout=1) + + assert factory_started.is_set() + + def second_factory(): + nonlocal second_factory_calls + second_factory_calls += 1 + return _Session() + + with pytest.raises(RAGRetrievalConnectionAcquisitionTimeout): + acquirer.acquire( + session_factory=second_factory, + cancellation=NativeThreadRAGRetrievalCancellation(), + deadline_ns=time.monotonic_ns() + 1_000_000_000, + monotonic_ns=time.monotonic_ns, + ) + + assert second_factory_calls == 0 + + release_factory.set() + assert first_session.closed.wait(timeout=1) + + +def test_acquirer_rejects_cancelled_request_before_factory_call() -> None: + factory_calls = 0 + cancellation = NativeThreadRAGRetrievalCancellation() + cancellation.cancel() + acquirer = NativeThreadRAGRetrievalConnectionAcquirer(max_workers=1) + + def session_factory(): + nonlocal factory_calls + factory_calls += 1 + return _Session() + + with pytest.raises(RAGRetrievalConnectionAcquisitionTimeout): + acquirer.acquire( + session_factory=session_factory, + cancellation=cancellation, + deadline_ns=time.monotonic_ns() + 1_000_000_000, + monotonic_ns=time.monotonic_ns, + ) + + assert factory_calls == 0 + + +def test_acquirer_cancellation_releases_caller_and_late_session_is_cleaned() -> None: + factory_started = threading.Event() + release_factory = threading.Event() + acquisition_finished = threading.Event() + session = _Session() + cancellation = NativeThreadRAGRetrievalCancellation() + acquirer = NativeThreadRAGRetrievalConnectionAcquirer(max_workers=1) + failures = [] + + def delayed_factory(): + factory_started.set() + release_factory.wait(timeout=1) + return session + + def acquire(): + try: + acquirer.acquire( + session_factory=delayed_factory, + cancellation=cancellation, + deadline_ns=time.monotonic_ns() + 1_000_000_000, + monotonic_ns=time.monotonic_ns, + ) + except Exception as exc: # pragma: no cover - asserted below + failures.append(exc) + finally: + acquisition_finished.set() + + caller = threading.Thread(target=acquire) + caller.start() + assert factory_started.wait(timeout=1) + + cancellation.cancel() + + assert acquisition_finished.wait(timeout=0.5) + caller.join(timeout=0.5) + assert caller.is_alive() is False + assert len(failures) == 1 + assert isinstance(failures[0], RAGRetrievalConnectionAcquisitionTimeout) + assert session.close_count == 0 + + release_factory.set() + assert session.closed.wait(timeout=1) + assert session.rollback_count == 1 + assert session.close_count == 1 diff --git a/apps/workflow_engine/tests/adapters/test_rag_retrieval_executor.py b/apps/workflow_engine/tests/adapters/test_rag_retrieval_executor.py new file mode 100644 index 000000000..742c4362d --- /dev/null +++ b/apps/workflow_engine/tests/adapters/test_rag_retrieval_executor.py @@ -0,0 +1,299 @@ +import os +import subprocess +import sys +import textwrap +from pathlib import Path + + +def _assert_isolated_script_succeeds(script: str) -> None: + repository_root = Path(__file__).resolve().parents[4] + environment = os.environ.copy() + environment["PYTHONPATH"] = str(repository_root) + + completed = subprocess.run( + [sys.executable, "-c", textwrap.dedent(script)], + cwd=repository_root, + env=environment, + capture_output=True, + text=True, + timeout=10, + check=False, + ) + + assert completed.returncode == 0 + + +def test_executor_uses_native_thread_after_gevent_monkey_patch() -> None: + _assert_isolated_script_succeeds( + """ + from gevent import monkey + monkey.patch_all() + + import gevent + import time + + allocate_native_lock = monkey.get_original('_thread', 'allocate_lock') + original_get_ident = monkey.get_original('_thread', 'get_ident') + original_sleep = monkey.get_original('time', 'sleep') + from apps.workflow_engine.adapters.rag_retrieval_executor import ( + GeventNativeThreadRAGRetrievalExecutor, + NativeThreadRAGRetrievalCancellation, + ) + + main_thread_id = original_get_ident() + executor = GeventNativeThreadRAGRetrievalExecutor(1) + try: + worker_thread_id = executor.submit(original_get_ident).result() + cancellation = NativeThreadRAGRetrievalCancellation() + registered = allocate_native_lock() + registered.acquire() + observed = allocate_native_lock() + observed.acquire() + + def wait_for_cancellation(): + unregister = cancellation.register(observed.release) + registered.release() + deadline = time.monotonic() + 1 + try: + while not cancellation.cancelled and time.monotonic() < deadline: + original_sleep(0.001) + if observed.acquire(timeout=1): + observed.release() + finally: + unregister() + + job = executor.submit(wait_for_cancellation) + gevent.sleep(0) + if not registered.acquire(timeout=1): + raise SystemExit(2) + cancellation.cancel() + job.result() + if not observed.acquire(timeout=1): + raise SystemExit(3) + finally: + executor.close() + if worker_thread_id == main_thread_id: + raise SystemExit(1) + """ + ) + + +def test_executors_share_process_wide_native_worker_cap() -> None: + _assert_isolated_script_succeeds( + """ + from gevent import monkey + monkey.patch_all() + + import gevent + + allocate_native_lock = monkey.get_original('_thread', 'allocate_lock') + from apps.workflow_engine.adapters.rag_retrieval_executor import ( + PROCESS_RAG_RETRIEVAL_MAX_WORKERS, + GeventNativeThreadRAGRetrievalExecutor, + ) + + gate = allocate_native_lock() + gate.acquire() + state_lock = allocate_native_lock() + state = {'active': 0, 'max_active': 0} + + def work(): + with state_lock: + state['active'] += 1 + state['max_active'] = max( + state['max_active'], + state['active'], + ) + try: + gate.acquire() + gate.release() + return True + finally: + with state_lock: + state['active'] -= 1 + + first = GeventNativeThreadRAGRetrievalExecutor( + PROCESS_RAG_RETRIEVAL_MAX_WORKERS + ) + second = GeventNativeThreadRAGRetrievalExecutor( + PROCESS_RAG_RETRIEVAL_MAX_WORKERS + ) + first_job_count = PROCESS_RAG_RETRIEVAL_MAX_WORKERS // 2 + jobs = [ + first.submit(work) + for _ in range(first_job_count) + ] + [ + second.submit(work) + for _ in range( + PROCESS_RAG_RETRIEVAL_MAX_WORKERS - first_job_count + ) + ] + gevent.spawn_later(0.1, gate.release) + try: + if not all(job.result() for job in jobs): + raise SystemExit(2) + finally: + first.close() + second.close() + + if state['max_active'] > PROCESS_RAG_RETRIEVAL_MAX_WORKERS: + raise SystemExit(1) + """ + ) + + +def test_process_pool_rejects_overflow_before_spawning_and_recovers() -> None: + _assert_isolated_script_succeeds( + """ + from gevent import monkey + monkey.patch_all() + + allocate_native_lock = monkey.get_original('_thread', 'allocate_lock') + from apps.workflow_engine.adapters.rag_retrieval_executor import ( + PROCESS_RAG_RETRIEVAL_MAX_WORKERS, + GeventNativeThreadRAGRetrievalExecutor, + ) + from apps.workflow_engine.application.rag_retrieval_fanout import ( + RAGRetrievalFanoutError, + ) + + gate = allocate_native_lock() + gate.acquire() + first = GeventNativeThreadRAGRetrievalExecutor( + PROCESS_RAG_RETRIEVAL_MAX_WORKERS + ) + + def blocked_work(): + gate.acquire() + gate.release() + return True + + jobs = [ + first.submit(blocked_work) + for _ in range(PROCESS_RAG_RETRIEVAL_MAX_WORKERS) + ] + + overflow_callback_calls = [] + + def overflow_callback(): + overflow_callback_calls.append(True) + return True + + overflow = GeventNativeThreadRAGRetrievalExecutor(1) + try: + try: + overflow.submit(overflow_callback) + except RAGRetrievalFanoutError: + pass + else: + raise SystemExit(1) + if overflow_callback_calls: + raise SystemExit(2) + + gate.release() + if not all(job.result() for job in jobs): + raise SystemExit(3) + + recovered = overflow.submit(lambda: True) + if not recovered.result(): + raise SystemExit(4) + finally: + first.close() + overflow.close() + """ + ) + + +def test_cancellation_callbacks_do_not_block_gevent_hub() -> None: + _assert_isolated_script_succeeds( + """ + from gevent import monkey + monkey.patch_all() + + import gevent + import time + + allocate_native_lock = monkey.get_original('_thread', 'allocate_lock') + original_get_ident = monkey.get_original('_thread', 'get_ident') + original_sleep = monkey.get_original('time', 'sleep') + from apps.workflow_engine.adapters.rag_retrieval_executor import ( + GeventNativeThreadRAGRetrievalExecutor, + NativeThreadRAGRetrievalCancellation, + ) + + main_thread_id = original_get_ident() + callback_started = allocate_native_lock() + callback_started.acquire() + callback_blocker = allocate_native_lock() + callback_blocker.acquire() + callback_finished = allocate_native_lock() + callback_finished.acquire() + callback_thread_ids = [] + + def blocking_callback(): + callback_thread_ids.append(original_get_ident()) + callback_started.release() + callback_blocker.acquire() + callback_blocker.release() + callback_finished.release() + + def release_callback(): + if not callback_started.acquire(timeout=1): + return False + original_sleep(0.2) + callback_blocker.release() + return True + + executor = GeventNativeThreadRAGRetrievalExecutor(1) + cancellation = NativeThreadRAGRetrievalCancellation() + cancellation.register(blocking_callback) + helper_job = executor.submit(release_callback) + gevent.sleep(0) + + started_at = time.monotonic() + cancellation.cancel() + cancel_elapsed = time.monotonic() - started_at + gevent.sleep(0) + + try: + if not callback_finished.acquire(timeout=1): + raise SystemExit(2) + if not helper_job.result(): + raise SystemExit(3) + finally: + executor.close() + + if cancel_elapsed >= 0.1: + raise SystemExit(1) + if callback_thread_ids == [main_thread_id]: + raise SystemExit(4) + """ + ) + + +def test_unregister_disarms_queued_cancellation_callback() -> None: + _assert_isolated_script_succeeds( + """ + from gevent import monkey + monkey.patch_all() + + import gevent + + allocate_native_lock = monkey.get_original('_thread', 'allocate_lock') + from apps.workflow_engine.adapters.rag_retrieval_executor import ( + NativeThreadRAGRetrievalCancellation, + ) + + invoked = allocate_native_lock() + invoked.acquire() + cancellation = NativeThreadRAGRetrievalCancellation() + unregister = cancellation.register(invoked.release) + + cancellation.cancel() + unregister() + gevent.sleep(0.05) + + if invoked.acquire(blocking=False): + raise SystemExit(1) + """ + ) diff --git a/apps/workflow_engine/tests/adapters/test_rag_retrieval_session.py b/apps/workflow_engine/tests/adapters/test_rag_retrieval_session.py new file mode 100644 index 000000000..c69601007 --- /dev/null +++ b/apps/workflow_engine/tests/adapters/test_rag_retrieval_session.py @@ -0,0 +1,344 @@ +import threading +from types import SimpleNamespace + +import gevent +import pytest +from sqlalchemy import create_engine, text + +from apps.workflow_engine.adapters.rag_retrieval_session import ( + RAGRetrievalSessionError, + RAGRetrievalSessionRunner, + RAGRetrievalSessionTimeout, + _StatementDeadlineGuard, +) +from apps.workflow_engine.adapters.rag_retrieval_connection_acquirer import ( + RAGRetrievalConnectionAcquisitionTimeout, +) +from apps.workflow_engine.adapters.rag_retrieval_executor import ( + NativeThreadRAGRetrievalCancellation, +) + + +class _DriverConnection: + def __init__(self) -> None: + self.cancel_count = 0 + self.cancelled = threading.Event() + + def cancel(self) -> None: + self.cancel_count += 1 + self.cancelled.set() + + +class _FakeSession: + def __init__(self, *, rollback_error: bool = False, close_error: bool = False): + self.events: list[object] = [] + self.driver = _DriverConnection() + self._connection = SimpleNamespace( + connection=SimpleNamespace(driver_connection=self.driver) + ) + self.rollback_error = rollback_error + self.close_error = close_error + + def connection(self): + self.events.append("connection") + return self._connection + + def execute(self, statement, parameters=None): + self.events.append((str(statement), parameters)) + return SimpleNamespace() + + def rollback(self): + self.events.append("rollback") + if self.rollback_error: + raise RuntimeError("private-rollback-detail") + + def close(self): + self.events.append("close") + if self.close_error: + raise RuntimeError("private-close-detail") + + +class _SQLAlchemyBackedSession: + def __init__(self) -> None: + self.engine = create_engine("sqlite:///:memory:") + self.sql_connection = self.engine.connect() + self.events: list[object] = [] + + def connection(self): + self.events.append("connection") + return self.sql_connection + + def execute(self, statement, parameters=None): + rendered = str(statement) + self.events.append((rendered, parameters)) + if rendered.startswith("SET TRANSACTION") or "set_config" in rendered: + return SimpleNamespace() + return self.sql_connection.execute(statement, parameters or {}) + + def rollback(self): + self.events.append("rollback") + self.sql_connection.rollback() + + def close(self): + self.events.append("close") + + def dispose(self) -> None: + self.sql_connection.close() + self.engine.dispose() + + +def test_session_runner_applies_read_only_timeout_and_always_rolls_back() -> None: + session = _FakeSession() + cancellation = NativeThreadRAGRetrievalCancellation() + + result = RAGRetrievalSessionRunner( + session_factory=lambda: session, + monotonic_ns=lambda: 0, + ).run( + timeout_ms=125, + cancellation=cancellation, + operation=lambda current: current is session, + ) + + assert result is True + assert session.events[0] == "connection" + assert "SET TRANSACTION READ ONLY" in session.events[1][0] + assert "set_config" in session.events[2][0] + assert session.events[2][1] == {"statement_timeout": "125ms"} + assert session.events[-2:] == ["rollback", "close"] + cancellation.cancel() + assert session.driver.cancel_count == 0 + + +def test_session_runner_opens_and_closes_one_session_per_invocation() -> None: + sessions = [_FakeSession(), _FakeSession()] + opened = [] + + def session_factory(): + session = sessions[len(opened)] + opened.append(session) + return session + + runner = RAGRetrievalSessionRunner(session_factory=session_factory) + observed = [ + runner.run( + timeout_ms=125, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=id, + ) + for _index in range(2) + ] + + assert observed == [id(session) for session in sessions] + assert opened == sessions + assert all(session.events[-2:] == ["rollback", "close"] for session in sessions) + + +def test_session_runner_cancellation_calls_driver_and_returns_safe_timeout() -> None: + session = _FakeSession() + cancellation = NativeThreadRAGRetrievalCancellation() + operation_started = threading.Event() + failures = [] + + def operation(_session): + operation_started.set() + session.driver.cancelled.wait(timeout=1) + return "late-result" + + def run(): + try: + RAGRetrievalSessionRunner(session_factory=lambda: session).run( + timeout_ms=125, + cancellation=cancellation, + operation=operation, + ) + except Exception as exc: # pragma: no cover - asserted below + failures.append(exc) + + worker = threading.Thread(target=run) + worker.start() + assert operation_started.wait(timeout=1) + + cancellation.cancel() + gevent.sleep(0) + worker.join(timeout=1) + + assert worker.is_alive() is False + assert session.driver.cancel_count == 1 + assert session.events[-2:] == ["rollback", "close"] + assert len(failures) == 1 + assert isinstance(failures[0], RAGRetrievalSessionTimeout) + assert "late-result" not in str(failures[0]) + + +def test_session_runner_redacts_operation_and_cleanup_failures() -> None: + session = _FakeSession(rollback_error=True, close_error=True) + + with pytest.raises(RAGRetrievalSessionError) as captured: + RAGRetrievalSessionRunner(session_factory=lambda: session).run( + timeout_ms=125, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=lambda _session: (_ for _ in ()).throw( + RuntimeError("private-query-detail") + ), + ) + + assert session.events[-2:] == ["rollback", "close"] + assert "private-query-detail" not in str(captured.value) + assert "private-rollback-detail" not in str(captured.value) + assert "private-close-detail" not in str(captured.value) + + +def test_session_runner_blocks_next_statement_after_cancellation() -> None: + session = _SQLAlchemyBackedSession() + cancellation = NativeThreadRAGRetrievalCancellation() + second_statement_reached = False + + def operation(current): + nonlocal second_statement_reached + assert current.execute(text("SELECT 1")).scalar_one() == 1 + cancellation.cancel() + current.execute(text("SELECT 2")) + second_statement_reached = True + + try: + with pytest.raises(RAGRetrievalSessionTimeout): + RAGRetrievalSessionRunner(session_factory=lambda: session).run( + timeout_ms=125, + cancellation=cancellation, + operation=operation, + ) + + assert second_statement_reached is False + assert session.events[-2:] == ["rollback", "close"] + finally: + session.dispose() + + +def test_session_runner_blocks_next_statement_after_absolute_deadline() -> None: + session = _SQLAlchemyBackedSession() + now_ns = 0 + + def monotonic_ns() -> int: + return now_ns + + def operation(current): + nonlocal now_ns + assert current.execute(text("SELECT 1")).scalar_one() == 1 + now_ns = 126_000_000 + current.execute(text("SELECT 2")) + + try: + with pytest.raises(RAGRetrievalSessionTimeout): + RAGRetrievalSessionRunner( + session_factory=lambda: session, + monotonic_ns=monotonic_ns, + ).run( + timeout_ms=125, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=operation, + ) + + assert session.events[-2:] == ["rollback", "close"] + finally: + session.dispose() + + +def test_postgres_statement_guard_uses_remaining_absolute_budget() -> None: + now_ns = 40_000_000 + cancellation = NativeThreadRAGRetrievalCancellation() + cursor_calls = [] + cursor = SimpleNamespace( + execute=lambda statement, parameters: cursor_calls.append( + (statement, parameters) + ) + ) + connection = SimpleNamespace(dialect=SimpleNamespace(name="postgresql")) + guard = _StatementDeadlineGuard( + cancellation=cancellation, + deadline_ns=125_000_000, + monotonic_ns=lambda: now_ns, + ) + + guard(connection, cursor, "SELECT 1", (), None, False) + + assert cursor_calls == [ + ( + "SELECT set_config('statement_timeout', %s, true)", + ("85ms",), + ) + ] + + now_ns = 126_000_000 + with pytest.raises(RAGRetrievalSessionTimeout): + guard(connection, cursor, "SELECT 2", (), None, False) + + assert len(cursor_calls) == 1 + + +def test_session_runner_removes_statement_guard_before_connection_reuse() -> None: + session = _SQLAlchemyBackedSession() + cancellation = NativeThreadRAGRetrievalCancellation() + + try: + result = RAGRetrievalSessionRunner(session_factory=lambda: session).run( + timeout_ms=125, + cancellation=cancellation, + operation=lambda current: current.execute(text("SELECT 1")).scalar_one(), + ) + cancellation.cancel() + + assert result == 1 + assert session.sql_connection.execute(text("SELECT 2")).scalar_one() == 2 + finally: + session.dispose() + + +def test_session_runner_maps_acquisition_timeout_to_safe_session_timeout() -> None: + class TimeoutAcquirer: + @staticmethod + def acquire(**_kwargs): + raise RAGRetrievalConnectionAcquisitionTimeout() + + with pytest.raises(RAGRetrievalSessionTimeout): + RAGRetrievalSessionRunner( + session_factory=lambda: (_ for _ in ()).throw( + RuntimeError("private-factory-detail") + ), + connection_acquirer=TimeoutAcquirer(), + ).run( + timeout_ms=125, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=lambda _session: None, + ) + + +def test_session_runner_redacts_session_factory_failure() -> None: + with pytest.raises(RAGRetrievalSessionError) as captured: + RAGRetrievalSessionRunner( + session_factory=lambda: (_ for _ in ()).throw( + RuntimeError("private-factory-detail") + ) + ).run( + timeout_ms=125, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=lambda _session: None, + ) + + assert "private-factory-detail" not in str(captured.value) + + +@pytest.mark.parametrize("timeout_ms", [0, -1, True, 30001]) +def test_session_runner_rejects_invalid_timeout_before_opening_session( + timeout_ms, +) -> None: + opened = [] + + with pytest.raises(RAGRetrievalSessionError): + RAGRetrievalSessionRunner(session_factory=lambda: opened.append(True)).run( + timeout_ms=timeout_ms, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=lambda _session: None, + ) + + assert opened == [] diff --git a/apps/workflow_engine/tests/adapters/test_rag_retrieval_session_postgres.py b/apps/workflow_engine/tests/adapters/test_rag_retrieval_session_postgres.py new file mode 100644 index 000000000..26f179690 --- /dev/null +++ b/apps/workflow_engine/tests/adapters/test_rag_retrieval_session_postgres.py @@ -0,0 +1,202 @@ +# ruff: noqa: E402 + +import os +import threading +import time +from uuid import uuid4 + +import pytest + +RUN_ENV = "NODEASE_RUN_DISPOSABLE_DB_TEST" +DB_PREFIX = "nodease_rag_fanout_test" + +if os.getenv(RUN_ENV) != "1": + pytest.skip( + f"set {RUN_ENV}=1 to run disposable RAG fan-out evidence", + allow_module_level=True, + ) + +from sqlalchemy import create_engine, text +from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy.orm import sessionmaker + +from apps.shared.tests.helpers.disposable_postgres import ( + DisposablePostgresConfig, + DisposablePostgresConfigurationError, + quote_disposable_database_name, +) +from apps.workflow_engine.adapters.rag_retrieval_session import ( + RAGRetrievalSessionRunner, + RAGRetrievalSessionTimeout, +) +from apps.workflow_engine.adapters.rag_retrieval_executor import ( + NativeThreadRAGRetrievalCancellation, +) + + +@pytest.fixture +def disposable_rag_database(): + try: + config = DisposablePostgresConfig.from_environment() + except DisposablePostgresConfigurationError: + raise pytest.fail.Exception( + "disposable PostgreSQL connection settings are not safely configured", + pytrace=False, + ) from None + + database = f"{DB_PREFIX}_{uuid4().hex[:12]}" + quoted_database = quote_disposable_database_name(database, prefix=DB_PREFIX) + admin_engine = create_engine( + config.database_url(config.maintenance_database), + isolation_level="AUTOCOMMIT", + ) + engine = None + database_created = False + try: + try: + with admin_engine.connect() as connection: + connection.execute(text(f"CREATE DATABASE {quoted_database}")) + database_created = True + engine = create_engine(config.database_url(database), pool_size=1) + except SQLAlchemyError: + raise pytest.fail.Exception( + "disposable PostgreSQL setup failed", + pytrace=False, + ) from None + yield engine + finally: + if engine is not None: + engine.dispose() + try: + if database_created: + with admin_engine.connect() as connection: + connection.execute( + text( + "SELECT pg_terminate_backend(pid) " + "FROM pg_stat_activity " + "WHERE datname = :database " + "AND pid <> pg_backend_pid()" + ), + {"database": database}, + ) + connection.execute(text(f"DROP DATABASE {quoted_database}")) + except SQLAlchemyError: + raise pytest.fail.Exception( + "disposable PostgreSQL cleanup failed", + pytrace=False, + ) from None + finally: + admin_engine.dispose() + + +def test_rag_session_is_read_only_and_local_timeout_does_not_leak( + disposable_rag_database, +): + session_factory = sessionmaker(bind=disposable_rag_database) + runner = RAGRetrievalSessionRunner(session_factory=session_factory) + + settings = runner.run( + timeout_ms=500, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=lambda session: ( + session.execute(text("SHOW transaction_read_only")).scalar_one(), + session.execute(text("SHOW statement_timeout")).scalar_one(), + ), + ) + + assert settings[0] == "on" + assert settings[1].endswith("ms") + assert 1 <= int(settings[1][:-2]) <= 500 + with session_factory() as session: + assert session.execute(text("SHOW statement_timeout")).scalar_one() == "0" + assert session.execute(text("SELECT 1")).scalar_one() == 1 + + +def test_rag_statement_timeout_rolls_back_before_connection_reuse( + disposable_rag_database, +): + session_factory = sessionmaker(bind=disposable_rag_database) + runner = RAGRetrievalSessionRunner(session_factory=session_factory) + + with pytest.raises(RAGRetrievalSessionTimeout): + runner.run( + timeout_ms=50, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=lambda session: session.execute( + text("SELECT pg_sleep(0.2)") + ).scalar_one(), + ) + + with session_factory() as session: + assert session.execute(text("SELECT 1")).scalar_one() == 1 + + +def test_rag_statement_timeout_shrinks_across_multiple_statements( + disposable_rag_database, +): + session_factory = sessionmaker(bind=disposable_rag_database) + runner = RAGRetrievalSessionRunner(session_factory=session_factory) + + def exceed_cumulative_budget(session): + session.execute(text("SELECT pg_sleep(0.03)")).scalar_one() + session.execute(text("SELECT pg_sleep(0.03)")).scalar_one() + + with pytest.raises(RAGRetrievalSessionTimeout): + runner.run( + timeout_ms=50, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=exceed_cumulative_budget, + ) + + with session_factory() as session: + assert session.execute(text("SELECT 1")).scalar_one() == 1 + + +def test_rag_connection_checkout_returns_before_shared_pool_timeout( + disposable_rag_database, +): + bounded_engine = create_engine( + disposable_rag_database.url, + pool_size=1, + max_overflow=0, + pool_timeout=2, + ) + session_factory = sessionmaker(bind=bounded_engine) + runner = RAGRetrievalSessionRunner(session_factory=session_factory) + held_connection = bounded_engine.connect() + caller_finished = threading.Event() + failures = [] + + def run(): + try: + runner.run( + timeout_ms=50, + cancellation=NativeThreadRAGRetrievalCancellation(), + operation=lambda session: session.execute( + text("SELECT 1") + ).scalar_one(), + ) + except Exception as exc: # pragma: no cover - asserted below + failures.append(exc) + finally: + caller_finished.set() + + caller = threading.Thread(target=run) + try: + caller.start() + + assert caller_finished.wait(timeout=0.5) + assert len(failures) == 1 + assert isinstance(failures[0], RAGRetrievalSessionTimeout) + finally: + held_connection.close() + caller.join(timeout=1) + + cleanup_deadline = time.monotonic() + 1 + while bounded_engine.pool.checkedout() and time.monotonic() < cleanup_deadline: + time.sleep(0.01) + checked_out = bounded_engine.pool.checkedout() + bounded_engine.dispose() + + assert caller.is_alive() is False + assert checked_out == 0 diff --git a/apps/workflow_engine/tests/application/test_rag_retrieval_fanout.py b/apps/workflow_engine/tests/application/test_rag_retrieval_fanout.py new file mode 100644 index 000000000..f8dcabb00 --- /dev/null +++ b/apps/workflow_engine/tests/application/test_rag_retrieval_fanout.py @@ -0,0 +1,318 @@ +import threading +import time + +import pytest + +from apps.workflow_engine.application.rag_retrieval_fanout import ( + RAGRetrievalFanoutConfigurationError, + RAGRetrievalFanoutError, + RAGRetrievalFanoutScheduler, + RAGRetrievalFanoutTask, +) +from apps.workflow_engine.adapters.rag_retrieval_executor import ( + GeventNativeThreadRAGRetrievalExecutor, + NativeThreadRAGRetrievalCancellation, +) + + +def _tasks(count: int) -> tuple[RAGRetrievalFanoutTask, ...]: + return tuple( + RAGRetrievalFanoutTask(ordinal=index, resource_ref=f"resource-{index}") + for index in range(count) + ) + + +def test_scheduler_overlaps_native_workers_caps_concurrency_and_orders_results() -> ( + None +): + barrier = threading.Barrier(2) + state_lock = threading.Lock() + active = 0 + max_active = 0 + thread_ids: set[int] = set() + + def search(task, _cancellation, _timeout_ms): + nonlocal active, max_active + with state_lock: + active += 1 + max_active = max(max_active, active) + thread_ids.add(threading.get_ident()) + try: + barrier.wait(timeout=1) + time.sleep(0.01 * (4 - task.ordinal)) + return f"result-{task.ordinal}" + finally: + with state_lock: + active -= 1 + + result = RAGRetrievalFanoutScheduler( + executor_factory=GeventNativeThreadRAGRetrievalExecutor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + max_workers=2, + per_task_timeout_seconds=1, + aggregate_timeout_seconds=2, + cleanup_reserve_seconds=0.1, + minimum_start_budget_ms=1, + ).execute(tasks=_tasks(4), worker=search) + + assert [task.ordinal for task, _value in result.results] == [0, 1, 2, 3] + assert [value for _task, value in result.results] == [ + "result-0", + "result-1", + "result-2", + "result-3", + ] + assert max_active == 2 + assert len(thread_ids) == 2 + assert result.failed_count == 0 + assert result.timeout_count == 0 + assert isinstance(result.slowest_search_latency_ms, int) + + +def test_scheduler_timeout_cancels_active_task_and_ignores_late_result() -> None: + cancellation_observed = threading.Event() + + def search(_task, cancellation, _timeout_ms): + unregister = cancellation.register(cancellation_observed.set) + try: + cancellation_observed.wait(timeout=1) + return "late-result" + finally: + unregister() + + started = time.monotonic() + result = RAGRetrievalFanoutScheduler( + executor_factory=GeventNativeThreadRAGRetrievalExecutor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + max_workers=1, + per_task_timeout_seconds=0.03, + aggregate_timeout_seconds=0.2, + cleanup_reserve_seconds=0.05, + minimum_start_budget_ms=1, + ).execute(tasks=_tasks(1), worker=search) + + assert time.monotonic() - started < 0.5 + assert cancellation_observed.is_set() + assert result.results == () + assert result.failed_count == 1 + assert result.timeout_count == 1 + + +def test_scheduler_hard_deadline_does_not_wait_for_uncooperative_worker() -> None: + worker_finished = threading.Event() + + def search(_task, _cancellation, _timeout_ms): + try: + time.sleep(0.3) + return "late-result" + finally: + worker_finished.set() + + started = time.monotonic() + result = RAGRetrievalFanoutScheduler( + executor_factory=GeventNativeThreadRAGRetrievalExecutor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + max_workers=1, + per_task_timeout_seconds=0.03, + aggregate_timeout_seconds=0.12, + cleanup_reserve_seconds=0.05, + minimum_start_budget_ms=1, + ).execute(tasks=_tasks(1), worker=search) + scheduler_elapsed = time.monotonic() - started + + assert scheduler_elapsed < 0.25 + assert result.results == () + assert result.failed_count == 1 + assert result.timeout_count == 1 + assert worker_finished.wait(timeout=1) + + +def test_scheduler_fail_fast_cancels_other_work_and_uses_safe_error() -> None: + barrier = threading.Barrier(2) + cancellation_observed = threading.Event() + + def search(task, cancellation, _timeout_ms): + barrier.wait(timeout=1) + if task.ordinal == 0: + raise RuntimeError("private-database-detail") + unregister = cancellation.register(cancellation_observed.set) + try: + cancellation_observed.wait(timeout=1) + return "late-result" + finally: + unregister() + + scheduler = RAGRetrievalFanoutScheduler( + executor_factory=GeventNativeThreadRAGRetrievalExecutor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + max_workers=2, + per_task_timeout_seconds=1, + aggregate_timeout_seconds=1, + cleanup_reserve_seconds=0.1, + minimum_start_budget_ms=1, + ) + + with pytest.raises(RAGRetrievalFanoutError) as captured: + scheduler.execute(tasks=_tasks(2), worker=search, fail_fast=True) + + assert cancellation_observed.is_set() + assert "private-database-detail" not in str(captured.value) + + +def test_scheduler_redacts_executor_coordination_failure() -> None: + class NeverReadyJob: + def ready(self): + return False + + def result(self): + raise AssertionError("unfinished job must not be consumed") + + class FailingExecutor: + def __init__(self): + self.closed = False + + def submit(self, _callback): + return NeverReadyJob() + + def wait(self, _jobs, *, timeout_seconds): + assert timeout_seconds > 0 + raise RuntimeError("private-executor-detail") + + def close(self): + self.closed = True + + executor = FailingExecutor() + scheduler = RAGRetrievalFanoutScheduler( + executor_factory=lambda _max_workers: executor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + aggregate_timeout_seconds=1, + cleanup_reserve_seconds=0.1, + minimum_start_budget_ms=1, + ) + + with pytest.raises(RAGRetrievalFanoutError) as captured: + scheduler.execute(tasks=_tasks(1), worker=lambda *_args: None) + + assert executor.closed is True + assert "private-executor-detail" not in str(captured.value) + + +def test_scheduler_stops_when_queued_task_has_insufficient_start_budget() -> None: + class RecordingCancellation: + def __init__(self): + self.cancelled = False + + def register(self, _callback): + return lambda: None + + def cancel(self): + self.cancelled = True + + class UnusedExecutor: + def __init__(self): + self.closed = False + + def submit(self, _callback): + raise AssertionError("insufficient-budget task must not be submitted") + + def wait(self, _jobs, *, timeout_seconds): + raise AssertionError("scheduler must not wait without running tasks") + + def close(self): + self.closed = True + + class FrozenBudgetClock: + def __init__(self): + self.calls = 0 + + def __call__(self): + self.calls += 1 + if self.calls > 8: + raise AssertionError("scheduler busy-spun without clock progress") + return 0.0 if self.calls == 1 else 0.45 + + executor = UnusedExecutor() + clock = FrozenBudgetClock() + result = RAGRetrievalFanoutScheduler( + executor_factory=lambda _max_workers: executor, + cancellation_factory=RecordingCancellation, + aggregate_timeout_seconds=1, + cleanup_reserve_seconds=0.1, + minimum_start_budget_ms=500, + clock=clock, + ).execute(tasks=_tasks(1), worker=lambda *_args: None) + + assert result.results == () + assert result.failed_count == 1 + assert result.timeout_count == 1 + assert executor.closed is True + assert clock.calls <= 8 + + +def test_scheduler_external_base_exception_cancels_all_running_tasks() -> None: + class ExternalAbort(BaseException): + pass + + class RecordingCancellation: + def __init__(self): + self.cancelled = False + + def register(self, _callback): + return lambda: None + + def cancel(self): + self.cancelled = True + + class NeverReadyJob: + def ready(self): + return False + + def result(self): + raise AssertionError("unfinished job must not be consumed") + + class AbortingExecutor: + def __init__(self): + self.closed = False + + def submit(self, _callback): + return NeverReadyJob() + + def wait(self, _jobs, *, timeout_seconds): + assert timeout_seconds > 0 + raise ExternalAbort() + + def close(self): + self.closed = True + + cancellations = [] + + def cancellation_factory(): + cancellation = RecordingCancellation() + cancellations.append(cancellation) + return cancellation + + executor = AbortingExecutor() + scheduler = RAGRetrievalFanoutScheduler( + executor_factory=lambda _max_workers: executor, + cancellation_factory=cancellation_factory, + aggregate_timeout_seconds=1, + cleanup_reserve_seconds=0.1, + minimum_start_budget_ms=1, + ) + + with pytest.raises(ExternalAbort): + scheduler.execute(tasks=_tasks(1), worker=lambda *_args: None) + + assert len(cancellations) == 2 + assert all(cancellation.cancelled for cancellation in cancellations) + assert executor.closed is True + + +def test_scheduler_rejects_more_than_runtime_candidate_cap() -> None: + scheduler = RAGRetrievalFanoutScheduler( + executor_factory=GeventNativeThreadRAGRetrievalExecutor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + ) + + with pytest.raises(RAGRetrievalFanoutConfigurationError): + scheduler.execute(tasks=_tasks(21), worker=lambda *_args: None) diff --git a/apps/workflow_engine/tests/nodes/test_llm_node_runtime.py b/apps/workflow_engine/tests/nodes/test_llm_node_runtime.py index fd6f0f605..e5aac3574 100644 --- a/apps/workflow_engine/tests/nodes/test_llm_node_runtime.py +++ b/apps/workflow_engine/tests/nodes/test_llm_node_runtime.py @@ -69,6 +69,16 @@ from apps.workflow_engine.application.provider_execution import ( # noqa: E402 ProviderExecutionConfigurationError, ) +from apps.workflow_engine.application.rag_retrieval_fanout import ( # noqa: E402 + RAGRetrievalFanoutError, + RAGRetrievalFanoutResult as ApplicationRAGRetrievalFanoutResult, +) +from apps.workflow_engine.adapters.rag_retrieval_executor import ( # noqa: E402 + NativeThreadRAGRetrievalCancellation, +) +from apps.workflow_engine.adapters.rag_retrieval_session import ( # noqa: E402 + RAGRetrievalSessionError, +) from apps.workflow_engine.composition.provider_execution import ( # noqa: E402 build_provider_execution_runtime, build_provider_usage_recorder, @@ -481,51 +491,54 @@ def resolve( def _patch_rag_gevent_inline(monkeypatch, node): - """RAG fanout unit test에서 gevent 의존성 없이 bounded path를 동기 실행한다.""" - - class FakeTimeout(Exception): - def __init__(self, seconds): - self.seconds = seconds - - def start(self): - return None - - def cancel(self): - return None - - class FakeJob: - def __init__(self, fn, kwargs): - try: - self.value = fn(**kwargs) - self.exception = None - except Exception as exc: # pragma: no cover - assertion에서 검증 - self.value = None - self.exception = exc - - def ready(self): - return True - - def kill(self, block=False): - return None - - class FakePool: - def __init__(self, size): - self.size = size - - def spawn(self, fn, **kwargs): - return FakeJob(fn, kwargs) - - def kill(self, block=False): - return None - - class FakeGevent: - Timeout = FakeTimeout + """기존 테스트 이름을 유지하면서 fanout scheduler를 동기 실행한다.""" + + class InlineScheduler: + def execute(self, *, tasks, worker, fail_fast=False): + results = [] + failed_count = 0 + timeout_count = 0 + for task in tasks: + try: + value = worker( + task, + NativeThreadRAGRetrievalCancellation(), + 10_000, + ) + except Exception as exc: + if fail_fast: + raise RAGRetrievalFanoutError() from None + failed_count += 1 + if isinstance(exc, TimeoutError): + timeout_count += 1 + continue + results.append((task, value)) + return ApplicationRAGRetrievalFanoutResult( + results=tuple(results), + failed_count=failed_count, + timeout_count=timeout_count, + slowest_search_latency_ms=0, + ) - @staticmethod - def joinall(jobs, timeout=None): - return jobs + class InlineSessionRunner: + def run(self, *, cancellation, operation, **_kwargs): + if cancellation.cancelled: + raise TimeoutError("cancelled") + session = node.execution_context.get("db", SimpleNamespace()) + return operation(session) - monkeypatch.setattr(node, "_rag_gevent_modules", lambda: (FakeGevent, FakePool)) + monkeypatch.setattr( + node, + "_rag_retrieval_fanout_scheduler_override", + InlineScheduler(), + raising=False, + ) + monkeypatch.setattr( + node, + "_rag_retrieval_session_runner_override", + InlineSessionRunner(), + raising=False, + ) class FakeRuntimePriorityQuery: @@ -886,9 +899,18 @@ def test_llm_node_rag_query_includes_bounded_client_history_for_follow_up( ) node._client_override = DummyClient() # noqa: SLF001 - provider isolation - def capture_search(query, db_session, *, candidate_resolution=None): + def capture_search( + query, + db_session, + *, + candidate_resolution=None, + candidate_resolution_latency_ms=None, + ): captured["query"] = query captured["candidate_resolution"] = candidate_resolution + captured["candidate_resolution_latency_ms"] = ( + candidate_resolution_latency_ms + ) return WorkflowRAGSearchResult( context="authorized evidence", metadata=[], @@ -907,6 +929,7 @@ def capture_search(query, db_session, *, candidate_resolution=None): assert "[REDACTED: possible prompt injection]" in captured["query"] assert len(captured["query"]) <= 1_000 assert captured["candidate_resolution"].candidates[0].knowledge_base_id == kb_id + assert isinstance(captured["candidate_resolution_latency_ms"], int) def test_llm_node_rag_query_without_client_history_preserves_current_query(): @@ -2905,7 +2928,8 @@ def test_llm_node_rag_no_evidence_skips_llm_call(monkeypatch): monkeypatch.setattr( LLMNode, "_execute_knowledge_search", - lambda self, query, db_session, *, candidate_resolution=None: ( + lambda self, query, db_session, *, candidate_resolution=None, + candidate_resolution_latency_ms=None: ( WorkflowRAGSearchResult( context="", metadata=[], @@ -2947,7 +2971,14 @@ def test_llm_node_rag_operational_failure_uses_safe_no_result(monkeypatch): client = DummyClient() node._client_override = client # noqa: SLF001 - 테스트용 주입 - def raise_retrieval_error(self, query, db_session, *, candidate_resolution=None): + def raise_retrieval_error( + self, + query, + db_session, + *, + candidate_resolution=None, + candidate_resolution_latency_ms=None, + ): raise RuntimeError("vector store unavailable") monkeypatch.setattr(LLMNode, "_execute_knowledge_search", raise_retrieval_error) @@ -3024,7 +3055,7 @@ def search_documents_sync(self, query, *, knowledge_base_id, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -3107,7 +3138,7 @@ def search_documents_sync(self, query, *, threshold, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -3194,7 +3225,7 @@ def search_documents_sync(self, query, *, knowledge_base_id, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -3284,7 +3315,7 @@ def search_documents_sync(self, query, *, knowledge_base_id, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -3359,7 +3390,7 @@ def search_documents_sync( ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -4086,7 +4117,7 @@ def test_rag_audit_omits_unscoped_invalid_organization(monkeypatch, audit_method assert audit_calls == [] -def test_llm_node_rag_fanout_uses_bounded_pool(monkeypatch): +def test_llm_node_rag_fanout_uses_scheduler_and_preserves_candidate_order(monkeypatch): node = LLMNode( "llm-1", LLMNodeData( @@ -4097,35 +4128,8 @@ def test_llm_node_rag_fanout_uses_bounded_pool(monkeypatch): ), ) kb_ids = [str(uuid.uuid4()) for _ in range(7)] - pool_sizes = [] search_calls = [] - class FakeJob: - def __init__(self, value): - self.value = value - self.exception = None - - def ready(self): - return True - - def kill(self, block=False): - return None - - class FakePool: - def __init__(self, size): - pool_sizes.append(size) - - def spawn(self, fn, **kwargs): - return FakeJob(fn(**kwargs)) - - def kill(self, block=False): - return None - - class FakeGevent: - @staticmethod - def joinall(jobs, timeout=None): - return jobs - def fake_search(**kwargs): search_calls.append(kwargs["knowledge_base_id"]) return [ @@ -4138,11 +4142,7 @@ def fake_search(**kwargs): ) ] - monkeypatch.setattr( - node, - "_rag_gevent_modules", - lambda: (FakeGevent, FakePool), - ) + _patch_rag_gevent_inline(monkeypatch, node) monkeypatch.setattr( node, "_search_single_rag_kb_with_new_session", @@ -4151,7 +4151,6 @@ def fake_search(**kwargs): result = node._run_rag_retrieval_fanout( # noqa: SLF001 query="query", - fallback_db_session=object(), user_id=uuid.uuid4(), organization_id=uuid.uuid4(), knowledge_base_ids=kb_ids, @@ -4159,20 +4158,29 @@ def fake_search(**kwargs): threshold=0.5, ) - assert pool_sizes == [5] assert search_calls == kb_ids + assert [kb_id for kb_id, _chunks in result.results] == kb_ids assert result.failed_count == 0 assert len(result.results) == len(kb_ids) def test_precomputed_fanout_log_uses_bucketed_candidate_count(monkeypatch, caplog): - node = LLMNode.__new__(LLMNode) + node = LLMNode( + "llm-1", + LLMNodeData( + title="LLM", + provider="openai", + model_id="gpt-4o", + user_prompt="user", + ), + ) kb_ids = [str(uuid.uuid4()) for _ in range(3)] - expected = WorkflowRAGFanoutResult(results=[], failed_count=0) + search_calls = [] + _patch_rag_gevent_inline(monkeypatch, node) monkeypatch.setattr( node, - "_run_rag_retrieval_fanout_sequential", - lambda **kwargs: expected, + "_search_single_rag_kb_with_new_session", + lambda **kwargs: search_calls.append(kwargs) or [], ) with caplog.at_level( @@ -4181,7 +4189,6 @@ def test_precomputed_fanout_log_uses_bucketed_candidate_count(monkeypatch, caplo ): result = node._run_rag_retrieval_fanout( # noqa: SLF001 - safe log contract query="query", - fallback_db_session=object(), user_id=uuid.uuid4(), organization_id=uuid.uuid4(), knowledge_base_ids=kb_ids, @@ -4190,7 +4197,9 @@ def test_precomputed_fanout_log_uses_bucketed_candidate_count(monkeypatch, caplo query_vectors_by_kb={kb_id: [0.1] for kb_id in kb_ids}, ) - assert result is expected + assert result.failed_count == 0 + assert [call["knowledge_base_id"] for call in search_calls] == kb_ids + assert [call["query_vector"] for call in search_calls] == [[0.1]] * 3 log_text = " ".join(caplog.messages) assert "kb_count_bucket=2-10" in log_text assert "kb_count=3" not in log_text @@ -4249,7 +4258,7 @@ def test_query_vector_precompute_log_uses_only_count_buckets(monkeypatch, caplog assert "failed_count=0" not in log_text -def test_llm_node_rag_single_kb_uses_bounded_pool(monkeypatch): +def test_llm_node_rag_single_kb_uses_scheduler(monkeypatch): node = LLMNode( "llm-1", LLMNodeData( @@ -4260,44 +4269,13 @@ def test_llm_node_rag_single_kb_uses_bounded_pool(monkeypatch): ), ) kb_id = str(uuid.uuid4()) - pool_sizes = [] search_calls = [] - class FakeJob: - def __init__(self, value): - self.value = value - self.exception = None - - def ready(self): - return True - - def kill(self, block=False): - return None - - class FakePool: - def __init__(self, size): - pool_sizes.append(size) - - def spawn(self, fn, **kwargs): - return FakeJob(fn(**kwargs)) - - def kill(self, block=False): - return None - - class FakeGevent: - @staticmethod - def joinall(jobs, timeout=None): - return jobs - def fake_search(**kwargs): search_calls.append(kwargs["knowledge_base_id"]) return [] - monkeypatch.setattr( - node, - "_rag_gevent_modules", - lambda: (FakeGevent, FakePool), - ) + _patch_rag_gevent_inline(monkeypatch, node) monkeypatch.setattr( node, "_search_single_rag_kb_with_new_session", @@ -4306,7 +4284,6 @@ def fake_search(**kwargs): result = node._run_rag_retrieval_fanout( # noqa: SLF001 query="query", - fallback_db_session=object(), user_id=uuid.uuid4(), organization_id=uuid.uuid4(), knowledge_base_ids=[kb_id], @@ -4314,12 +4291,14 @@ def fake_search(**kwargs): threshold=0.5, ) - assert pool_sizes == [1] assert search_calls == [kb_id] assert result.failed_count == 0 -def test_llm_node_rag_fails_closed_when_timeout_guard_unavailable(monkeypatch): +def test_llm_node_rag_session_runner_uses_injected_session_factory(monkeypatch): + def injected_factory(): + return object() + node = LLMNode( "llm-1", LLMNodeData( @@ -4328,23 +4307,103 @@ def test_llm_node_rag_fails_closed_when_timeout_guard_unavailable(monkeypatch): model_id="gpt-4o", user_prompt="user", ), + execution_context={"db_session_factory": injected_factory}, ) + captured_factories = [] + runner = object() - monkeypatch.setattr(node, "_rag_gevent_modules", lambda: None) + monkeypatch.setattr( + "apps.workflow_engine.adapters.rag_retrieval_session." + "RAGRetrievalSessionRunner", + lambda *, session_factory: captured_factories.append(session_factory) or runner, + ) - result = node._run_rag_retrieval_fanout( # noqa: SLF001 - query="query", - fallback_db_session=object(), - user_id=uuid.uuid4(), - organization_id=uuid.uuid4(), - knowledge_base_ids=[str(uuid.uuid4())], - top_k=3, - threshold=0.5, + assert node._get_rag_retrieval_session_runner() is runner # noqa: SLF001 + assert captured_factories == [injected_factory] + + +def test_llm_node_rag_session_runner_rejects_invalid_injected_factory(): + node = LLMNode( + "llm-1", + LLMNodeData( + title="LLM", + provider="openai", + model_id="gpt-4o", + user_prompt="user", + ), + execution_context={"db_session_factory": "invalid"}, ) - assert result.results == [] - assert result.failed_count == 1 - assert result.timeout_count == 1 + with pytest.raises(RAGRetrievalSessionError): + node._get_rag_retrieval_session_runner() # noqa: SLF001 + + +def test_llm_node_rag_fails_closed_when_scheduler_is_unavailable(monkeypatch): + node = LLMNode( + "llm-1", + LLMNodeData( + title="LLM", + provider="openai", + model_id="gpt-4o", + user_prompt="user", + ), + ) + + class FailingScheduler: + def execute(self, **_kwargs): + raise RAGRetrievalFanoutError() + + monkeypatch.setattr( + node, + "_rag_retrieval_fanout_scheduler_override", + FailingScheduler(), + raising=False, + ) + + with pytest.raises(RAGRetrievalFanoutError): + node._run_rag_retrieval_fanout( # noqa: SLF001 + query="query", + user_id=uuid.uuid4(), + organization_id=uuid.uuid4(), + knowledge_base_ids=[str(uuid.uuid4())], + top_k=3, + threshold=0.5, + ) + + +def test_rag_stage_latencies_are_trace_only_not_result_metadata(): + node = LLMNode( + "llm-1", + LLMNodeData(title="LLM", model_id="gpt-4o", user_prompt="user"), + ) + stage_latencies = { + "candidate_resolution_latency_ms": 1, + "query_embedding_latency_ms": 2, + "retrieval_fanout_latency_ms": 3, + "slowest_search_latency_ms": 4, + "evidence_policy_latency_ms": 5, + } + knowledge_result = WorkflowRAGSearchResult( + context="", + metadata=[], + evidence_decision=RAGEvidenceDecision( + evidence_sufficient=False, + insufficiency_reason="no_evidence", + ), + should_invoke_llm=False, + trace_summary={"retrieval_strategy": "safe", **stage_latencies}, + ) + + result_metadata = node._rag_result_metadata(knowledge_result) # noqa: SLF001 + trace_payload = node._rag_retrieval_trace_payload( # noqa: SLF001 + [], + evidence_decision=knowledge_result.evidence_decision, + runtime_summary=knowledge_result.trace_summary, + ) + + assert result_metadata["retrieval_strategy"] == "safe" + assert not set(stage_latencies).intersection(result_metadata) + assert {key: trace_payload[key] for key in stage_latencies} == stage_latencies def test_llm_node_rag_partial_retrieval_failure_respects_fail_node_policy( @@ -4374,7 +4433,7 @@ def search_documents_sync(self, *args, **kwargs): raise RuntimeError("vector store unavailable") monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) @@ -4400,8 +4459,9 @@ def search_documents_sync(self, *args, **kwargs): ) _patch_rag_gevent_inline(monkeypatch, node) - with pytest.raises(RuntimeError, match="vector store unavailable"): + with pytest.raises(RAGRetrievalFanoutError) as captured: node._execute_knowledge_search("query", db_session=FakeDb()) # noqa: SLF001 + assert "vector store unavailable" not in str(captured.value) def test_llm_runtime_permission_denied_uses_detailed_reason_and_unknown_target( @@ -6096,6 +6156,7 @@ def fake_search( db_session, *, candidate_resolution=None, + candidate_resolution_latency_ms=None, ): captured_resolutions.append(candidate_resolution) return WorkflowRAGSearchResult( @@ -6281,6 +6342,8 @@ def test_collection_only_zero_candidates_skips_retrieval_embedding_and_provider( assert result["metadata"]["rag"]["candidate_resolution_status"] == ( "safe_no_result" ) + trace_payload = node._trace_payloads[0]["payload"] # noqa: SLF001 + assert "candidate_resolution_latency_ms" not in trace_payload safe_output = json.dumps( { "result": result, @@ -6357,6 +6420,8 @@ def test_empty_rendered_rag_query_never_invokes_provider(): assert len(resolver.calls) == 1 assert client.calls == [] assert result["text"] == RAG_NO_EVIDENCE_MESSAGE + trace_payload = node._trace_payloads[0]["payload"] # noqa: SLF001 + assert "candidate_resolution_latency_ms" not in trace_payload def test_runtime_candidate_infrastructure_error_bypasses_rag_failure_policy(): @@ -6730,7 +6795,7 @@ def search_documents_sync(self, *args, **kwargs): return [] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) fake_db = _patch_allowed_knowledge_permissions(monkeypatch, [kb_id]) @@ -6795,7 +6860,7 @@ def search_documents_sync(self, *args, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -6874,7 +6939,7 @@ def search_documents_sync(self, *args, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -6963,7 +7028,7 @@ def search_documents_sync(self, *args, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -7037,7 +7102,7 @@ def search_documents_sync(self, *args, **kwargs): ) monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) fake_db = _patch_allowed_knowledge_permissions(monkeypatch, [kb_id]) @@ -7097,7 +7162,7 @@ def search_documents_sync(self, *args, **kwargs): raise AssertionError(f"unexpected KB: {kwargs['knowledge_base_id']}") monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -7178,7 +7243,7 @@ def search_documents_sync(self, *args, **kwargs): ), ], ) -def test_workflow_llm_node_fail_node_propagates_retrieval_failures( +def test_workflow_llm_node_fail_node_redacts_retrieval_failures( monkeypatch, exception, ): @@ -7195,7 +7260,7 @@ def search_documents_sync(self, *args, **kwargs): raise exception monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) node = LLMNode( @@ -7221,9 +7286,11 @@ def search_documents_sync(self, *args, **kwargs): _patch_rag_gevent_inline(monkeypatch, node) node._client_override = StaticTextClient("unused") # noqa: SLF001 - with pytest.raises(type(exception), match=str(exception)): + with pytest.raises(RAGRetrievalFanoutError) as captured: node.execute({}) + assert str(exception) not in str(captured.value) + def test_knowledge_search_partial_timeout_trace_summary_is_safe(monkeypatch): user_id = uuid.uuid4() @@ -7296,7 +7363,7 @@ def search_documents_sync(self, *args, **kwargs): return [_chunk_preview("현재 근거", filename="current.md")] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) monkeypatch.setattr( @@ -7373,7 +7440,7 @@ def search_documents_sync(self, *args, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) fake_db = _patch_allowed_knowledge_permissions(monkeypatch, [kb_id]) @@ -7441,7 +7508,7 @@ def search_documents_sync(self, *args, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) fake_db = _patch_allowed_knowledge_permissions(monkeypatch, [kb_id]) @@ -7519,7 +7586,7 @@ def search_documents_sync(self, *args, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) fake_db = _patch_allowed_knowledge_permissions(monkeypatch, [kb_id]) @@ -7585,7 +7652,7 @@ def search_documents_sync(self, *args, **kwargs): ] monkeypatch.setattr( - "apps.workflow_engine.workflow.nodes.llm.llm_node.RetrievalService", + "apps.workflow_engine.services.retrieval.RetrievalService", FakeRetrievalService, ) fake_db = _patch_allowed_knowledge_permissions(monkeypatch, [kb_id]) diff --git a/apps/workflow_engine/tests/services/test_retrieval_metadata.py b/apps/workflow_engine/tests/services/test_retrieval_metadata.py index 1e56eede1..4e48c5ea4 100644 --- a/apps/workflow_engine/tests/services/test_retrieval_metadata.py +++ b/apps/workflow_engine/tests/services/test_retrieval_metadata.py @@ -1,7 +1,12 @@ import asyncio +import logging +import sys +import threading import uuid from types import SimpleNamespace +import pytest + from apps.shared.db.models.knowledge import KnowledgeBase from apps.shared.db.models.llm import LLMCredential, LLMModel, LLMProvider from apps.shared.schemas.rag import ChunkPreview @@ -9,6 +14,96 @@ from apps.workflow_engine.services.retrieval import RetrievalService +@pytest.mark.parametrize("use_sync", [True, False]) +def test_retrieval_failure_log_does_not_include_raw_exception( + use_sync, + caplog, +): + class FailingDb: + def query(self, _model): + raise RuntimeError("private-query-and-resource-detail") + + service = RetrievalService( + FailingDb(), + uuid.uuid4(), + organization_id=uuid.uuid4(), + ) + + with caplog.at_level( + logging.ERROR, + logger="apps.workflow_engine.services.retrieval", + ): + with pytest.raises(RuntimeError): + if use_sync: + service.search_documents_sync( + "private-query", + knowledge_base_id=str(uuid.uuid4()), + ) + else: + asyncio.run( + service.search_documents( + "private-query", + knowledge_base_id=str(uuid.uuid4()), + ) + ) + + assert "private-query-and-resource-detail" not in " ".join(caplog.messages) + assert "error_type=RuntimeError" in " ".join(caplog.messages) + + +def test_sync_retrieval_requires_organization_before_database_access() -> None: + class UnexpectedDb: + def query(self, _model): + raise AssertionError("organization-less retrieval must not query") + + result = RetrievalService( + UnexpectedDb(), + uuid.uuid4(), + organization_id=None, + ).search_documents_sync( + "query", + knowledge_base_id=str(uuid.uuid4()), + ) + + assert result == [] + + +def test_sync_retrieval_scopes_knowledge_base_lookup_to_organization() -> None: + criteria = [] + + class ScopedQuery: + def filter(self, *values): + criteria.extend(values) + return self + + def first(self): + return None + + class ScopedDb: + def query(self, model): + assert model is KnowledgeBase + return ScopedQuery() + + knowledge_base_id = str(uuid.uuid4()) + organization_id = uuid.uuid4() + result = RetrievalService( + ScopedDb(), + uuid.uuid4(), + organization_id=organization_id, + ).search_documents_sync( + "query", + knowledge_base_id=knowledge_base_id, + ) + + assert result == [] + assert _criterion_compares_column(criteria, "id", knowledge_base_id) + assert _criterion_compares_column( + criteria, + "organization_id", + organization_id, + ) + + class _FakeQuery: def __init__(self, rows, criteria): self.rows = rows @@ -40,7 +135,10 @@ def _criterion_compares_column(criteria, column_name, value): for criterion in criteria: left = getattr(criterion, "left", None) right = getattr(criterion, "right", None) - if getattr(left, "name", None) == column_name and getattr(right, "value", None) == value: + if ( + getattr(left, "name", None) == column_name + and getattr(right, "value", None) == value + ): return True return False @@ -134,6 +232,56 @@ def test_search_method_labels_hierarchical_paths(monkeypatch): ) +def test_cross_encoder_model_initializes_once_across_native_workers(monkeypatch): + constructor_entered = threading.Event() + duplicate_constructor_entered = threading.Event() + release_constructor = threading.Event() + constructor_calls = [] + + class FakeCrossEncoder: + def __init__(self, model_name, *, max_length): + constructor_calls.append((model_name, max_length)) + constructor_entered.set() + if len(constructor_calls) > 1: + duplicate_constructor_entered.set() + release_constructor.wait(timeout=1) + + monkeypatch.setitem( + sys.modules, + "sentence_transformers", + SimpleNamespace(CrossEncoder=FakeCrossEncoder), + ) + monkeypatch.setattr(RetrievalService, "_cross_encoder_model", None) + monkeypatch.setattr(RetrievalService, "_cross_encoder_model_name", None) + + barrier = threading.Barrier(2) + results = [] + errors = [] + + def load_model(): + try: + barrier.wait(timeout=1) + results.append(RetrievalService._get_cross_encoder_model()) + except Exception as exc: # pragma: no cover - asserted below + errors.append(exc) + + workers = [threading.Thread(target=load_model) for _ in range(2)] + for worker in workers: + worker.start() + + assert constructor_entered.wait(timeout=1) + duplicate_started = duplicate_constructor_entered.wait(timeout=0.1) + release_constructor.set() + for worker in workers: + worker.join(timeout=1) + + assert duplicate_started is False + assert errors == [] + assert len(constructor_calls) == 1 + assert len(results) == 2 + assert results[0] is results[1] + + def test_rewrite_query_passes_active_organization_to_llm(monkeypatch): captured = {} user_id = uuid.uuid4() @@ -180,18 +328,24 @@ def test_rewrite_model_selection_filters_credentials_by_active_organization(): organization_id=organization_id, ) - assert service._get_efficient_rewrite_model() == LLMService.EFFICIENT_MODELS["openai"] + assert ( + service._get_efficient_rewrite_model() == LLMService.EFFICIENT_MODELS["openai"] + ) assert _criterion_compares_column(criteria, "organization_id", organization_id) -def test_generate_answer_preserves_references_when_generation_model_missing(monkeypatch): +def test_generate_answer_preserves_references_when_generation_model_missing( + monkeypatch, +): chunk = ChunkPreview( content="검색 결과", document_id=uuid.uuid4(), filename="guide.md", similarity_score=0.9, ) - service = RetrievalService(db=object(), user_id=uuid.uuid4(), organization_id=uuid.uuid4()) + service = RetrievalService( + db=object(), user_id=uuid.uuid4(), organization_id=uuid.uuid4() + ) async def fake_search_documents(*args, **kwargs): return [chunk] @@ -200,7 +354,9 @@ async def fake_search_documents(*args, **kwargs): monkeypatch.setattr( LLMService, "get_client_for_user", - lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("missing credential")), + lambda *args, **kwargs: (_ for _ in ()).throw( + RuntimeError("missing credential") + ), ) response = asyncio.run(service.generate_answer("query", "kb-1")) @@ -254,7 +410,11 @@ def fake_chunk(content): low_chunk = fake_chunk("낮은 점수 근거") high_chunk = fake_chunk("충분한 점수 근거") - service = RetrievalService(db=FakeDb(), user_id=uuid.uuid4()) + service = RetrievalService( + db=FakeDb(), + user_id=uuid.uuid4(), + organization_id=uuid.uuid4(), + ) monkeypatch.setattr(service, "_has_valid_hierarchy", lambda *_args: False) monkeypatch.setattr( service, diff --git a/apps/workflow_engine/workflow/nodes/llm/__init__.py b/apps/workflow_engine/workflow/nodes/llm/__init__.py index 982af2365..7ae626416 100644 --- a/apps/workflow_engine/workflow/nodes/llm/__init__.py +++ b/apps/workflow_engine/workflow/nodes/llm/__init__.py @@ -1,4 +1,19 @@ +from typing import TYPE_CHECKING + from .entities import KnowledgeCollectionRef, LLMNodeData -from .llm_node import LLMNode + +if TYPE_CHECKING: + from .llm_node import LLMNode + + +def __getattr__(name: str) -> object: + """Load worker runtime implementations only when explicitly requested.""" + + if name == "LLMNode": + from .llm_node import LLMNode + + return LLMNode + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + __all__ = ["KnowledgeCollectionRef", "LLMNode", "LLMNodeData"] diff --git a/apps/workflow_engine/workflow/nodes/llm/llm_node.py b/apps/workflow_engine/workflow/nodes/llm/llm_node.py index 389918af3..f73dd2136 100644 --- a/apps/workflow_engine/workflow/nodes/llm/llm_node.py +++ b/apps/workflow_engine/workflow/nodes/llm/llm_node.py @@ -5,7 +5,7 @@ import time import uuid from dataclasses import dataclass, field -from typing import Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any, Dict, List, Optional from jinja2 import Environment from sqlalchemy.exc import SQLAlchemyError @@ -99,6 +99,12 @@ QueryEmbeddingPlan, QueryEmbeddingPreflight, ) +from apps.workflow_engine.application.rag_retrieval_fanout import ( + DEFAULT_RAG_FANOUT_MAX_WORKERS, + RAGRetrievalCancellation, + RAGRetrievalFanoutScheduler, + RAGRetrievalFanoutTask, +) from apps.workflow_engine.application.runtime_retrieval.knowledge_candidates import ( KnowledgeRuntimeCandidateInfrastructureError, KnowledgeRuntimeCandidateResolver, @@ -117,7 +123,6 @@ build_judge_first_active_policy, select_runtime_judge_model_id, ) -from apps.workflow_engine.services.retrieval import RetrievalService from apps.workflow_engine.workflow.errors import ( NonRetryableWorkflowError, ProviderOutcomeUnknownWorkflowError, @@ -130,14 +135,30 @@ LLMNodeData, ) +if TYPE_CHECKING: + from apps.workflow_engine.adapters.rag_retrieval_session import ( + RAGRetrievalSessionRunner, + ) + from apps.workflow_engine.services.retrieval import RetrievalService + logger = logging.getLogger(__name__) _jinja_env = Environment(autoescape=False) MEMORY_RUN_LIMIT = 5 # 최근 실행 몇 건을 기억 컨텍스트에 반영할지 결정 MAX_RAG_TRACE_RETRIEVED_CHUNKS = 20 -MAX_RAG_FANOUT_CONCURRENCY = 5 +MAX_RAG_FANOUT_CONCURRENCY = DEFAULT_RAG_FANOUT_MAX_WORKERS RAG_FANOUT_AGGREGATE_TIMEOUT_SECONDS = 30.0 RAG_FANOUT_PER_KB_TIMEOUT_SECONDS = 10.0 +MAX_RAG_STAGE_LATENCY_MS = 300_000 +RAG_TRACE_STAGE_LATENCY_FIELDS = frozenset( + { + "candidate_resolution_latency_ms", + "query_embedding_latency_ms", + "retrieval_fanout_latency_ms", + "slowest_search_latency_ms", + "evidence_policy_latency_ms", + } +) MAX_RAG_REWRITTEN_QUERY_LENGTH = 1000 QUERY_REWRITE_PLACEHOLDER_RE = re.compile(r"\{\{\s*query\s*\}\}|\{query\}") PROVIDER_HTTP_STATUS_RE = re.compile(r"\bstatus\s*[=:]?\s*(\d{3})\b", re.IGNORECASE) @@ -232,6 +253,7 @@ class WorkflowRAGFanoutResult: results: List[tuple[str, List[ChunkPreview]]] failed_count: int timeout_count: int = 0 + slowest_search_latency_ms: int | None = None @dataclass(frozen=True) @@ -1346,13 +1368,18 @@ def _run(self, inputs: Dict[str, Any]) -> Dict[str, Any]: temp_session = db_session candidate_resolution: KnowledgeRuntimeCandidateResolution | None = None + candidate_resolution_latency_ms: int | None = None if knowledge_enabled: + candidate_resolution_started_at = time.perf_counter() try: candidate_resolution = self._resolve_runtime_knowledge_candidates() except Exception: if temp_session is not None: temp_session.close() raise + candidate_resolution_latency_ms = self._elapsed_stage_latency_ms( + candidate_resolution_started_at + ) if not candidate_resolution.candidates: knowledge_result = self._knowledge_candidate_safe_no_result( candidate_resolution @@ -1445,6 +1472,9 @@ def _run(self, inputs: Dict[str, Any]) -> Dict[str, Any]: query=rag_search_query, db_session=db_session, candidate_resolution=candidate_resolution, + candidate_resolution_latency_ms=( + candidate_resolution_latency_ms + ), ) knowledge_context = knowledge_result.context knowledge_metadata = knowledge_result.metadata @@ -2603,6 +2633,7 @@ def _execute_knowledge_search( *, candidate_resolution: KnowledgeRuntimeCandidateResolution | None = None, query_embedding_plan: QueryEmbeddingPlan | None = None, + candidate_resolution_latency_ms: int | None = None, ) -> WorkflowRAGSearchResult: """ 연결된 지식 베이스에서 문서를 검색합니다. @@ -2624,11 +2655,14 @@ def _execute_knowledge_search( "RAG retrieval requires an active organization context." ) from exc - resolution = ( - candidate_resolution - if candidate_resolution is not None - else self._resolve_runtime_knowledge_candidates() - ) + if candidate_resolution is None: + candidate_resolution_started_at = time.perf_counter() + resolution = self._resolve_runtime_knowledge_candidates() + candidate_resolution_latency_ms = self._elapsed_stage_latency_ms( + candidate_resolution_started_at + ) + else: + resolution = candidate_resolution candidate_summary = self._knowledge_candidate_trace_summary(resolution) candidate_kind_by_kb_id = { str(candidate.knowledge_base_id): candidate.provenance.kind @@ -2638,6 +2672,7 @@ def _execute_knowledge_search( kb_ids = list(candidate_kind_by_kb_id) if not kb_ids: return self._knowledge_candidate_safe_no_result(resolution) + query_embedding_started_at = time.perf_counter() if query_embedding_plan is None: query_embedding_plan = self._preflight_query_embedding( organization_id=organization_uuid, @@ -2674,13 +2709,16 @@ def _execute_knowledge_search( organization_id=organization_uuid, knowledge_base_ids=kb_ids, ) + query_embedding_latency_ms = self._elapsed_stage_latency_ms( + query_embedding_started_at + ) fanout_kb_ids = kb_ids if precomputed_vectors: fanout_kb_ids = [kb_id for kb_id in kb_ids if kb_id in query_vectors_by_kb] + fanout_started_at = time.perf_counter() fanout = self._run_rag_retrieval_fanout( query=search_query, - fallback_db_session=db_session, user_id=retrieval_user_id, organization_id=organization_uuid, knowledge_base_ids=fanout_kb_ids, @@ -2691,11 +2729,15 @@ def _execute_knowledge_search( model_bindings_by_kb if precomputed_vectors else None ), ) + retrieval_fanout_latency_ms = self._elapsed_stage_latency_ms( + fanout_started_at + ) if embedding_failed_count: fanout = WorkflowRAGFanoutResult( results=fanout.results, failed_count=fanout.failed_count + embedding_failed_count, timeout_count=fanout.timeout_count, + slowest_search_latency_ms=fanout.slowest_search_latency_ms, ) for kb_id, chunks in fanout.results: @@ -2703,6 +2745,7 @@ def _execute_knowledge_search( for chunk in chunks: all_chunks.append((kb_id, chunk)) + evidence_policy_started_at = time.perf_counter() source_tier_policy = getattr(self.data, "sourceTierPolicy", "tie_break") # source_tier는 권한을 통과한 evidence 안에서만 동점 정렬 힌트로 사용한다. sorted_chunks = sorted( @@ -2742,6 +2785,10 @@ def _execute_knowledge_search( fanout.failed_count, successful_candidate_count=len(kb_ids) - fanout.failed_count, ) + policy_block_reason = blocked_evidence_reason_for_chunks(selected_chunks) + evidence_policy_latency_ms = self._elapsed_stage_latency_ms( + evidence_policy_started_at + ) trace_summary = self._rag_runtime_trace_summary( authorized_kb_count=len(kb_ids), retrieved_chunk_count=len(top_chunks), @@ -2753,7 +2800,15 @@ def _execute_knowledge_search( bucket_kb_counts=bucket_kb_counts, ) trace_summary.update(candidate_summary) - policy_block_reason = blocked_evidence_reason_for_chunks(selected_chunks) + trace_summary.update( + self._rag_stage_latency_summary( + candidate_resolution_latency_ms=candidate_resolution_latency_ms, + query_embedding_latency_ms=query_embedding_latency_ms, + retrieval_fanout_latency_ms=retrieval_fanout_latency_ms, + slowest_search_latency_ms=fanout.slowest_search_latency_ms, + evidence_policy_latency_ms=evidence_policy_latency_ms, + ) + ) if policy_block_reason: # 정책상 외부 LLM에 전달할 수 없는 evidence는 근거 충분성과 무관하게 차단한다. self._record_rag_policy_block_audit( @@ -2769,12 +2824,24 @@ def _execute_knowledge_search( results=[], failed_count=fanout.failed_count, timeout_count=fanout.timeout_count, + slowest_search_latency_ms=fanout.slowest_search_latency_ms, ), query_rewrite_applied=query_rewrite_applied, query_rewrite_strategy=query_rewrite_strategy, bucket_kb_counts=bucket_kb_counts, ) blocked_trace_summary.update(candidate_summary) + blocked_trace_summary.update( + self._rag_stage_latency_summary( + candidate_resolution_latency_ms=( + candidate_resolution_latency_ms + ), + query_embedding_latency_ms=query_embedding_latency_ms, + retrieval_fanout_latency_ms=retrieval_fanout_latency_ms, + slowest_search_latency_ms=fanout.slowest_search_latency_ms, + evidence_policy_latency_ms=evidence_policy_latency_ms, + ) + ) blocked_trace_summary["safe_exclusion_summary"] = { "policy_filtered": True, "reason_code": policy_block_reason, @@ -2901,7 +2968,6 @@ def _run_rag_retrieval_fanout( self, *, query: str, - fallback_db_session, user_id: uuid.UUID | None, organization_id: uuid.UUID, knowledge_base_ids: List[str], @@ -2919,131 +2985,69 @@ def _run_rag_retrieval_fanout( "kb_count_bucket=%s", self._bucket_count(len(knowledge_base_ids)), ) - return self._run_rag_retrieval_fanout_sequential( - query=query, - db_session=fallback_db_session, - user_id=user_id, - organization_id=organization_id, - knowledge_base_ids=knowledge_base_ids, - top_k=top_k, - threshold=threshold, - query_vectors_by_kb=query_vectors_by_kb, - model_bindings_by_kb=model_bindings_by_kb, - ) - gevent_modules = self._rag_gevent_modules() - if gevent_modules is None: - return self._run_rag_retrieval_fanout_sequential( + tasks = tuple( + RAGRetrievalFanoutTask(ordinal=index, resource_ref=knowledge_base_id) + for index, knowledge_base_id in enumerate(knowledge_base_ids) + ) + + def search( + task: RAGRetrievalFanoutTask, + cancellation: RAGRetrievalCancellation, + timeout_ms: int, + ) -> List[ChunkPreview]: + knowledge_base_id = task.resource_ref + return self._search_single_rag_kb_with_new_session( query=query, - db_session=fallback_db_session, user_id=user_id, organization_id=organization_id, - knowledge_base_ids=knowledge_base_ids, + knowledge_base_id=knowledge_base_id, top_k=top_k, threshold=threshold, - query_vectors_by_kb=query_vectors_by_kb, - model_bindings_by_kb=model_bindings_by_kb, + query_vector=( + query_vectors_by_kb.get(knowledge_base_id) + if query_vectors_by_kb + else None + ), + embedding_model_binding=( + model_bindings_by_kb.get(knowledge_base_id) + if model_bindings_by_kb + else None + ), + cancellation=cancellation, + timeout_ms=timeout_ms, ) - gevent, pool_cls = gevent_modules - pool = pool_cls(size=min(MAX_RAG_FANOUT_CONCURRENCY, len(knowledge_base_ids))) - jobs = { - pool.spawn( - self._search_single_rag_kb_with_new_session, - query=query, - user_id=user_id, - organization_id=organization_id, - knowledge_base_id=kb_id, - top_k=top_k, - threshold=threshold, - query_vector=query_vectors_by_kb.get(kb_id) - if query_vectors_by_kb - else None, - embedding_model_binding=model_bindings_by_kb.get(kb_id) - if model_bindings_by_kb - else None, - ): kb_id - for kb_id in knowledge_base_ids - } - gevent.joinall( - list(jobs), - timeout=RAG_FANOUT_AGGREGATE_TIMEOUT_SECONDS, + scheduled = self._get_rag_retrieval_fanout_scheduler().execute( + tasks=tasks, + worker=search, + fail_fast=self.data.ragFailurePolicy == "fail_node", ) - - results: List[tuple[str, List[ChunkPreview]]] = [] - failed_count = 0 - timeout_count = 0 - for job, kb_id in jobs.items(): - if not job.ready(): - timeout_count += 1 - failed_count += 1 - job.kill(block=False) - continue - if job.exception is not None: - if self.data.ragFailurePolicy == "fail_node": - pool.kill(block=False) - raise job.exception - if isinstance(job.exception, TimeoutError): - timeout_count += 1 - failed_count += 1 - continue - results.append((kb_id, job.value or [])) - - pool.kill(block=False) return WorkflowRAGFanoutResult( - results=results, - failed_count=failed_count, - timeout_count=timeout_count, + results=[ + (task.resource_ref, chunks) for task, chunks in scheduled.results + ], + failed_count=scheduled.failed_count, + timeout_count=scheduled.timeout_count, + slowest_search_latency_ms=scheduled.slowest_search_latency_ms, ) - def _run_rag_retrieval_fanout_sequential( - self, - *, - query: str, - db_session, - user_id: uuid.UUID | None, - organization_id: uuid.UUID, - knowledge_base_ids: List[str], - top_k: int, - threshold: float, - query_vectors_by_kb: Optional[Dict[str, List[float]]] = None, - model_bindings_by_kb: Optional[Dict[str, EmbeddingModelBinding]] = None, - ) -> WorkflowRAGFanoutResult: - retrieval = RetrievalService( - db_session, - user_id, - organization_id=organization_id, + def _get_rag_retrieval_fanout_scheduler(self) -> RAGRetrievalFanoutScheduler: + override = getattr(self, "_rag_retrieval_fanout_scheduler_override", None) + if override is not None: + return override + from apps.workflow_engine.adapters.rag_retrieval_executor import ( + GeventNativeThreadRAGRetrievalExecutor, + NativeThreadRAGRetrievalCancellation, ) - results: List[tuple[str, List[ChunkPreview]]] = [] - failed_count = 0 - timeout_count = 0 - for kb_id in knowledge_base_ids: - try: - chunks = self._search_single_rag_kb_with_timeout( - retrieval, - query=query, - knowledge_base_id=kb_id, - top_k=top_k, - threshold=threshold, - query_vector=query_vectors_by_kb.get(kb_id) - if query_vectors_by_kb - else None, - embedding_model_binding=model_bindings_by_kb.get(kb_id) - if model_bindings_by_kb - else None, - ) - except Exception as exc: - if self.data.ragFailurePolicy == "fail_node": - raise - if isinstance(exc, TimeoutError): - timeout_count += 1 - failed_count += 1 - continue - results.append((kb_id, chunks)) - return WorkflowRAGFanoutResult( - results=results, - failed_count=failed_count, - timeout_count=timeout_count, + + return RAGRetrievalFanoutScheduler( + executor_factory=GeventNativeThreadRAGRetrievalExecutor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + max_workers=MAX_RAG_FANOUT_CONCURRENCY, + per_task_timeout_seconds=RAG_FANOUT_PER_KB_TIMEOUT_SECONDS, + aggregate_timeout_seconds=RAG_FANOUT_AGGREGATE_TIMEOUT_SECONDS, + cleanup_reserve_seconds=1.0, ) def _search_single_rag_kb_with_new_session( @@ -3057,46 +3061,17 @@ def _search_single_rag_kb_with_new_session( threshold: float, query_vector: Optional[List[float]] = None, embedding_model_binding: Optional[EmbeddingModelBinding] = None, + cancellation: RAGRetrievalCancellation, + timeout_ms: int, ) -> List[ChunkPreview]: - session = SessionLocal() - try: + def search(session) -> List[ChunkPreview]: + from apps.workflow_engine.services.retrieval import RetrievalService + retrieval = RetrievalService( session, user_id, organization_id=organization_id, ) - return self._search_single_rag_kb_with_timeout( - retrieval, - query=query, - knowledge_base_id=knowledge_base_id, - top_k=top_k, - threshold=threshold, - query_vector=query_vector, - embedding_model_binding=embedding_model_binding, - ) - finally: - session.close() - - def _search_single_rag_kb_with_timeout( - self, - retrieval: RetrievalService, - *, - query: str, - knowledge_base_id: str, - top_k: int, - threshold: float, - query_vector: Optional[List[float]] = None, - embedding_model_binding: Optional[EmbeddingModelBinding] = None, - ) -> List[ChunkPreview]: - gevent_modules = self._rag_gevent_modules() - if gevent_modules is None: - # workflow_engine은 gevent 의존성을 갖는다. guard가 없으면 RAG 호출을 - # 무제한으로 붙잡지 않도록 operational failure로 닫는다. - raise TimeoutError("RAG retrieval timeout guard is unavailable.") - gevent, _pool_cls = gevent_modules - timer = gevent.Timeout(RAG_FANOUT_PER_KB_TIMEOUT_SECONDS) - timer.start() - try: return self._search_single_rag_kb( retrieval, query=query, @@ -3106,16 +3081,32 @@ def _search_single_rag_kb_with_timeout( query_vector=query_vector, embedding_model_binding=embedding_model_binding, ) - except gevent.Timeout as exc: - if exc is timer: - raise TimeoutError("RAG retrieval timed out.") from None - raise - finally: - timer.cancel() + + return self._get_rag_retrieval_session_runner().run( + timeout_ms=timeout_ms, + cancellation=cancellation, + operation=search, + ) + + def _get_rag_retrieval_session_runner(self) -> "RAGRetrievalSessionRunner": + override = getattr(self, "_rag_retrieval_session_runner_override", None) + if override is not None: + return override + from apps.workflow_engine.adapters.rag_retrieval_session import ( + RAGRetrievalSessionError, + RAGRetrievalSessionRunner, + ) + + session_factory = self.execution_context.get("db_session_factory") + if session_factory is None: + session_factory = SessionLocal + elif not callable(session_factory): + raise RAGRetrievalSessionError() + return RAGRetrievalSessionRunner(session_factory=session_factory) def _search_single_rag_kb( self, - retrieval: RetrievalService, + retrieval: "RetrievalService", *, query: str, knowledge_base_id: str, @@ -3242,16 +3233,6 @@ def _render_rag_query_template(self, template: str, query: str) -> str: rendered = f"{query} {template}" return " ".join(rendered.split()) - @staticmethod - def _rag_gevent_modules(): - try: - import gevent - from gevent.pool import Pool - - return gevent, Pool - except ImportError: - return None - def _rag_runtime_trace_summary( self, *, @@ -3481,6 +3462,7 @@ def _knowledge_candidate_safe_no_result( self, resolution: KnowledgeRuntimeCandidateResolution, ) -> WorkflowRAGSearchResult: + trace_summary = self._knowledge_candidate_trace_summary(resolution) return WorkflowRAGSearchResult( context="", metadata=[], @@ -3490,9 +3472,36 @@ def _knowledge_candidate_safe_no_result( ), should_invoke_llm=False, answer_override=RAG_NO_EVIDENCE_MESSAGE, - trace_summary=self._knowledge_candidate_trace_summary(resolution), + trace_summary=trace_summary, ) + @staticmethod + def _elapsed_stage_latency_ms(started_at: float) -> int: + elapsed_ms = max(0, int((time.perf_counter() - started_at) * 1000)) + return min(elapsed_ms, MAX_RAG_STAGE_LATENCY_MS) + + @staticmethod + def _rag_stage_latency_summary( + *, + candidate_resolution_latency_ms: int | None = None, + query_embedding_latency_ms: int | None = None, + retrieval_fanout_latency_ms: int | None = None, + slowest_search_latency_ms: int | None = None, + evidence_policy_latency_ms: int | None = None, + ) -> Dict[str, int]: + values = { + "candidate_resolution_latency_ms": candidate_resolution_latency_ms, + "query_embedding_latency_ms": query_embedding_latency_ms, + "retrieval_fanout_latency_ms": retrieval_fanout_latency_ms, + "slowest_search_latency_ms": slowest_search_latency_ms, + "evidence_policy_latency_ms": evidence_policy_latency_ms, + } + return { + key: min(value, MAX_RAG_STAGE_LATENCY_MS) + for key, value in values.items() + if type(value) is int and value >= 0 + } + @staticmethod def _knowledge_candidate_trace_summary( resolution: KnowledgeRuntimeCandidateResolution, @@ -3663,6 +3672,8 @@ def _rag_result_metadata( if isinstance(knowledge_result.trace_summary, dict) else {} ) + for key in RAG_TRACE_STAGE_LATENCY_FIELDS: + summary.pop(key, None) summary.update(self._rag_evidence_summary(knowledge_result.evidence_decision)) return summary diff --git a/docs/features/audit-tracing/requirements.md b/docs/features/audit-tracing/requirements.md index 77fa99128..6e9bdd4f4 100644 --- a/docs/features/audit-tracing/requirements.md +++ b/docs/features/audit-tracing/requirements.md @@ -74,8 +74,9 @@ Audit와 trace는 workflow 실행, RAG retrieval, LLM 호출, permission/policy ## Policies And Edge Cases - Success/authorized path와 hidden/denied/resource-hidden path의 allowlist를 분리한다. -- Authorized retrieval/answer path는 redaction-safe id와 summary만 저장할 수 있다. 허용되는 값은 KB id, document version id, chunk id, citation id, optional collection id, score/rank summary, safe metadata summary, policy result, latency/cost/token aggregate, retryability, opaque correlation/request id, `retrieval_strategy`, `rag_mode`, authorized/selected/retrieved count summary, `context_token_estimate`, `permission_filter_applied`, `safe_exclusion_summary`, `query_rewrite_applied`, `query_rewrite_strategy`, `evidence_sufficient`, `insufficiency_reason`, `source_tier_policy`, `source_tier_used`, `fanout_concurrency`, `fanout_timeout_seconds`, `failure_policy`다. Raw rewritten query는 저장하지 않는다. +- Authorized retrieval/answer path는 redaction-safe id와 summary만 저장할 수 있다. 허용되는 값은 KB id, document version id, chunk id, citation id, optional collection id, score/rank summary, safe metadata summary, policy result, latency/cost/token aggregate, retryability, opaque correlation/request id, `retrieval_strategy`, `rag_mode`, authorized/selected/retrieved count summary, `context_token_estimate`, `permission_filter_applied`, `safe_exclusion_summary`, `query_rewrite_applied`, `query_rewrite_strategy`, `evidence_sufficient`, `insufficiency_reason`, `source_tier_policy`, `source_tier_used`, `fanout_concurrency`, `fanout_timeout_seconds`, `failure_policy`와 RAG aggregate stage latency인 `candidate_resolution_latency_ms`, `query_embedding_latency_ms`, `retrieval_fanout_latency_ms`, `slowest_search_latency_ms`, `evidence_policy_latency_ms`다. Stage latency는 0~300,000 범위의 finite non-negative integer만 허용하고 실행하지 않은 stage는 생략한다. Raw rewritten query, per-KB timing과 hidden target identity는 저장하지 않는다. - Hidden/denied/resource-hidden path는 sanitized reason class, request/correlation id, actor/org scope, 필요한 경우 coarse retryability, audit action/status만 저장한다. Raw source id/url/path/title, raw source principal, source distribution, raw ACL fact, exact hidden/denied count, raw query, raw answer, raw prompt/completion, content preview, raw exception은 저장하지 않는다. +- Candidate 0건, empty query 등 hidden/resource-hidden 상태와 구분하지 않는 RAG `safe_no_result` 경로는 exact stage latency를 저장하지 않는다. Stage latency는 하나 이상의 authorized candidate가 실제 retrieval 경로에 진입한 경우에만 허용한다. - Partial result audit/trace는 `partial_result=true`, bucketed reason summary, retryability, correlation/request id만 저장한다. - Prompt trace와 RAG retrieval trace는 서로 다른 payload kind를 사용할 수 있지만 raw evidence 저장 금지 기준은 동일하게 적용한다. RAG context block, citation preview, chunk body를 디버깅 편의 목적으로 durable trace에 복사하지 않는다. - Source ACL mapping audit는 safe principal reference만 저장하고 raw email/path/title/url은 저장하지 않는다. diff --git a/docs/features/audit-tracing/test_cases.md b/docs/features/audit-tracing/test_cases.md index 5b106c676..936a4c97f 100644 --- a/docs/features/audit-tracing/test_cases.md +++ b/docs/features/audit-tracing/test_cases.md @@ -21,11 +21,12 @@ Status: Draft - 유효한 `organization_id`와 `workflow_run_id`가 있지만 비동기 `log.create_run`이 아직 WorkflowRun을 저장하지 않은 경우 Outbox worker는 safe reason으로 bounded retry한다. 재시도 중 Run이 보이면 typed correlation을 보존하고, 마지막 시도에도 없으면 correlation만 NULL로 내려 canonical AuditLog 자체는 저장한다. - 수동/API/Webhook 사용자 WorkflowRun 완료·실패 audit은 Workflow의 canonical `organization_id`와 `workflow_run_id`를 함께 기록해 Outbox 저장 뒤에도 typed run 연결을 유지한다. System schedule은 기존 exact claim provenance 검증을 계속 사용한다. - Audit metadata sanitizer는 raw source id/url/path/title, raw source principal, raw source ACL row, raw chunk content, raw prompt/completion, credential value, `encrypted_config`, raw exception을 제거한다. -- Authorized retrieval summary allowlist는 KB id, document version id, chunk id, citation id, optional collection id, rank/score, safe metadata summary, policy result, latency/cost/token aggregate, retryability, opaque correlation/request id, `retrieval_strategy`, `rag_mode`, authorized/selected/retrieved count summary, `context_token_estimate`, `permission_filter_applied`, `safe_exclusion_summary`, `query_rewrite_applied`, `query_rewrite_strategy`, `evidence_sufficient`, `insufficiency_reason`, `source_tier_policy`, `source_tier_used`, `fanout_concurrency`, `fanout_timeout_seconds`, `failure_policy`만 허용한다. +- Authorized retrieval summary allowlist는 KB id, document version id, chunk id, citation id, optional collection id, rank/score, safe metadata summary, policy result, latency/cost/token aggregate, retryability, opaque correlation/request id, `retrieval_strategy`, `rag_mode`, authorized/selected/retrieved count summary, `context_token_estimate`, `permission_filter_applied`, `safe_exclusion_summary`, `query_rewrite_applied`, `query_rewrite_strategy`, `evidence_sufficient`, `insufficiency_reason`, `source_tier_policy`, `source_tier_used`, `fanout_concurrency`, `fanout_timeout_seconds`, `failure_policy`와 `candidate_resolution_latency_ms`, `query_embedding_latency_ms`, `retrieval_fanout_latency_ms`, `slowest_search_latency_ms`, `evidence_policy_latency_ms`만 허용한다. RAG stage latency는 0~300,000 범위 finite non-negative integer만 보존하며 bool, 음수, NaN/Infinity, 문자열, 상한 초과, unknown/per-KB timing을 제거한다. - Audit/trace metadata sanitizer는 raw rewritten query를 raw prompt와 같은 민감 입력으로 보고 durable metadata와 log에서 제거한다. - 공통 Trace redaction은 숫자 또는 null인 명시적 token limit/usage allowlist(`max_tokens`, prompt/completion/input/output/total token count 등)는 보존하되, 인증 token 문자열과 알 수 없는 `*_token`, 문자열로 들어온 token count, 정책이 지정한 민감 경로는 계속 마스킹한다. - Conversation Memory sanitizer는 raw transcript/Memory content, Access Grant token/hash, prompt, private source identity와 provider raw error를 제거하고 session/grant lifecycle의 safe opaque reference, audience, status, reason과 bucketed count만 허용한다. - Hidden/denied/resource-hidden summary allowlist는 sanitized reason class, actor/org scope, request/correlation id, coarse retryability만 허용하고 exact hidden/denied count나 hidden KB id를 거부한다. +- Candidate 0건 또는 empty query의 RAG `safe_no_result` trace는 `candidate_resolution_latency_ms`를 포함한 exact stage latency를 모두 생략한다. - Partial result summary는 `partial_result=true`, bucketed reason summary, retryability, request/correlation id만 허용한다. - Run trigger pure policy는 `manual`/`test`/`manual_compare`/`cost_optimizer_compare`를 MANUAL, `api`/`app`/`deployed`/`api_secret`을 API, `webhook`을 WEBHOOK, `schedule`/`scheduler`를 SCHEDULER로 정규화한다. String은 trim/lowercase하고, Log System adapter가 이미 받은 `RunTriggerMode` enum은 string alias policy를 거치지 않고 exact 값을 보존한다. - Trigger 누락 또는 `None`은 legacy compatibility로 deployed면 API, 아니면 MANUAL이지만 blank/unknown/non-string explicit input은 fallback하지 않고 permanent contract error다. diff --git a/docs/features/knowledge/component_spec.md b/docs/features/knowledge/component_spec.md index cad415bcb..e09af0681 100644 --- a/docs/features/knowledge/component_spec.md +++ b/docs/features/knowledge/component_spec.md @@ -1,7 +1,7 @@ # Knowledge Component Spec Status: Draft -Verified Against: feature/mba-359 @ 504f418ac708a2dc541a5283f1ad8e97da0869a2 +Verified Against: feature/mba-354 @ 0dcc92377bd9632a1cb2c218472c164462f37810 MBA-105 구현 baseline, 운영 기본값, permission helper output, active version finalization, resource hiding matrix는 [implementation_baseline.md](implementation_baseline.md)를 따른다. Workflow RAG에서 `execution_subject`가 없는 MVP public-only runtime은 [ADR-0018](../../decisions/ADR-0018-workflow-rag-anonymous-public-only-runtime.md)을 따른다. MCP/API source connector와 incremental sync 경계는 [ADR-0020](../../decisions/ADR-0020-knowledge-mcp-incremental-sync-boundary.md)을 따른다. Direct KB와 명시 selected Collection의 Workflow runtime candidate 해석은 [ADR-0036](../../decisions/ADR-0036-knowledge-runtime-candidate-resolution.md)을 따른다. KC lifecycle, item 순서와 권한 운영 경계는 [ADR-0044](../../decisions/ADR-0044-knowledge-collection-operational-management-boundary.md)을 따른다. KC sync의 Gateway application, durable repository, Workflow executor와 Client polling 경계는 [ADR-0048](../../decisions/ADR-0048-knowledge-collection-sync-execution-boundary.md)을 따른다. Organization Detector Provider와 embedding 전 local masking Target은 [ADR-0070](../../decisions/ADR-0070-organization-detector-provider-and-pre-embedding-local-masking-boundary.md)을 따른다. 현재 component/runtime 구현 완료를 뜻하지 않는다. @@ -761,8 +761,12 @@ tombstone cleanup은 구현 전에 별도 retention policy, audit action/reason - 초기 candidate cap은 `max_candidate_kbs=5000`, `max_route_collections=20`, `max_retrieval_kbs=20`, `max_chunks_per_kb=8`, `max_total_chunks=50`이다. 이 값은 운영 baseline이며 제품의 고정 계약이 아니다. - Candidate cap, fanout concurrency, timeout, partial failure behavior는 [implementation_baseline.md](implementation_baseline.md)의 baseline을 시작점으로 삼고, operations policy로 조정 가능해야 하며 운영 배포 전에 load test를 거쳐야 한다. - 가능한 경우 KB/version filter를 포함한 단일 vector/keyword query를 우선한다. Backend가 지원하지 못하면 concurrency와 timeout cap이 있는 bounded per-KB fanout을 사용한다. +- Workflow의 bounded per-KB fanout은 authorized candidate와 사전 계산 query vector만 받는 application scheduler가 소유한다. 현재 baseline은 invocation당 동시 검색 최대 5개이면서 Workflow Engine 프로세스 전체 native blocking-I/O data worker와 제출 admission도 각각 최대 5개다. Executor는 callback greenlet을 만들기 전에 invocation/process slot을 모두 예약하고, 프로세스 포화 시 대기 greenlet이나 native queue를 추가하지 않은 채 해당 child를 기존 partial-failure 정책으로 닫는다. 완료 또는 실패한 job은 두 slot을 반드시 반환한다. KB별 최대 10초, 최초 task 제출 전부터 시작하는 caller aggregate 30초이며 마지막 1초는 cancellation과 transaction 정리를 기다리는 데 예약한다. Scheduler queue 대기와 cleanup 대기를 aggregate deadline에 포함하고 deadline 뒤 새 KB search를 시작하지 않는다. +- 각 시작된 KB search는 candidate resolution과 query embedding runtime에 사용한 것과 동일한 주입 `db_session_factory`에서 독립 SQLAlchemy session을 만들고 PostgreSQL read-only transaction에서 `organization_id + knowledge_base_id`로 KB를 조회한다. 명시적으로 주입된 factory가 유효하지 않으면 전역 `SessionLocal`로 바꾸지 않고 fail-closed한다. Session factory 실행과 connection checkout은 프로세스 전체 최대 5개의 bounded native acquisition worker가 소유한다. Task cancellation/deadline 또는 acquisition slot 포화는 caller와 RAG data worker를 즉시 typed timeout으로 반환하고, 늦게 획득된 session은 acquisition worker가 rollback/close한다. 강제 thread 종료나 공용 engine의 `pool_timeout` mutation은 하지 않는다. Adapter는 task 시작 시 단조시계 기준 절대 deadline을 고정하고 operation이 발생시키는 모든 SQL 직전에 cancellation과 남은 budget을 재검증한다. PostgreSQL transaction-local `statement_timeout`은 각 SQL마다 남은 정수 millisecond 이하로 축소하므로 여러 statement가 각각 최초 timeout을 새로 사용할 수 없다. Guard는 operation 뒤 connection에서 제거한 다음 transaction을 rollback하고 session을 close한다. Timeout/cancel의 DBAPI query cancel은 data worker와 분리된 프로세스 전체 최대 2개의 native control worker에서 실행해 gevent hub를 막지 않는다. 등록 해제된 callback은 실행하지 않고, 이미 실행 중인 callback은 완료된 뒤에만 해당 session을 rollback/close하여 pool로 반환된 연결에 늦은 cancel이 도달하지 않게 한다. Outer Workflow session은 native worker와 공유하지 않는다. Scheduler는 cleanup reserve 동안 종료를 기다리되 협조하지 않는 non-DB 작업 때문에 caller hard deadline을 연장하지 않는다. Hard deadline 뒤 완료된 task는 evidence를 게시할 수 없다. +- Fanout completion 순서는 evidence 결과를 바꾸지 않는다. Scheduler는 candidate ordinal 기준으로 결과를 반환하고, 기존 global score/source-tier 정렬, dedupe, top-k, final evidence policy와 citation projection이 최종 순서를 결정한다. `safe_no_result`는 성공 evidence를 유지할 수 있지만 `fail_node`는 남은 task를 취소하고 partial evidence를 사용하지 않는다. +- Authorized RAG retrieval trace는 `candidate_resolution_latency_ms`, `query_embedding_latency_ms`, `retrieval_fanout_latency_ms`, `slowest_search_latency_ms`, `evidence_policy_latency_ms`만 aggregate stage latency로 허용한다. 값은 0~300,000 범위의 finite non-negative integer millisecond이며 실행하지 않은 stage는 생략한다. Candidate 0건, empty query 등 hidden/resource-hidden 상태와 구분하지 않는 `safe_no_result` 경로는 exact stage latency를 모두 생략한다. Per-KB timing, raw query/vector, hidden resource identity와 provider/DB raw error는 저장하지 않고 이 다섯 필드는 일반 Workflow result metadata, chatbot/SSE와 citation projection에 포함하지 않는다. - Workflow LLM node는 `context_variable`이 지정된 경우 해당 referenced variable의 정제된 값만 ephemeral retrieval query로 사용한다. 설정이 없는 legacy graph만 렌더링된 user prompt 전체를 사용하며 raw query는 durable trace, audit, log 또는 cache key에 저장하지 않는다. -- Opt-in CrossEncoder rerank는 권한을 통과한 chunk의 redacted canonical text를 메모리에서 복호화한 뒤 사용한다. 저장 암호문을 ranking model input으로 전달하거나 복호화 실패 때 암호문으로 fallback하지 않는다. +- Opt-in CrossEncoder rerank는 권한을 통과한 chunk의 redacted canonical text를 메모리에서 복호화한 뒤 사용한다. 프로세스 cache의 최초 또는 model 변경 초기화는 원본 native lock으로 직렬화하고 완성된 model만 게시한다. 저장 암호문을 ranking model input으로 전달하거나 복호화 실패 때 암호문으로 fallback하지 않는다. - MBA-232 runtime resolver는 candidate ID/authorization을 invocation 사이에 cache하지 않는다. 향후 candidate cache를 별도 승인할 경우 permission/freshness revision을 포함해 ACL revocation이 stale candidate를 무효화해야 한다. - Skill candidate cache key에는 skill version, freshness state, eval state, source version reference를 포함해 stale skill이나 source tier 변경이 즉시 무효화되어야 한다. - Query rewrite cache를 둘 경우 key에는 rewrite mode, safe template id, skill version, permission/freshness epoch를 포함해야 하며 raw rewritten query를 durable cache key나 trace key로 사용하지 않는다. diff --git a/docs/features/knowledge/implementation_baseline.md b/docs/features/knowledge/implementation_baseline.md index 4efebc804..2e59eca2a 100644 --- a/docs/features/knowledge/implementation_baseline.md +++ b/docs/features/knowledge/implementation_baseline.md @@ -46,8 +46,8 @@ MBA-105에서 구현하지 않는 범위: | Chunk size / overlap | child chunk 800-1,200 tokens, overlap 10-20% | | Parent/child hierarchy | parent 2,000-4,000 tokens, child 500-1,000 tokens | | Candidate caps | `max_candidate_kbs=5000`, `max_route_collections=20`, `max_retrieval_kbs=20`, `max_chunks_per_kb=8`, `max_total_chunks=50`. Collection/KB candidate cap은 임의 row를 먼저 자른 뒤 authorization하는 방식이 아니라, route/use/source ACL helper를 통과한 authorized subset에 적용한다 | -| Fanout | 단일 filtered vector/keyword query 우선, 불가하면 concurrency 5 bounded fanout | -| Retrieval timeout | 호출당 5-10s, aggregate interactive path 15-30s | +| Fanout | 단일 filtered vector/keyword query 우선. Workflow의 per-KB fallback은 invocation당 동시 검색, 프로세스 전체 native blocking-I/O data worker와 제출 admission을 각각 최대 5개로 제한하고 authorized candidate ordinal을 보존한다. Executor는 greenlet 생성 전에 invocation/process slot을 예약하고 포화 시 추가 대기열을 만들지 않은 채 기존 partial-failure 정책으로 닫는다. Session factory/connection checkout도 별도 프로세스 전체 최대 5개의 bounded native acquisition worker로 제한하며 deadline 뒤 늦은 session은 획득 thread가 정리한다. DB cancel은 별도 bounded native control worker에서 처리한다 | +| Retrieval timeout | Workflow per-KB fallback은 호출당 10s, 최초 제출 전부터 caller aggregate 30s, 마지막 1s cleanup reserve. Queue 대기와 cleanup 대기를 aggregate에 포함한다. 각 DB worker는 모든 SQL 직전에 동일한 task 절대 deadline과 cancellation을 재검증하고 statement timeout을 남은 budget으로 축소한다. Non-DB 작업이 cancellation에 협조하지 않아도 caller deadline을 연장하지 않고 late result를 폐기한다 | | Runtime authorization batch | `check_access_batch` 50-200 source item 후보. Batch 미지원 source는 bounded single check fallback만 허용 | | Runtime authorization fallback | per-source concurrency 3-5, per-call timeout 3-5s, aggregate timeout 10-20s 후보. Timeout/unknown은 private evidence fail-closed | | Runtime access cache | `allowed`/`denied` 1-5m, `unknown`/timeout 30-60s 후보. Source ACL/mapping epoch 변경 시 즉시 무효화 | @@ -326,7 +326,7 @@ Hidden/resource-hidden path의 external `reason_code`는 항상 `resource.hidden - Legal/regulatory evidence는 strict use 전에 jurisdiction, effective date, version, review-required state를 가져야 한다. - Query rewrite 기본값은 `off`다. MBA-105 runtime은 deterministic/template rewrite를 opt-in으로 구현하고, LLM-assisted rewrite는 LLMOps/cost/security gate가 닫히기 전까지 범위 밖이다. - Raw rewritten query는 raw prompt와 같은 민감 입력으로 보고 durable audit, trace, usage, cache key, summary에 저장하지 않는다. -- CrossEncoder rerank는 기본값 `off`다. Worker image에서 reranker dependency를 포함하려면 `INSTALL_RAG_RERANKER=true` build arg를 명시하고, runtime에서 `RAG_CROSS_ENCODER_RERANK_ENABLED=true`를 설정해야 한다. dependency가 없거나 model load/predict가 실패하면 retrieval은 원래 candidate 순서로 fallback하며, 이 fallback은 operational error가 아니라 degraded ranking path로 기록한다. 기본 local/demo/production path는 cold start와 image size를 피하기 위해 reranker를 로드하지 않는다. +- CrossEncoder rerank는 기본값 `off`다. Worker image에서 reranker dependency를 포함하려면 `INSTALL_RAG_RERANKER=true` build arg를 명시하고, runtime에서 `RAG_CROSS_ENCODER_RERANK_ENABLED=true`를 설정해야 한다. 프로세스 cache 초기화는 native worker 사이에서 직렬화하고 완성된 model만 게시한다. Dependency가 없거나 model load/predict가 실패하면 retrieval은 원래 candidate 순서로 fallback하며, 이 fallback은 operational error가 아니라 degraded ranking path로 기록한다. 기본 local/demo/production path는 cold start와 image size를 피하기 위해 reranker를 로드하지 않는다. Slack/meeting source item의 effective ACL baseline: diff --git a/docs/features/knowledge/test_cases.md b/docs/features/knowledge/test_cases.md index 3857cb9dd..627b08bf3 100644 --- a/docs/features/knowledge/test_cases.md +++ b/docs/features/knowledge/test_cases.md @@ -1,7 +1,7 @@ # Knowledge Test Cases Status: Draft -Verified Against: feature/mba-359 @ 504f418ac708a2dc541a5283f1ad8e97da0869a2 +Verified Against: feature/mba-354 @ 0dcc92377bd9632a1cb2c218472c164462f37810 이 문서는 현재 RAG 동작과 목표 KB 통합 모델에 필요한 테스트 범위를 함께 기록한다. MBA-105 목표 모델 테스트는 [ADR-0017](../../decisions/ADR-0017-knowledge-integration-provisional-implementation-baseline.md)과 [implementation_baseline.md](implementation_baseline.md)의 임시 baseline을 기준으로 구현 blocker가 된다. KC sync의 실행·복구·snapshot·versioned finalization 검증은 [ADR-0048](../../decisions/ADR-0048-knowledge-collection-sync-execution-boundary.md)을 따른다. Organization Detector Provider와 embedding 전 local masking Target 테스트는 [ADR-0070](../../decisions/ADR-0070-organization-detector-provider-and-pre-embedding-local-masking-boundary.md)을 따른다. MBA-362가 runtime/persistence/provider adapter 테스트를 TDD로 구현하기 전에는 완료 증거가 아니다. @@ -525,6 +525,19 @@ Organization Detector Provider와 embedding 전 local masking Target 테스트 - Query count는 회귀 차단 기준이며, 실제 PostgreSQL 환경에서는 변경 전 per-KB lookup 대비 query 수와 latency를 관찰한다. 환경 변동이 큰 단일 latency threshold는 merge gate로 사용하지 않는다. - MBA-289가 Authorized Retrieval Port를 도입할 때 projection은 권한 판정을 복제하지 않고 authorized candidate 이후 adapter 내부 단계로 이동해야 한다. +### MBA-354 Precomputed Query Embedding KB Fanout + +- Authorized candidate가 0개면 query embedding, fanout scheduler, child DB session과 generation provider를 모두 호출하지 않는다. Candidate가 있으면 ADR-0071의 distinct-model capability/provider attempt와 invocation-local vector를 그대로 소비하며 KB별 embedding을 다시 만들지 않는다. +- 사전 계산 query vector를 사용하는 2개 이상 KB search는 실제 native worker에서 겹쳐 실행된다. Invocation당 동시 검색과 여러 invocation을 합친 프로세스 전체 native data worker 및 제출 admission이 모두 5개를 넘지 않아야 한다. Callback greenlet 생성 전에 두 admission을 예약하고, 공용 pool 포화 시 추가 callback을 실행하거나 대기 greenlet을 만들지 않아야 한다. 기존 job이 끝나면 slot이 회수되어 후속 제출이 성공해야 한다. Barrier 기반 test로 overlap을, 두 executor 동시 제출과 포화·회복 test로 process-wide 상한을 검증하며 우연한 wall-clock 단축만 성공 기준으로 사용하지 않는다. +- Aggregate deadline은 executor/task 제출 전에 시작한다. Queue 대기, active search, cancellation과 cleanup 대기를 포함해 caller 기준 30초를 넘기지 않고 마지막 1초에는 새 DB search를 시작하지 않는다. Running task 없이 queued task의 start budget만 부족한 경우 busy-spin 없이 timeout으로 수렴한다. Cancellation에 협조하지 않는 worker를 기다리느라 caller deadline을 연장하지 않으며 late result를 폐기하는 test를 포함한다. `gevent.Timeout` 등 `BaseException` 계열 외부 종료도 실행 중인 모든 child cancellation을 요청하고 원래 예외를 전파해야 한다. +- 각 worker는 candidate resolution/query embedding과 동일한 주입 `db_session_factory`에서 별도 session을 만들고 `organization_id + knowledge_base_id` 조건과 read-only transaction을 적용한다. Invalid explicit factory는 전역 DB fallback 없이 차단한다. Session factory와 connection checkout은 프로세스 전체 최대 5개의 native acquisition worker로 제한한다. Pool을 소진한 PostgreSQL test에서 per-KB deadline이 공용 `pool_timeout`보다 먼저 caller를 반환하고, 늦게 획득된 session은 획득 thread가 rollback/close하며 acquisition slot과 checked-out connection이 복구되어야 한다. 단조시계 절대 deadline을 기준으로 operation의 모든 SQL 직전에 cancellation과 남은 budget을 재검증하고 transaction-local statement timeout은 각 SQL마다 남은 budget 이하로 축소한다. 단일 timeout 안에서는 각각 성공할 두 `pg_sleep`의 합이 절대 deadline을 넘는 경우 후속 SQL이 timeout되고, 취소 뒤 다음 SQL은 DBAPI 실행 전에 차단되어야 한다. DB cancel callback은 프로세스 전체 최대 2개의 별도 native control worker에서 실행해 gevent hub를 막지 않아야 한다. Dispatch 뒤 등록 해제된 callback은 실행하지 않고 이미 실행 중인 callback은 완료 뒤에만 statement guard를 제거하고 session을 rollback/close한다. Success, empty, DB error, per-KB timeout, aggregate cancel과 rollback error 모두 close를 시도한다. Disposable PostgreSQL test는 checkout deadline, cumulative statement timeout, `pg_sleep` cancel 뒤 rollback, connection 재사용과 timeout/guard 비누출을 확인한다. +- Reverse completion, partial timeout과 mixed success에서도 candidate ordinal을 거쳐 기존 global score/source-tier 정렬, dedupe, top-k, evidence sufficiency와 citation 결과가 순차 기준 fixture와 같아야 한다. `fail_node`는 partial evidence를 사용하지 않고 남은 task에 cancellation을 요청한다. +- RetrievalService의 sync/async exception log와 fanout error projection에는 raw query, SQL/parameter, KB/document/chunk identity, provider payload와 raw exception 문자열이 없어야 한다. `fail_node`도 실패 정책은 유지하지만 child 원문 예외 대신 고정된 safe fanout error를 반환한다. +- 하나 이상의 authorized candidate가 retrieval에 진입한 RAG trace에는 실행된 stage의 `candidate_resolution_latency_ms`, `query_embedding_latency_ms`, `retrieval_fanout_latency_ms`, `slowest_search_latency_ms`, `evidence_policy_latency_ms`만 0~300,000 범위 integer로 남긴다. Candidate 0건 또는 empty query의 구분 불가능한 `safe_no_result`는 exact stage latency를 모두 생략한다. Unknown/negative/non-finite/bool/per-KB timing은 제거하고 일반 Workflow result metadata, chatbot/SSE와 citation에는 이 필드가 없어야 한다. +- 첫 multi-KB CrossEncoder rerank에서 여러 native worker가 동시에 model을 요청해도 constructor는 한 번만 실행되고 완성된 동일 cache instance를 사용해야 한다. +- Synthetic benchmark는 KB 1/2/4/10개에서 동일 delay profile의 p50/p95, 프로세스 전체 max active data worker와 search/provider call count를 기록한다. 절대 latency threshold는 merge gate가 아니며 overlap, concurrency, 호출 수와 deterministic 결과를 회귀 gate로 사용한다. Slow/timeout 동작은 별도의 deterministic scheduler test로 검증한다. +- `rag_retrieval_connection_acquirer.py` 또는 `rag_retrieval_executor.py`만 바뀌어도 Knowledge disposable PostgreSQL checkout/cancel/connection 복구 계약을 선택한다. LLM node import 경계, `NodeFactory` 또는 RAG adapter 변경은 Gateway architecture import-boundary test도 선택해 Gateway graph 검증이 worker-only `gevent` 의존성을 요구하지 않는지 확인한다. + ### MBA-288 Durable Document Ingestion - Process/sync/resume/reindex admission은 Document 또는 KB row lock 아래 설정·queued projection·job을 한 transaction에 저장한다. DB source process/sync는 owner Connection lock 뒤 fresh KB/Document lock과 revision compare를 수행하고 같은 lifecycle UoW로 reference metadata와 job을 commit한다. Approval resume의 첫 요청이 Document를 `waiting_for_approval`에서 `indexing`으로 바꾼 뒤 같은 요청을 반복해도 active 동일 intent를 재사용하며, 다른 intent만 conflict로 닫는다. Commit 실패에는 job과 Document 변경이 모두 없고, commit 뒤 broker publish 실패에는 pending job이 남아 recovery로 실행 가능해야 한다. diff --git a/scripts/ci/changed_scope.py b/scripts/ci/changed_scope.py index 71768fd36..94a06fb0f 100644 --- a/scripts/ci/changed_scope.py +++ b/scripts/ci/changed_scope.py @@ -32,12 +32,25 @@ "apps/shared/tests/domain/test_knowledge_runtime_candidates.py", "apps/shared/tests/services/test_knowledge_permission_runtime_bulk.py", "apps/workflow_engine/adapters/knowledge_runtime_candidates.py", + "apps/workflow_engine/adapters/rag_retrieval_connection_acquirer.py", + "apps/workflow_engine/adapters/rag_retrieval_executor.py", + "apps/workflow_engine/adapters/rag_retrieval_session.py", + "apps/workflow_engine/application/rag_retrieval_fanout.py", "apps/workflow_engine/application/runtime_retrieval/**", "apps/workflow_engine/composition/runtime_retrieval.py", "apps/workflow_engine/tests/adapters/test_postgres_knowledge_runtime_candidate_adapter.py", + "apps/workflow_engine/tests/adapters/test_rag_retrieval_session_postgres.py", ".github/workflows/test-knowledge-runtime-postgres.yml", ) +_GATEWAY_WORKFLOW_IMPORT_BOUNDARY_PATTERNS = ( + # Gateway graph validation imports these Workflow definitions without + # installing worker-only runtime dependencies such as gevent. + "apps/workflow_engine/adapters/rag_retrieval_*.py", + "apps/workflow_engine/workflow/core/workflow_node_factory.py", + "apps/workflow_engine/workflow/nodes/llm/**", +) + _WORKFLOW_POSTGRES_PATTERNS = ( "apps/gateway/adapters/db/schedule_dispatch_repository.py", "apps/gateway/adapters/audit/sqlalchemy_schedule_dispatch_audit.py", @@ -436,6 +449,8 @@ def classify_paths(raw_paths: Iterable[str]) -> ChangeScope: if _matches_any(path, _KNOWLEDGE_POSTGRES_PATTERNS): scope.knowledge_postgres = True + if _matches_any(path, _GATEWAY_WORKFLOW_IMPORT_BOUNDARY_PATTERNS): + scope.gateway_tests = True if _matches_any(path, _WORKFLOW_POSTGRES_PATTERNS): scope.workflow_postgres = True if _matches_any(path, _AGENT_BUILDER_POSTGRES_PATTERNS): diff --git a/tests/ci/test_changed_scope.py b/tests/ci/test_changed_scope.py index 7a6c3c42d..5eb75ef04 100644 --- a/tests/ci/test_changed_scope.py +++ b/tests/ci/test_changed_scope.py @@ -137,16 +137,41 @@ def test_agent_builder_usage_change_selects_agent_builder_postgres(path: str): assert scope.agent_builder_postgres is True -def test_knowledge_runtime_change_selects_knowledge_postgres(): - scope = classify_paths( - ["apps/workflow_engine/application/runtime_retrieval/knowledge_candidates.py"] - ) +@pytest.mark.parametrize( + "path", + [ + "apps/workflow_engine/application/runtime_retrieval/knowledge_candidates.py", + "apps/workflow_engine/application/rag_retrieval_fanout.py", + "apps/workflow_engine/adapters/rag_retrieval_connection_acquirer.py", + "apps/workflow_engine/adapters/rag_retrieval_executor.py", + "apps/workflow_engine/adapters/rag_retrieval_session.py", + "apps/workflow_engine/tests/adapters/test_rag_retrieval_session_postgres.py", + ], +) +def test_knowledge_runtime_change_selects_knowledge_postgres(path: str): + scope = classify_paths([path]) assert scope.workflow_tests is True assert scope.knowledge_postgres is True assert scope.workflow_postgres is False +@pytest.mark.parametrize( + "path", + [ + "apps/workflow_engine/workflow/core/workflow_node_factory.py", + "apps/workflow_engine/workflow/nodes/llm/__init__.py", + "apps/workflow_engine/adapters/rag_retrieval_connection_acquirer.py", + "apps/workflow_engine/adapters/rag_retrieval_executor.py", + ], +) +def test_worker_import_boundary_change_selects_gateway_contract(path: str): + scope = classify_paths([path]) + + assert scope.workflow_tests is True + assert scope.gateway_tests is True + + @pytest.mark.parametrize( "path", [ diff --git a/tests/ci/test_pr_quality_gate_workflow.py b/tests/ci/test_pr_quality_gate_workflow.py index 4565f91f2..e1f3b2ad9 100644 --- a/tests/ci/test_pr_quality_gate_workflow.py +++ b/tests/ci/test_pr_quality_gate_workflow.py @@ -1,7 +1,10 @@ from pathlib import Path import re -from scripts.ci.changed_scope import _WORKFLOW_POSTGRES_PATTERNS +from scripts.ci.changed_scope import ( + _KNOWLEDGE_POSTGRES_PATTERNS, + _WORKFLOW_POSTGRES_PATTERNS, +) REPOSITORY_ROOT = Path(__file__).resolve().parents[2] @@ -408,6 +411,19 @@ def test_knowledge_postgres_dev_push_tracks_all_durable_ingestion_services(): assert '- "apps/shared/services/knowledge_ingestion_*.py"' in push_paths +def test_knowledge_postgres_dev_push_covers_selector_patterns(): + workflow = KNOWLEDGE_POSTGRES_PATH.read_text(encoding="utf-8") + push_paths = workflow.split(" push:", maxsplit=1)[1].split( + "permissions:", + maxsplit=1, + )[0] + configured_paths = set( + re.findall(r'^\s*- "([^"]+)"', push_paths, flags=re.MULTILINE) + ) + + assert set(_KNOWLEDGE_POSTGRES_PATTERNS) <= configured_paths + + def test_workflow_postgres_dev_push_covers_selector_patterns(): workflow = WORKFLOW_POSTGRES_PATH.read_text(encoding="utf-8") push_paths = workflow.split(" push:", maxsplit=1)[1].split( diff --git a/tests/performance/README.md b/tests/performance/README.md index 6b0a4df87..48c62c714 100644 --- a/tests/performance/README.md +++ b/tests/performance/README.md @@ -15,3 +15,15 @@ Audit cursor 항목은 운영 목록과 동일하게 `audit_metadata ->> 'organi Trace 항목은 실제 visibility policy join 전체가 아니라 `5,000건 fetch 후 필터`와 `SQL에서 visible 20건 제한`의 조회량 차이를 단순화해 측정한다. 실제 `/api/v1/traces` 응답 시간으로 해석하지 않는다. 측정 결과는 `reports/`에 데이터 규모와 실행 날짜별 JSON으로 보관한다. 같은 날짜에 데이터 규모가 다르면 파일명에 `100k`처럼 행 수를 표시한다. + +## RAG Retrieval Fan-out + +사전 계산 query embedding을 사용하는 KB 1/2/4/10/20개 검색의 순차 기준과 bounded native-thread fan-out을 합성 blocking I/O로 비교한다. Provider, PostgreSQL과 실제 query/vector는 사용하지 않으며 절대 latency는 merge gate가 아니다. Overlap, 최대 active worker 5개, candidate당 검색 1회와 query embedding provider model당 1회 계약을 확인하기 위한 보조 측정이다. + +```bash +PYTHONPATH=$(git rev-parse --show-toplevel) apps/workflow_engine/.venv/bin/python \ + tests/performance/benchmark_rag_retrieval_fanout.py \ + --candidates 1,2,4,10,20 --delay-ms 20 --iterations 20 +``` + +출력에는 p50/p95, 최대 active worker와 호출 수만 포함한다. Raw query, vector, resource identifier와 provider/DB 오류는 기록하지 않는다. diff --git a/tests/performance/benchmark_rag_retrieval_fanout.py b/tests/performance/benchmark_rag_retrieval_fanout.py new file mode 100644 index 000000000..0597794a3 --- /dev/null +++ b/tests/performance/benchmark_rag_retrieval_fanout.py @@ -0,0 +1,150 @@ +"""Synthetic benchmark for bounded blocking RAG retrieval fan-out.""" + +from __future__ import annotations + +import argparse +import json +import math +import statistics +import threading +import time + +from apps.workflow_engine.application.rag_retrieval_fanout import ( + RAGRetrievalFanoutScheduler, + RAGRetrievalFanoutTask, +) +from apps.workflow_engine.adapters.rag_retrieval_executor import ( + GeventNativeThreadRAGRetrievalExecutor, + NativeThreadRAGRetrievalCancellation, +) + + +def _percentile(values: list[float], percentile: float) -> float: + ordered = sorted(values) + index = max(0, math.ceil(percentile * len(ordered)) - 1) + return ordered[index] + + +def _summarize(values: list[float]) -> dict[str, float]: + return { + "p50_ms": round(statistics.median(values), 3), + "p95_ms": round(_percentile(values, 0.95), 3), + } + + +def _sequential_elapsed_ms(candidate_count: int, delay_seconds: float) -> float: + started = time.perf_counter() + for _index in range(candidate_count): + time.sleep(delay_seconds) + return (time.perf_counter() - started) * 1000 + + +def _fanout_elapsed_ms( + candidate_count: int, + delay_seconds: float, +) -> tuple[float, int, int]: + state_lock = threading.Lock() + active = 0 + max_active = 0 + search_calls = 0 + + def search(_task, _cancellation, _timeout_ms): + nonlocal active, max_active, search_calls + with state_lock: + active += 1 + search_calls += 1 + max_active = max(max_active, active) + try: + time.sleep(delay_seconds) + return None + finally: + with state_lock: + active -= 1 + + tasks = tuple( + RAGRetrievalFanoutTask(index, f"benchmark-{index}") + for index in range(candidate_count) + ) + started = time.perf_counter() + result = RAGRetrievalFanoutScheduler( + executor_factory=GeventNativeThreadRAGRetrievalExecutor, + cancellation_factory=NativeThreadRAGRetrievalCancellation, + aggregate_timeout_seconds=5, + cleanup_reserve_seconds=0.1, + per_task_timeout_seconds=2, + minimum_start_budget_ms=1, + ).execute(tasks=tasks, worker=search) + elapsed_ms = (time.perf_counter() - started) * 1000 + if result.failed_count or len(result.results) != candidate_count: + raise RuntimeError("synthetic RAG fan-out benchmark did not complete") + return elapsed_ms, max_active, search_calls + + +def run(*, candidates: tuple[int, ...], delay_ms: int, iterations: int) -> dict: + delay_seconds = delay_ms / 1000 + results = [] + for candidate_count in candidates: + _sequential_elapsed_ms(candidate_count, delay_seconds) + _fanout_elapsed_ms(candidate_count, delay_seconds) + sequential_timings = [] + fanout_timings = [] + max_active = 0 + search_calls = 0 + for _index in range(iterations): + sequential_timings.append( + _sequential_elapsed_ms(candidate_count, delay_seconds) + ) + elapsed_ms, observed_active, observed_calls = _fanout_elapsed_ms( + candidate_count, + delay_seconds, + ) + fanout_timings.append(elapsed_ms) + max_active = max(max_active, observed_active) + search_calls += observed_calls + results.append( + { + "candidate_count": candidate_count, + "sequential": _summarize(sequential_timings), + "bounded_fanout": _summarize(fanout_timings), + "max_active_worker": max_active, + "search_calls": search_calls, + "simulated_query_embedding_provider_calls": iterations, + } + ) + return { + "delay_ms": delay_ms, + "iterations": iterations, + "query_embedding_precomputed": True, + "results": results, + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--candidates", default="1,2,4,10,20") + parser.add_argument("--delay-ms", type=int, default=20) + parser.add_argument("--iterations", type=int, default=20) + args = parser.parse_args() + candidates = tuple(int(value) for value in args.candidates.split(",")) + if ( + not candidates + or any(value < 1 or value > 20 for value in candidates) + or args.delay_ms < 1 + or args.iterations < 1 + ): + raise SystemExit("benchmark arguments are outside the safe bounded range") + print( + json.dumps( + run( + candidates=candidates, + delay_ms=args.delay_ms, + iterations=args.iterations, + ), + ensure_ascii=True, + indent=2, + ) + ) + + +if __name__ == "__main__": + main()