diff --git a/src/backend/base/langflow/alembic/versions/mb00a1b2c3d4_add_memory_base_schema.py b/src/backend/base/langflow/alembic/versions/mb00a1b2c3d4_add_memory_base_schema.py new file mode 100644 index 0000000000..e8817fb808 --- /dev/null +++ b/src/backend/base/langflow/alembic/versions/mb00a1b2c3d4_add_memory_base_schema.py @@ -0,0 +1,234 @@ +"""add_memory_base_schema + +Consolidates all Memory Base schema changes into a single migration: + - job.dedupe_key (nullable String) + ix_job_dedupe_key + - message.run_id (nullable UUID) + ix_message_run_id + - message.is_output (bool, default false) + - memory_base table + ix_memory_base_flow_id + ix_memory_base_user_id + - memory_base_session table + three indexes + - message_ingestion_record table + three indexes + - memory_base_workflow_run table + two indexes + +Phase: EXPAND + +Revision ID: mb00a1b2c3d4 +Revises: d306e5c17c41 +Create Date: 2026-04-14 00:00:00.000000 +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op +from langflow.utils import migration + +# revision identifiers, used by Alembic. +revision: str = "mb00a1b2c3d4" # pragma: allowlist secret +down_revision: str | None = "d306e5c17c41" # pragma: allowlist secret +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + conn = op.get_bind() + + # ------------------------------------------------------------------ # + # job.dedupe_key # + # ------------------------------------------------------------------ # + inspector = sa.inspect(conn) + existing_job_indexes = {idx["name"] for idx in inspector.get_indexes("job")} + with op.batch_alter_table("job", schema=None) as batch_op: + if not migration.column_exists("job", "dedupe_key", conn): + batch_op.add_column(sa.Column("dedupe_key", sa.String(), nullable=True)) + if "ix_job_dedupe_key" not in existing_job_indexes: + batch_op.create_index(batch_op.f("ix_job_dedupe_key"), ["dedupe_key"], unique=False) + + # ------------------------------------------------------------------ # + # message.run_id + message.is_output # + # ------------------------------------------------------------------ # + with op.batch_alter_table("message", schema=None) as batch_op: + if not migration.column_exists("message", "run_id", conn): + batch_op.add_column(sa.Column("run_id", sa.Uuid(), nullable=True)) + if not migration.column_exists("message", "is_output", conn): + batch_op.add_column(sa.Column("is_output", sa.Boolean(), nullable=False, server_default=sa.text("false"))) + + existing_message_indexes = {idx["name"] for idx in sa.inspect(conn).get_indexes("message")} + if "ix_message_run_id" not in existing_message_indexes: + op.create_index("ix_message_run_id", "message", ["run_id"]) + + # ------------------------------------------------------------------ # + # memory_base # + # ------------------------------------------------------------------ # + if not migration.table_exists("memory_base", conn): + op.create_table( + "memory_base", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column("name", sa.String(), nullable=False), + sa.Column("flow_id", sa.Uuid(), nullable=False), + sa.Column("user_id", sa.Uuid(), nullable=False), + sa.Column("threshold", sa.Integer(), nullable=False, server_default=sa.text("50")), + sa.Column("auto_capture", sa.Boolean(), nullable=False, server_default=sa.text("true")), + sa.Column("embedding_model", sa.String(), nullable=False, server_default=sa.text("''")), + sa.Column("preprocessing", sa.Boolean(), nullable=False, server_default=sa.text("false")), + sa.Column("preproc_model", sa.String(), nullable=True), + sa.Column("preproc_instructions", sa.String(), nullable=True), + sa.Column("kb_name", sa.String(), nullable=False), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("ix_memory_base_flow_id", "memory_base", ["flow_id"]) + op.create_index("ix_memory_base_user_id", "memory_base", ["user_id"]) + + # ------------------------------------------------------------------ # + # memory_base_session # + # ------------------------------------------------------------------ # + if not migration.table_exists("memory_base_session", conn): + op.create_table( + "memory_base_session", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column( + "memory_base_id", + sa.Uuid(), + sa.ForeignKey("memory_base.id", ondelete="CASCADE"), + nullable=False, + ), + sa.Column("session_id", sa.String(), nullable=False), + sa.Column("cursor_id", sa.Uuid(), nullable=True), + sa.Column("total_processed", sa.Integer(), nullable=False, server_default=sa.text("0")), + sa.Column("last_sync_at", sa.DateTime(timezone=True), nullable=True), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("memory_base_id", "session_id", name="uq_memory_base_session"), + ) + op.create_index("ix_memory_base_session_memory_base_id", "memory_base_session", ["memory_base_id"]) + op.create_index("ix_memory_base_session_session_id", "memory_base_session", ["session_id"]) + op.create_index( + "ix_memory_base_session_lookup", + "memory_base_session", + ["memory_base_id", "session_id"], + ) + + # ------------------------------------------------------------------ # + # message_ingestion_record # + # ------------------------------------------------------------------ # + if not migration.table_exists("message_ingestion_record", conn): + op.create_table( + "message_ingestion_record", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column( + "message_id", + sa.Uuid(), + sa.ForeignKey("message.id", ondelete="CASCADE"), + nullable=False, + ), + sa.Column( + "memory_base_id", + sa.Uuid(), + sa.ForeignKey("memory_base.id", ondelete="CASCADE"), + nullable=False, + ), + sa.Column( + "job_id", + sa.Uuid(), + sa.ForeignKey("job.job_id", ondelete="SET NULL"), + nullable=True, + ), + sa.Column("session_id", sa.String(), nullable=False), + sa.Column("ingested_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint( + "message_id", + "session_id", + "memory_base_id", + name="uq_mir_message_session_mb", + ), + ) + op.create_index("ix_mir_message_id", "message_ingestion_record", ["message_id"]) + op.create_index("ix_mir_job_id", "message_ingestion_record", ["job_id"]) + op.create_index( + "ix_mir_memory_base_session", + "message_ingestion_record", + ["memory_base_id", "session_id"], + ) + + # ------------------------------------------------------------------ # + # memory_base_workflow_run # + # ------------------------------------------------------------------ # + if not migration.table_exists("memory_base_workflow_run", conn): + op.create_table( + "memory_base_workflow_run", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column( + "memory_base_id", + sa.Uuid(), + sa.ForeignKey("memory_base.id", ondelete="CASCADE"), + nullable=False, + ), + sa.Column("session_id", sa.String(), nullable=False), + sa.Column( + "workflow_job_id", + sa.Uuid(), + sa.ForeignKey("job.job_id", ondelete="SET NULL"), + nullable=True, + ), + sa.Column( + "ingestion_job_id", + sa.Uuid(), + sa.ForeignKey("job.job_id", ondelete="SET NULL"), + nullable=True, + ), + sa.Column("recorded_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint( + "memory_base_id", + "session_id", + "workflow_job_id", + name="uq_mbwr_mb_session_wf_job", + ), + ) + op.create_index("ix_mbwr_mb_session", "memory_base_workflow_run", ["memory_base_id", "session_id"]) + op.create_index("ix_mbwr_ingestion_job_id", "memory_base_workflow_run", ["ingestion_job_id"]) + + +def downgrade() -> None: + conn = op.get_bind() + + # Children first (FK dependencies) ----------------------------------- # + if migration.table_exists("memory_base_workflow_run", conn): + op.drop_index("ix_mbwr_ingestion_job_id", table_name="memory_base_workflow_run") + op.drop_index("ix_mbwr_mb_session", table_name="memory_base_workflow_run") + op.drop_table("memory_base_workflow_run") + + if migration.table_exists("message_ingestion_record", conn): + op.drop_index("ix_mir_memory_base_session", table_name="message_ingestion_record") + op.drop_index("ix_mir_job_id", table_name="message_ingestion_record") + op.drop_index("ix_mir_message_id", table_name="message_ingestion_record") + op.drop_table("message_ingestion_record") + + if migration.table_exists("memory_base_session", conn): + op.drop_index("ix_memory_base_session_lookup", table_name="memory_base_session") + op.drop_index("ix_memory_base_session_session_id", table_name="memory_base_session") + op.drop_index("ix_memory_base_session_memory_base_id", table_name="memory_base_session") + op.drop_table("memory_base_session") + + if migration.table_exists("memory_base", conn): + op.drop_index("ix_memory_base_user_id", table_name="memory_base") + op.drop_index("ix_memory_base_flow_id", table_name="memory_base") + op.drop_table("memory_base") + + # Message column/index ----------------------------------------------- # + existing_message_indexes = {idx["name"] for idx in sa.inspect(conn).get_indexes("message")} + if "ix_message_run_id" in existing_message_indexes: + op.drop_index("ix_message_run_id", table_name="message") + with op.batch_alter_table("message", schema=None) as batch_op: + if migration.column_exists("message", "is_output", conn): + batch_op.drop_column("is_output") + if migration.column_exists("message", "run_id", conn): + batch_op.drop_column("run_id") + + # Job column/index --------------------------------------------------- # + with op.batch_alter_table("job", schema=None) as batch_op: + existing_job_indexes = {idx["name"] for idx in sa.inspect(conn).get_indexes("job")} + if "ix_job_dedupe_key" in existing_job_indexes: + batch_op.drop_index(batch_op.f("ix_job_dedupe_key")) + if migration.column_exists("job", "dedupe_key", conn): + batch_op.drop_column("dedupe_key") diff --git a/src/backend/base/langflow/api/build.py b/src/backend/base/langflow/api/build.py index 7ed4feffd6..a35c69c6d4 100644 --- a/src/backend/base/langflow/api/build.py +++ b/src/backend/base/langflow/api/build.py @@ -29,7 +29,15 @@ from langflow.exceptions.component import ComponentBuildError from langflow.schema.message import ErrorMessage from langflow.schema.schema import OutputValue from langflow.services.database.models.flow.model import Flow -from langflow.services.deps import get_chat_service, get_telemetry_service, session_scope +from langflow.services.database.models.jobs.model import JobType +from langflow.services.deps import ( + get_chat_service, + get_job_service, + get_memory_base_service, + get_task_service, + get_telemetry_service, + session_scope, +) from langflow.services.job_queue.service import JobQueueNotFoundError, JobQueueService from langflow.services.telemetry.schema import ComponentInputsPayload, ComponentPayload, PlaygroundPayload @@ -527,35 +535,81 @@ async def generate_flow_events( event_manager.on_error(data=error_message.data) raise + # Create a WORKFLOW job record so memory-base on_flow_output can track this run. + # Best-effort: failures here must never break the build path. + _build_job_svc = None + _build_run_id: uuid.UUID | None = None + try: + _build_run_id = uuid.UUID(graph.run_id) if graph.run_id else None + if _build_run_id is not None: + _build_job_svc = get_job_service() + await _build_job_svc.create_job( + job_id=_build_run_id, + flow_id=flow_id, + user_id=current_user.id, + job_type=JobType.WORKFLOW, + ) + except Exception: # noqa: BLE001 + await logger.awarning( + "Failed to create workflow job for /build — memory base tracking disabled for flow %s", + flow_id, + exc_info=True, + ) + _build_job_svc = None + event_manager.on_vertices_sorted(data={"ids": ids, "to_run": vertices_to_run}) vertex_timedeltas: list[float] = [] event_manager.on_build_start(data={}) - tasks = [] - for vertex_id in ids: - task = asyncio.create_task(build_vertices(vertex_id, graph, event_manager, vertex_timedeltas)) - tasks.append(task) - try: - await asyncio.gather(*tasks) - except asyncio.CancelledError: - background_tasks.add_task(graph.end_all_traces_in_context()) - raise - except Exception as e: - await logger.aerror(f"Error building vertices: {e}") - custom_component = graph.get_vertex(vertex_id).custom_component - trace_name = getattr(custom_component, "trace_name", None) - error_message = ErrorMessage( - flow_id=flow_id, - exception=e, - session_id=graph.session_id, - trace_name=trace_name, - ) - event_manager.on_error(data=error_message.data) - raise + + async def _run_vertex_build() -> None: + tasks = [] + for vertex_id in ids: + task = asyncio.create_task(build_vertices(vertex_id, graph, event_manager, vertex_timedeltas)) + tasks.append(task) + try: + await asyncio.gather(*tasks) + except asyncio.CancelledError: + background_tasks.add_task(graph.end_all_traces_in_context()) + raise + except Exception as e: + await logger.aerror(f"Error building vertices: {e}") + custom_component = graph.get_vertex(vertex_id).custom_component + trace_name = getattr(custom_component, "trace_name", None) + error_message = ErrorMessage( + flow_id=flow_id, + exception=e, + session_id=graph.session_id, + trace_name=trace_name, + ) + event_manager.on_error(data=error_message.data) + raise + + if _build_job_svc and _build_run_id: + await _build_job_svc.execute_with_status(_build_run_id, _run_vertex_build) + else: + await _run_vertex_build() build_duration = sum(vertex_timedeltas) event_manager.on_end(data={"build_duration": build_duration}) await graph.end_all_traces() + + # Fire memory-base auto-capture hook — non-blocking background effect. + # Must use fire_and_forget_task (not background_tasks.add_task) because + # generate_flow_events runs as an asyncio task; by the time the flow + # finishes, FastAPI has already drained the background_tasks queue and any + # tasks added after that point are silently dropped. + try: + _run_id_uuid = uuid.UUID(graph.run_id) if graph.run_id else None # type-cast only; same run_id set on graph + await get_task_service().fire_and_forget_task( + get_memory_base_service().on_flow_output, + flow_id=flow_id, + session_id=graph.session_id or str(flow_id), + job_id=_run_id_uuid, + ) + except (RuntimeError, ValueError, OSError): + await logger.awarning("Memory base hook scheduling failed for flow %s", flow_id, exc_info=True) + await event_manager.queue.put((None, None, time.time())) diff --git a/src/backend/base/langflow/api/router.py b/src/backend/base/langflow/api/router.py index c4a38d44fa..bfb8b45c97 100644 --- a/src/backend/base/langflow/api/router.py +++ b/src/backend/base/langflow/api/router.py @@ -15,6 +15,7 @@ from langflow.api.v1 import ( login_router, mcp_projects_router, mcp_router, + memories_router, model_options_router, models_router, monitor_router, @@ -68,6 +69,7 @@ router_v1.include_router(folders_router) router_v1.include_router(projects_router) router_v1.include_router(starter_projects_router) router_v1.include_router(knowledge_bases_router) +router_v1.include_router(memories_router) router_v1.include_router(mcp_router) router_v1.include_router(voice_mode_router) router_v1.include_router(mcp_projects_router) diff --git a/src/backend/base/langflow/api/utils/kb_helpers.py b/src/backend/base/langflow/api/utils/kb_helpers.py index b195777f19..63a1c0190b 100644 --- a/src/backend/base/langflow/api/utils/kb_helpers.py +++ b/src/backend/base/langflow/api/utils/kb_helpers.py @@ -445,7 +445,7 @@ class KBIngestionHelper: splitter_kwargs["separators"] = [resolved_separator] text_splitter = RecursiveCharacterTextSplitter(**splitter_kwargs) - embeddings = await KBIngestionHelper._build_embeddings(embedding_provider, embedding_model, current_user) + embeddings = await KBIngestionHelper.build_embeddings(embedding_provider, embedding_model, current_user) client = KBStorageHelper.get_fresh_chroma_client(kb_path) chroma = Chroma( @@ -462,40 +462,30 @@ class KBIngestionHelper: continue chunks = text_splitter.split_text(content) - for i in range(0, len(chunks), INGESTION_BATCH_SIZE): - if await KBIngestionHelper._is_job_cancelled(job_service, task_job_id): - raise IngestionCancelledError + docs = [ + Document( + page_content=c, + metadata={ + "source": source_name or file_name, + "file_name": file_name, + "chunk_index": i, + "total_chunks": len(chunks), + "ingested_at": datetime.now(timezone.utc).isoformat(), + "job_id": job_id_str, + }, + ) + for i, c in enumerate(chunks) + ] - batch = chunks[i : i + INGESTION_BATCH_SIZE] - docs = [ - Document( - page_content=c, - metadata={ - "source": source_name or file_name, - "file_name": file_name, - "chunk_index": i + j, - "total_chunks": len(chunks), - "ingested_at": datetime.now(timezone.utc).isoformat(), - "job_id": job_id_str, - }, - ) - for j, c in enumerate(batch) - ] - - for attempt in range(MAX_RETRY_ATTEMPTS): - if await KBIngestionHelper._is_job_cancelled(job_service, task_job_id): - raise IngestionCancelledError - try: - await chroma.aadd_documents(docs) - break - except Exception as e: - if attempt == MAX_RETRY_ATTEMPTS - 1: - raise - wait = (attempt + 1) * EXPONENTIAL_BACKOFF_MULTIPLIER - await logger.awarning("Write failed, retrying in %ds: %s", wait, e) - await asyncio.sleep(wait) - - await asyncio.sleep(0.01) + written = await KBIngestionHelper.write_documents_to_chroma( + documents=docs, + chroma=chroma, + task_job_id=task_job_id, + job_service=job_service, + ) + if written < len(docs): + # Job was cancelled mid-file + raise IngestionCancelledError total_chunks_created += len(chunks) processed_files.append(file_name) @@ -555,13 +545,69 @@ class KBIngestionHelper: KBStorageHelper.release_chroma_resources(kb_path) @staticmethod - async def _is_job_cancelled(job_service: JobService, job_id: uuid.UUID) -> bool: + async def write_documents_to_chroma( + *, + documents: list[Document], + chroma: Chroma, + task_job_id: uuid.UUID, + job_service: JobService, + ) -> int: + """Write pre-built Documents into an open Chroma collection. + + This is the shared primitive used by both file-based KB ingestion + (``perform_ingestion``) and message-based Memory Base ingestion. + + Documents must already be chunked and have their metadata populated + by the caller — this method only handles the batched write, cancellation + checking, and retry logic. + + Args: + documents: LangChain Document objects ready for embedding. + chroma: An already-constructed ``Chroma`` instance pointing at the + target collection. + task_job_id: Job ID used to poll for cancellation. + job_service: Service for checking job status. + + Returns: + Number of documents successfully written. If the job is cancelled + mid-batch this will be less than ``len(documents)``. + + Raises: + Exception: Re-raises any non-cancellation write failure after the + retry budget is exhausted. + """ + written = 0 + for i in range(0, len(documents), INGESTION_BATCH_SIZE): + if await KBIngestionHelper.is_job_cancelled(job_service, task_job_id): + return written + + batch = documents[i : i + INGESTION_BATCH_SIZE] + for attempt in range(MAX_RETRY_ATTEMPTS): + if await KBIngestionHelper.is_job_cancelled(job_service, task_job_id): + return written + try: + await chroma.aadd_documents(batch) + break + except Exception as e: + if attempt == MAX_RETRY_ATTEMPTS - 1: + raise + wait = (attempt + 1) * EXPONENTIAL_BACKOFF_MULTIPLIER + await logger.awarning("Write failed, retrying in %ds: %s", wait, e) + await asyncio.sleep(wait) + + written += len(batch) + await asyncio.sleep(0.01) + + return written + + @staticmethod + async def is_job_cancelled(job_service: JobService, job_id: uuid.UUID) -> bool: """Internal helper to check if a job has been cancelled.""" job = await job_service.get_job_by_job_id(job_id) return job is not None and job.status == JobStatus.CANCELLED @staticmethod - async def _build_embeddings(provider: str, model: str, current_user): + async def build_embeddings(provider: str, model: str, current_user): """Internal helper to build embeddings object.""" options = get_embedding_model_options(user_id=current_user.id) selected_option = next((o for o in options if o["provider"] == provider and o["name"] == model), None) diff --git a/src/backend/base/langflow/api/v1/__init__.py b/src/backend/base/langflow/api/v1/__init__.py index 177fc34df0..11dcb2816b 100644 --- a/src/backend/base/langflow/api/v1/__init__.py +++ b/src/backend/base/langflow/api/v1/__init__.py @@ -10,6 +10,7 @@ from langflow.api.v1.knowledge_bases import router as knowledge_bases_router from langflow.api.v1.login import router as login_router from langflow.api.v1.mcp import router as mcp_router from langflow.api.v1.mcp_projects import router as mcp_projects_router +from langflow.api.v1.memories import router as memories_router from langflow.api.v1.model_options import router as model_options_router from langflow.api.v1.models import router as models_router from langflow.api.v1.monitor import router as monitor_router @@ -36,6 +37,7 @@ __all__ = [ "login_router", "mcp_projects_router", "mcp_router", + "memories_router", "model_options_router", "models_router", "monitor_router", diff --git a/src/backend/base/langflow/api/v1/endpoints.py b/src/backend/base/langflow/api/v1/endpoints.py index 10cb82d683..5a38e73caa 100644 --- a/src/backend/base/langflow/api/v1/endpoints.py +++ b/src/backend/base/langflow/api/v1/endpoints.py @@ -61,8 +61,17 @@ from langflow.services.auth.utils import ( from langflow.services.cache.utils import save_uploaded_file from langflow.services.database.models.flow.model import Flow, FlowRead from langflow.services.database.models.flow.utils import get_all_webhook_components_in_flow +from langflow.services.database.models.jobs.model import JobType from langflow.services.database.models.user.model import User, UserRead -from langflow.services.deps import get_auth_service, get_session_service, get_settings_service, get_telemetry_service +from langflow.services.deps import ( + get_auth_service, + get_job_service, + get_memory_base_service, + get_session_service, + get_settings_service, + get_task_service, + get_telemetry_service, +) from langflow.services.event_manager import create_webhook_event_manager, webhook_event_manager from langflow.services.telemetry.schema import RunPayload from langflow.utils.compression import compress_response @@ -172,8 +181,8 @@ async def simple_run_flow( graph = Graph.from_payload( graph_data, flow_id=flow_id_str, user_id=str(user_id), flow_name=flow.name, context=context ) - if run_id is None: - run_id = str(uuid4()) + run_id_uuid = uuid4() if run_id is None else UUID(run_id) + run_id = str(run_id_uuid) graph.set_run_id(run_id) inputs = None if input_request.input_value is not None: @@ -196,15 +205,61 @@ async def simple_run_flow( and (input_request.output_type == "any" or input_request.output_type in vertex.id.lower()) # type: ignore[operator] ) ] - task_result, session_id = await run_graph_internal( - graph=graph, - flow_id=flow_id_str, - session_id=input_request.session_id, - inputs=inputs, - outputs=outputs, - stream=stream, - event_manager=event_manager, - ) + + # Create a WORKFLOW job record so memory-base on_flow_output can track this run. + # Best-effort: if job creation fails (e.g. no authenticated user), fall back to + # direct execution without memory-base tracking. + _job_created = False + if user_id is not None: + try: + _job_svc = get_job_service() + await _job_svc.create_job( + job_id=run_id_uuid, + flow_id=flow.id, + user_id=user_id, + job_type=JobType.WORKFLOW, + ) + task_result, session_id = await _job_svc.execute_with_status( + run_id_uuid, + run_graph_internal, + graph=graph, + flow_id=flow_id_str, + session_id=input_request.session_id, + inputs=inputs, + outputs=outputs, + stream=stream, + event_manager=event_manager, + ) + _job_created = True + except Exception: # noqa: BLE001 + await logger.awarning( + "Failed to create workflow job for /run — memory base tracking disabled for flow %s", + flow.id, + exc_info=True, + ) + + if not _job_created: + task_result, session_id = await run_graph_internal( + graph=graph, + flow_id=flow_id_str, + session_id=input_request.session_id, + inputs=inputs, + outputs=outputs, + stream=stream, + event_manager=event_manager, + ) + + # Fire memory-base auto-capture hook — non-blocking background effect. + try: + _run_id_uuid = UUID(graph.run_id) if graph.run_id else None # type-cast only + await get_task_service().fire_and_forget_task( + get_memory_base_service().on_flow_output, + flow_id=flow.id, + session_id=session_id, + job_id=_run_id_uuid, + ) + except (RuntimeError, ValueError, OSError): + await logger.awarning("Memory base hook scheduling failed for flow %s", flow.id, exc_info=True) return RunResponse(outputs=task_result, session_id=session_id) @@ -942,6 +997,18 @@ async def experimental_run_flow( except Exception as exc: raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=str(exc)) from exc + # Fire memory-base auto-capture hook — non-blocking background effect. + try: + _run_id_uuid = UUID(graph.run_id) if graph.run_id else None # type-cast only + await get_task_service().fire_and_forget_task( + get_memory_base_service().on_flow_output, + flow_id=flow.id, + session_id=session_id, + job_id=_run_id_uuid, + ) + except (RuntimeError, ValueError, OSError): + await logger.awarning("Memory base hook scheduling failed for flow %s", flow.id, exc_info=True) + return RunResponse(outputs=task_result, session_id=session_id) diff --git a/src/backend/base/langflow/api/v1/knowledge_bases.py b/src/backend/base/langflow/api/v1/knowledge_bases.py index 62ddbf41a1..6c9490edcc 100644 --- a/src/backend/base/langflow/api/v1/knowledge_bases.py +++ b/src/backend/base/langflow/api/v1/knowledge_bases.py @@ -74,6 +74,30 @@ def _resolve_kb_path(kb_name: str, current_user: CurrentActiveUser) -> Path: return kb_path +def _is_memory_base_associated(metadata: dict[str, Any]) -> bool: + """Return True if the KB metadata indicates an association with a Memory Base.""" + source_types = metadata.get("source_types") + return isinstance(source_types, list) and "memory" in source_types + + +def _check_memory_base_association(kb_name: str, current_user: CurrentActiveUser) -> None: + """Raise 403 if the KB is associated with a Memory Base. + + Designed as a FastAPI dependency for per-KB routes — FastAPI injects + ``kb_name`` from the path parameter and ``current_user`` via its own + dependency. The list endpoint filters memory KBs inline using + ``_is_memory_base_associated`` directly. + """ + kb_path = _resolve_kb_path(kb_name, current_user) + + metadata = KBAnalysisHelper.get_metadata(kb_path, fast=True) + if _is_memory_base_associated(metadata): + raise HTTPException( + status_code=403, + detail=f"Access denied: knowledge base '{kb_name}' is managed by a Memory Base.", + ) + + @router.post("", status_code=HTTPStatus.CREATED) @router.post("/", status_code=HTTPStatus.CREATED) async def create_knowledge_base( @@ -275,7 +299,7 @@ async def preview_chunks( return {"files": file_previews} -@router.post("/{kb_name}/ingest", status_code=HTTPStatus.OK) +@router.post("/{kb_name}/ingest", status_code=HTTPStatus.OK, dependencies=[Depends(_check_memory_base_association)]) async def ingest_files_to_knowledge_base( kb_name: str, current_user: CurrentActiveUser, @@ -427,6 +451,8 @@ async def list_knowledge_bases( try: # Use deep update (fast=False) to ensure legacy KBs are migrated on first view metadata = KBAnalysisHelper.get_metadata(kb_dir, fast=False) + if _is_memory_base_associated(metadata): + continue # Skip KBs that are associated with a Memory Base # Extract KB ID from metadata (stored as string, convert to UUID) kb_id_str = metadata.get("id") @@ -503,7 +529,7 @@ async def list_knowledge_bases( return knowledge_bases -@router.get("/{kb_name}", status_code=HTTPStatus.OK) +@router.get("/{kb_name}", status_code=HTTPStatus.OK, dependencies=[Depends(_check_memory_base_association)]) async def get_knowledge_base(kb_name: str, current_user: CurrentActiveUser) -> KnowledgeBaseInfo: """Get detailed information about a specific knowledge base.""" try: @@ -543,7 +569,7 @@ async def get_knowledge_base(kb_name: str, current_user: CurrentActiveUser) -> K raise HTTPException(status_code=500, detail="Error getting knowledge base.") from e -@router.get("/{kb_name}/chunks", status_code=HTTPStatus.OK) +@router.get("/{kb_name}/chunks", status_code=HTTPStatus.OK, dependencies=[Depends(_check_memory_base_association)]) async def get_knowledge_base_chunks( kb_name: str, current_user: CurrentActiveUser, @@ -636,7 +662,7 @@ async def get_knowledge_base_chunks( KBStorageHelper.release_chroma_resources(kb_path) -@router.delete("/{kb_name}", status_code=HTTPStatus.OK) +@router.delete("/{kb_name}", status_code=HTTPStatus.OK, dependencies=[Depends(_check_memory_base_association)]) async def delete_knowledge_base(kb_name: str, current_user: CurrentActiveUser) -> dict[str, str]: """Delete a specific knowledge base.""" try: @@ -703,7 +729,7 @@ async def delete_knowledge_bases_bulk(request: BulkDeleteRequest, current_user: return result -@router.post("/{kb_name}/cancel", status_code=HTTPStatus.OK) +@router.post("/{kb_name}/cancel", status_code=HTTPStatus.OK, dependencies=[Depends(_check_memory_base_association)]) async def cancel_ingestion( kb_name: str, current_user: CurrentActiveUser, diff --git a/src/backend/base/langflow/api/v1/memories.py b/src/backend/base/langflow/api/v1/memories.py new file mode 100644 index 0000000000..83d3b79649 --- /dev/null +++ b/src/backend/base/langflow/api/v1/memories.py @@ -0,0 +1,343 @@ +"""REST API for Memory Base management. + +Endpoints: + POST /memories - Create + GET /memories - List (current user, paginated) + GET /memories/{id} - Get one + GET /memories/{id}/sessions - List sessions (tracked + untracked from MessageTable) + PATCH /memories/{id} - Update (name / threshold / auto_capture / preprocessing) + DELETE /memories/{id} - Delete (cancels active tasks + removes KB from disk) + POST /memories/{id}/flush - Manual flush / trigger ingestion + POST /memories/{id}/regenerate - Regenerate from mismatch + +Edge cases enforced: + 409 Conflict - name already in use for this user (on create). + 409 Conflict - active ingestion task already running for same (mb, session). + 404 Not Found - memory base does not belong to the current user. + 422 Unprocessable - preprocessing=true but preproc_model missing. +""" + +from __future__ import annotations + +import uuid +from datetime import datetime +from http import HTTPStatus +from typing import Annotated + +from fastapi import APIRouter, Body, Depends, HTTPException +from fastapi_pagination import Page, Params +from fastapi_pagination.ext.sqlmodel import apaginate +from pydantic import BaseModel +from sqlmodel import col, select + +from langflow.api.utils import CurrentActiveUser +from langflow.services.database.models.memory_base.model import ( + MemoryBase, + MemoryBaseCreate, + MemoryBaseRead, + MemoryBaseSessionRead, + MemoryBaseUpdate, +) +from langflow.services.database.models.message.model import MessageTable +from langflow.services.deps import get_memory_base_service, session_scope +from langflow.services.jobs import DuplicateJobError + +router = APIRouter(tags=["Memories"], prefix="/memories", include_in_schema=False) + + +# ------------------------------------------------------------------ # +# Request / Response schemas # +# ------------------------------------------------------------------ # + + +class MessageReadResponse(BaseModel): + """Slim message projection for Memory Base session message listings. + + Only messages that have been ingested into the requested Memory Base are returned. + ``job_id`` and ``ingested_at`` are sourced from MessageIngestionRecord. + """ + + model_config = {"from_attributes": True} + + id: uuid.UUID + timestamp: datetime | None = None + sender: str + sender_name: str + session_id: str + text: str + content_blocks: list = [] + job_id: uuid.UUID | None = None + ingested_at: datetime | None = None + + +class FlushRequest(BaseModel): + session_id: str + + +class MismatchResponse(BaseModel): + mismatch_detected: bool + + +class RegenerateResponse(BaseModel): + job_ids: list[str] + + +# ------------------------------------------------------------------ # +# CRUD # +# ------------------------------------------------------------------ # + + +@router.post("", status_code=HTTPStatus.CREATED) +@router.post("/", status_code=HTTPStatus.CREATED) +async def create_memory_base( + current_user: CurrentActiveUser, + payload: Annotated[MemoryBaseCreate, Body(embed=False)] = ..., +) -> MemoryBaseRead: + """Create a new Memory Base. + + - kb_name is auto-generated as `{sanitized_name}_{8hex}`. + - KB directory and embedding_metadata.json are created on disk immediately. + - Returns 409 if a Memory Base with the same name already exists for this user. + - Returns 422 if preprocessing=true but preproc_model is missing. + """ + try: + mb = await get_memory_base_service().create(payload, user_id=current_user.id) + except PermissionError as exc: + # Flow not found or belongs to another user — return 404 to avoid info-leak + raise HTTPException(status_code=404, detail=str(exc)) from exc + except ValueError as exc: + raise HTTPException(status_code=409, detail=str(exc)) from exc + return MemoryBaseRead.model_validate(mb) + + +@router.get("", status_code=HTTPStatus.OK) +@router.get("/", status_code=HTTPStatus.OK) +async def list_memory_bases( + current_user: CurrentActiveUser, + params: Annotated[Params, Depends()], + flow_id: uuid.UUID | None = None, +) -> Page[MemoryBaseRead]: + """List all Memory Bases owned by the current user (paginated) for a flow_id. + + Query params (from fastapi-pagination): + page - 1-based page number (default 1) + size - page size (default 50) + """ + async with session_scope() as db: + stmt = get_memory_base_service().list_for_user_stmt(user_id=current_user.id, flow_id=flow_id) + return await apaginate( + db, stmt, params=params, transformer=lambda items: [MemoryBaseRead.model_validate(m) for m in items] + ) + + +@router.get("/{memory_base_id}", status_code=HTTPStatus.OK) +async def get_memory_base( + memory_base_id: uuid.UUID, + current_user: CurrentActiveUser, +) -> MemoryBaseRead: + """Get details for a specific Memory Base.""" + mb = await get_memory_base_service().get(memory_base_id, user_id=current_user.id) + if mb is None: + raise HTTPException(status_code=404, detail="Memory base not found") + return MemoryBaseRead.model_validate(mb) + + +@router.get("/{memory_base_id}/sessions", status_code=HTTPStatus.OK) +async def list_sessions( + memory_base_id: uuid.UUID, + current_user: CurrentActiveUser, + params: Annotated[Params, Depends()], +) -> Page[MemoryBaseSessionRead]: + """List persisted sessions for this Memory Base (paginated). + + Only sessions that have been synced at least once (i.e. have a + MemoryBaseSession row) are returned. Results are ordered by + last_sync_at descending. + + Each item includes ``pending_count``: the number of completed flow runs + remaining before the next auto-capture ingestion is triggered. + """ + async with session_scope() as db: + try: + mb = await get_memory_base_service()._get_mb_or_raise(db, memory_base_id, current_user.id) # noqa: SLF001 + except ValueError as exc: + raise HTTPException(status_code=404, detail=str(exc)) from exc + + stmt = get_memory_base_service().sessions_stmt(memory_base_id) + raw_page = await apaginate(db, stmt, params=params) + + items: list[MemoryBaseSessionRead] = [] + for s in raw_page.items: + pending_count = await get_memory_base_service()._count_pending(db, mb, s) # noqa: SLF001 + read = MemoryBaseSessionRead.model_validate(s) + read.pending_count = pending_count + items.append(read) + + return raw_page.model_copy(update={"items": items}) + + +@router.get("/{memory_base_id}/sessions/{session_id}/messages", status_code=HTTPStatus.OK) +async def list_session_messages( + memory_base_id: uuid.UUID, + session_id: str, + current_user: CurrentActiveUser, + params: Annotated[Params, Depends()], +) -> Page[MessageReadResponse]: + """List messages ingested into this Memory Base session (paginated). + + Only messages that have been successfully ingested into the requested Memory Base + are returned. Messages are ordered by timestamp ascending. + Each item includes ``job_id`` and ``ingested_at`` from the MessageIngestionRecord. + + Returns 404 if the Memory Base does not belong to the current user. + """ + from sqlalchemy import and_ + + from langflow.services.database.models.memory_base.model import MessageIngestionRecord + + async with session_scope() as db: + mb_stmt = select(MemoryBase).where(MemoryBase.id == memory_base_id).where(MemoryBase.user_id == current_user.id) + result = await db.exec(mb_stmt) + if result.first() is None: + raise HTTPException(status_code=404, detail="Memory base not found") + + # INNER JOIN — only messages that were actually ingested into this MB/session pair. + # No extra WHERE filters needed: + # - mir.session_id == session_id in the JOIN guarantees msg.session_id == session_id + # (session_id is denormalized from the message at ingestion time — immutable). + # - flow_id is implicitly correct: ingestion only ever touches messages from mb.flow_id, + # and MB ownership is already verified above. + msg_stmt = ( + select(MessageTable, MessageIngestionRecord) + .join( + MessageIngestionRecord, + and_( + MessageIngestionRecord.message_id == MessageTable.id, + MessageIngestionRecord.memory_base_id == memory_base_id, + MessageIngestionRecord.session_id == session_id, + ), + ) + .order_by(col(MessageTable.timestamp).asc()) + ) + return await apaginate( + db, + msg_stmt, + params=params, + transformer=lambda rows: [ + MessageReadResponse( + id=msg.id, + timestamp=msg.timestamp, + sender=msg.sender, + sender_name=msg.sender_name, + session_id=msg.session_id, + text=msg.text, + content_blocks=msg.content_blocks or [], + job_id=mir.job_id, + ingested_at=mir.ingested_at, + ) + for msg, mir in rows + ], + ) + + +@router.patch("/{memory_base_id}", status_code=HTTPStatus.OK) +async def update_memory_base( + memory_base_id: uuid.UUID, + current_user: CurrentActiveUser, + patch: Annotated[MemoryBaseUpdate, Body(embed=False)] = ..., +) -> MemoryBaseRead: + """Update mutable parameters (threshold, auto_capture, preprocessing, etc.). + + Threshold changes only take effect at the next auto-capture trigger. + Any already-running ingestion task continues with its original arguments. + """ + mb = await get_memory_base_service().update(memory_base_id, user_id=current_user.id, patch=patch) + if mb is None: + raise HTTPException(status_code=404, detail="Memory base not found") + return MemoryBaseRead.model_validate(mb) + + +@router.delete("/{memory_base_id}", status_code=HTTPStatus.NO_CONTENT) +async def delete_memory_base( + memory_base_id: uuid.UUID, + current_user: CurrentActiveUser, +) -> None: + """Delete a Memory Base. + + Active ingestion tasks are forcefully cancelled before the DB record is + removed. The associated KB directory is deleted from disk afterwards + (best-effort — a disk failure will not affect the 204 response). + """ + deleted = await get_memory_base_service().delete(memory_base_id, user_id=current_user.id) + if not deleted: + raise HTTPException(status_code=404, detail="Memory base not found") + + +# ------------------------------------------------------------------ # +# Ingestion trigger # +# ------------------------------------------------------------------ # + + +@router.post("/{memory_base_id}/flush", status_code=HTTPStatus.ACCEPTED) +async def flush_memory_base( + memory_base_id: uuid.UUID, + current_user: CurrentActiveUser, + body: Annotated[FlushRequest, Body(embed=False)] = ..., +) -> dict: + """Manually trigger an ingestion / sync job regardless of the threshold. + + Returns 409 Conflict if an ingestion task is already in progress for the + given (memory_base_id, session_id) pair to prevent concurrent indexing. + """ + try: + job_id = await get_memory_base_service().trigger_ingestion( + memory_base_id=memory_base_id, + user_id=current_user.id, + session_id=body.session_id, + ) + except ValueError as exc: + raise HTTPException(status_code=404, detail=str(exc)) from exc + except DuplicateJobError as exc: + raise HTTPException(status_code=409, detail=str(exc)) from exc + except RuntimeError as exc: + raise HTTPException(status_code=409, detail=str(exc)) from exc + + return {"job_id": job_id} + + +# ------------------------------------------------------------------ # +# Mismatch detection & regeneration # +# ------------------------------------------------------------------ # + + +@router.get("/{memory_base_id}/mismatch", status_code=HTTPStatus.OK) +async def check_mismatch( + memory_base_id: uuid.UUID, + current_user: CurrentActiveUser, +) -> MismatchResponse: + """Detect if the vector store is empty while metadata records processed messages. + + The UI should surface a "Mismatch Detected" warning and offer a Regenerate button. + """ + try: + detected = await get_memory_base_service().check_mismatch(memory_base_id, user_id=current_user.id) + except ValueError as exc: + raise HTTPException(status_code=404, detail=str(exc)) from exc + return MismatchResponse(mismatch_detected=detected) + + +@router.post("/{memory_base_id}/regenerate", status_code=HTTPStatus.ACCEPTED) +async def regenerate_memory_base( + memory_base_id: uuid.UUID, + current_user: CurrentActiveUser, +) -> RegenerateResponse: + """Regenerate the Knowledge Base by resetting all session cursors and re-ingesting. + + Use this to recover from external Chroma directory deletions or vector DB corruption. + All MemoryBaseSession.cursor_id values are set to None before re-running ingestion. + """ + try: + job_ids = await get_memory_base_service().regenerate(memory_base_id, user_id=current_user.id) + except ValueError as exc: + raise HTTPException(status_code=404, detail=str(exc)) from exc + return RegenerateResponse(job_ids=job_ids) diff --git a/src/backend/base/langflow/api/v2/workflow.py b/src/backend/base/langflow/api/v2/workflow.py index 8bd7749ed1..3d513a2be6 100644 --- a/src/backend/base/langflow/api/v2/workflow.py +++ b/src/backend/base/langflow/api/v2/workflow.py @@ -29,6 +29,7 @@ from uuid import UUID, uuid4 from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, Request, status from fastapi.responses import StreamingResponse from lfx.graph.graph.base import Graph +from lfx.log.logger import logger from lfx.schema.workflow import ( WORKFLOW_EXECUTION_RESPONSES, WORKFLOW_STATUS_RESPONSES, @@ -65,7 +66,7 @@ from langflow.services.auth.utils import api_key_security from langflow.services.database.models.flow.model import FlowRead from langflow.services.database.models.jobs.model import JobType from langflow.services.database.models.user.model import UserRead -from langflow.services.deps import get_job_service, get_task_service +from langflow.services.deps import get_job_service, get_memory_base_service, get_task_service # Configuration constants EXECUTION_TIMEOUT = 300 # 5 minutes default timeout for sync execution @@ -391,6 +392,18 @@ async def execute_sync_workflow( stream=False, ) + # Fire memory-base auto-capture hook — non-blocking background effect. + try: + _run_id_uuid = UUID(graph.run_id) if graph.run_id else None # type-cast only; same run_id set on graph + await get_task_service().fire_and_forget_task( + get_memory_base_service().on_flow_output, + flow_id=flow.id, + session_id=execution_session_id, + job_id=_run_id_uuid, + ) + except (RuntimeError, ValueError, OSError): + await logger.awarning("Memory base hook scheduling failed for flow %s", flow.id, exc_info=True) + # Build RunResponse run_response = RunResponse(outputs=task_result, session_id=execution_session_id) # Convert to WorkflowExecutionResponse @@ -470,10 +483,39 @@ async def execute_workflow_background( user_id=api_key_user.id, ) + # Closure captures flow identity for the memory-base hook. + # run_id is the same as job_id — graph.set_run_id(job_id) was called above. + _hook_flow_id = flow.id + _hook_run_id = job_id + + async def _run_and_notify(**kwargs): + """Thin wrapper: execute graph then fire memory-base hook as a background effect. + + The hook is dispatched non-blocking after graph completion. Any failure in + the hook is swallowed so it never affects the job status of the graph run. + """ + result = await run_graph_internal(**kwargs) + _, _effective_session_id = result + try: + # Direct await — we are already inside a background task; awaiting here + # is non-blocking from the client's perspective and avoids the race + # condition that arises when dispatching a second fire_and_forget from + # within an already-running fire_and_forget task. + await get_memory_base_service().on_flow_output( + flow_id=_hook_flow_id, + session_id=_effective_session_id, + job_id=_hook_run_id, + ) + except Exception: # noqa: BLE001 + await logger.awarning( + "Memory base hook failed for flow %s, but workflow succeeded.", _hook_flow_id, exc_info=True + ) + return result + await task_service.fire_and_forget_task( job_service.execute_with_status, job_id=job_id, - run_coro_func=run_graph_internal, + run_coro_func=_run_and_notify, graph=graph, flow_id=flow_id_str, session_id=session_id, @@ -673,7 +715,7 @@ async def stop_workflow( task_service = get_task_service() try: - # 1. Fetch Job + # 1. Fetch Job and verify ownership job = await job_service.get_job_by_job_id(job_id, user_id=api_key_user.id) except Exception as exc: raise HTTPException( diff --git a/src/backend/base/langflow/services/database/models/__init__.py b/src/backend/base/langflow/services/database/models/__init__.py index 2ecc8d0153..b00e515af3 100644 --- a/src/backend/base/langflow/services/database/models/__init__.py +++ b/src/backend/base/langflow/services/database/models/__init__.py @@ -8,6 +8,7 @@ from .flow_version import FlowVersion from .flow_version_deployment_attachment import FlowVersionDeploymentAttachment from .folder import Folder from .jobs import Job +from .memory_base import MemoryBase, MemoryBaseSession, MemoryBaseWorkflowRun, MessageIngestionRecord from .message import MessageTable from .traces.model import SpanTable, TraceTable from .transactions import TransactionTable @@ -24,6 +25,10 @@ __all__ = [ "FlowVersionDeploymentAttachment", "Folder", "Job", + "MemoryBase", + "MemoryBaseSession", + "MemoryBaseWorkflowRun", + "MessageIngestionRecord", "MessageTable", "SSOConfig", "SSOUserProfile", diff --git a/src/backend/base/langflow/services/database/models/jobs/model.py b/src/backend/base/langflow/services/database/models/jobs/model.py index c18b4415f1..a58213d651 100644 --- a/src/backend/base/langflow/services/database/models/jobs/model.py +++ b/src/backend/base/langflow/services/database/models/jobs/model.py @@ -63,6 +63,9 @@ class JobBase(SQLModel): asset_type: str | None = Field( index=False, nullable=True ) # Polymorphic: records if job is related to an entity like a KB, workflow, etc. + dedupe_key: str | None = Field( + index=True, nullable=True + ) # Optional idempotency key to prevent duplicate jobs for the same asset and operation. class Job(JobBase, table=True): # type: ignore[call-arg] diff --git a/src/backend/base/langflow/services/database/models/memory_base/__init__.py b/src/backend/base/langflow/services/database/models/memory_base/__init__.py new file mode 100644 index 0000000000..842eb8826f --- /dev/null +++ b/src/backend/base/langflow/services/database/models/memory_base/__init__.py @@ -0,0 +1,21 @@ +from langflow.services.database.models.memory_base.model import ( + MemoryBase, + MemoryBaseCreate, + MemoryBaseRead, + MemoryBaseSession, + MemoryBaseSessionRead, + MemoryBaseUpdate, + MemoryBaseWorkflowRun, + MessageIngestionRecord, +) + +__all__ = [ + "MemoryBase", + "MemoryBaseCreate", + "MemoryBaseRead", + "MemoryBaseSession", + "MemoryBaseSessionRead", + "MemoryBaseUpdate", + "MemoryBaseWorkflowRun", + "MessageIngestionRecord", +] diff --git a/src/backend/base/langflow/services/database/models/memory_base/model.py b/src/backend/base/langflow/services/database/models/memory_base/model.py new file mode 100644 index 0000000000..f1bb766204 --- /dev/null +++ b/src/backend/base/langflow/services/database/models/memory_base/model.py @@ -0,0 +1,198 @@ +from datetime import datetime, timezone +from uuid import UUID, uuid4 + +import sqlalchemy as sa +from pydantic import model_validator +from sqlalchemy import Column, DateTime, ForeignKey, Index, UniqueConstraint +from sqlmodel import Field, Relationship, SQLModel + + +class MemoryBaseBase(SQLModel): + name: str = Field(index=False) + flow_id: UUID = Field(index=True) + user_id: UUID = Field(index=True) + threshold: int = Field(default=50) + auto_capture: bool = Field(default=True) + # Preprocessing config — accepted in payload but logic deferred to future scope + embedding_model: str = Field(default="") + preprocessing: bool = Field(default=False) + preproc_model: str | None = Field(default=None) + preproc_instructions: str | None = Field(default=None) + + +class MemoryBase(MemoryBaseBase, table=True): # type: ignore[call-arg] + __tablename__ = "memory_base" + + id: UUID = Field(default_factory=uuid4, primary_key=True) + # kb_name is auto-generated at creation time — not user-supplied + kb_name: str = Field(default="") + created_at: datetime = Field( + default_factory=lambda: datetime.now(timezone.utc), + sa_column=Column(DateTime(timezone=True), nullable=False), + ) + + sessions: list["MemoryBaseSession"] = Relationship( + back_populates="memory_base", + sa_relationship_kwargs={"cascade": "all, delete-orphan"}, + ) + + +class MemoryBaseCreate(MemoryBaseBase): + user_id: UUID | None = None # Derived from auth token in the endpoint; not required in request body + + @model_validator(mode="after") + def preproc_model_required_when_preprocessing(self) -> "MemoryBaseCreate": + if self.preprocessing and not self.preproc_model: + msg = "preproc_model is required when preprocessing is enabled" + raise ValueError(msg) + return self + + +class MemoryBaseUpdate(SQLModel): + name: str | None = None + threshold: int | None = None + auto_capture: bool | None = None + preprocessing: bool | None = None + preproc_model: str | None = None + preproc_instructions: str | None = None + + +class MemoryBaseRead(MemoryBaseBase): + id: UUID + kb_name: str + created_at: datetime + + +class MemoryBaseSessionBase(SQLModel): + """Fields shared between the table class and response schemas.""" + + session_id: str = Field(index=True) + cursor_id: UUID | None = Field(default=None) + total_processed: int = Field(default=0) + last_sync_at: datetime | None = Field( + default=None, + sa_column=Column(DateTime(timezone=True), nullable=True), + ) + + +class MemoryBaseSession(MemoryBaseSessionBase, table=True): # type: ignore[call-arg] + __tablename__ = "memory_base_session" + + __table_args__ = ( + UniqueConstraint("memory_base_id", "session_id", name="uq_memory_base_session"), + Index("ix_memory_base_session_lookup", "memory_base_id", "session_id"), + ) + + id: UUID = Field(default_factory=uuid4, primary_key=True) + + # FK defined via sa_column so Alembic sees the same shape as the migration: + # inline ForeignKey on the column with ondelete="CASCADE". + # This matches the pattern used by the File model (ForeignKey on sa_column). + memory_base_id: UUID = Field( + sa_column=Column( + sa.Uuid(), + ForeignKey("memory_base.id", ondelete="CASCADE"), + nullable=False, + index=True, + ) + ) + + memory_base: MemoryBase = Relationship(back_populates="sessions") + + +class MemoryBaseSessionRead(MemoryBaseSessionBase): + id: UUID + memory_base_id: UUID # Explicit — not in base to keep base free of DB-layer FK + pending_count: int = Field(default=0) + + +class MemoryBaseWorkflowRun(SQLModel, table=True): # type: ignore[call-arg] + """Tracks WORKFLOW job runs per (memory_base, session) for threshold-based ingestion. + + One row per WORKFLOW job, per session, per memory base. + - ``workflow_job_id``: the WORKFLOW job that produced this run (SET NULL on job deletion). + - ``ingestion_job_id``: set only after the ingestion job that processed this run completes + successfully. NULL means the run is still pending (not yet counted toward an ingestion). + + Count pending = COUNT(*) WHERE ingestion_job_id IS NULL for a given (memory_base_id, session_id). + """ + + __tablename__ = "memory_base_workflow_run" + __table_args__ = ( + UniqueConstraint("memory_base_id", "session_id", "workflow_job_id", name="uq_mbwr_mb_session_wf_job"), + Index("ix_mbwr_mb_session", "memory_base_id", "session_id"), + Index("ix_mbwr_ingestion_job_id", "ingestion_job_id"), + ) + + id: UUID = Field(default_factory=uuid4, primary_key=True) + memory_base_id: UUID = Field( + sa_column=Column( + sa.Uuid(), + ForeignKey("memory_base.id", ondelete="CASCADE"), + nullable=False, + ) + ) + session_id: str = Field(sa_column=Column(sa.String(), nullable=False)) + workflow_job_id: UUID | None = Field( + default=None, + sa_column=Column( + sa.Uuid(), + ForeignKey("job.job_id", ondelete="SET NULL"), + nullable=True, + ), + ) + ingestion_job_id: UUID | None = Field( + default=None, + sa_column=Column( + sa.Uuid(), + ForeignKey("job.job_id", ondelete="SET NULL"), + nullable=True, + ), + ) + recorded_at: datetime = Field(sa_column=Column(DateTime(timezone=True), nullable=False)) + + +class MessageIngestionRecord(SQLModel, table=True): # type: ignore[call-arg] + """M-N join table recording which messages were ingested into which Memory Base by which job. + + One record per (message, session, memory_base) — enforced by the unique constraint. + Records are written only after a confirmed successful Chroma write (write-on-success). + On regenerate, all records for the memory_base are deleted atomically alongside the + cursor reset so that re-ingestion starts clean. + """ + + __tablename__ = "message_ingestion_record" + __table_args__ = ( + UniqueConstraint("message_id", "session_id", "memory_base_id", name="uq_mir_message_session_mb"), + Index("ix_mir_message_id", "message_id"), + Index("ix_mir_job_id", "job_id"), + Index("ix_mir_memory_base_session", "memory_base_id", "session_id"), + ) + + id: UUID = Field(default_factory=uuid4, primary_key=True) + + message_id: UUID = Field( + sa_column=Column( + sa.Uuid(), + ForeignKey("message.id", ondelete="CASCADE"), + nullable=False, + ) + ) + memory_base_id: UUID = Field( + sa_column=Column( + sa.Uuid(), + ForeignKey("memory_base.id", ondelete="CASCADE"), + nullable=False, + ) + ) + job_id: UUID | None = Field( + default=None, + sa_column=Column( + sa.Uuid(), + ForeignKey("job.job_id", ondelete="SET NULL"), + nullable=True, + ), + ) + # Denormalized from MessageTable.session_id — immutable, avoids JOIN on the hot query path + session_id: str = Field(sa_column=Column(sa.String(), nullable=False)) + ingested_at: datetime = Field(sa_column=Column(DateTime(timezone=True), nullable=False)) diff --git a/src/backend/base/langflow/services/database/models/message/model.py b/src/backend/base/langflow/services/database/models/message/model.py index 52bfd68267..1a8b71d580 100644 --- a/src/backend/base/langflow/services/database/models/message/model.py +++ b/src/backend/base/langflow/services/database/models/message/model.py @@ -32,6 +32,7 @@ class MessageBase(SQLModel): properties: Properties = Field(default_factory=Properties) category: str = Field(default="message") content_blocks: list[ContentBlock] = Field(default_factory=list) + session_metadata: dict | None = Field(default=None) @field_serializer("timestamp") def serialize_timestamp(self, value): @@ -149,8 +150,11 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg] id: UUID = Field(default_factory=uuid4, primary_key=True) flow_id: UUID | None = Field(default=None) + run_id: UUID | None = Field(default=None, index=True) + is_output: bool = Field(default=False) files: list[str] = Field(sa_column=Column(JSON)) + session_metadata: dict | None = Field(default=None, sa_column=Column(JSON)) properties: dict | Properties = Field( # type: ignore[assignment] default_factory=lambda: Properties().model_dump(), sa_column=Column(JSON), @@ -161,13 +165,6 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg] sa_column=Column(JSON), ) - # Enterprise session metadata - flexible JSON column for client-provided context - session_metadata: dict | None = Field( - default=None, - sa_column=Column(JSON), - description="Session context data (e.g., user roles, custom tags, or analytics data).", - ) - @field_validator("flow_id", mode="before") @classmethod def validate_flow_id(cls, value): @@ -193,7 +190,7 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg] return value - @field_validator("properties", "content_blocks", "session_metadata", mode="before") + @field_validator("properties", "content_blocks", mode="before") @classmethod def validate_properties_or_content_blocks(cls, value): if isinstance(value, list): @@ -205,13 +202,11 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg] return cls._sanitize_json(value) - @field_serializer("properties", "content_blocks", "session_metadata") + @field_serializer("properties", "content_blocks") @classmethod - def serialize_properties_or_content_blocks(cls, value) -> dict | list[dict] | None: + def serialize_properties_or_content_blocks(cls, value) -> dict | list[dict]: # Redundant sanitization here acts as a defensive measure for rows # already in the database that might contain NaN/Infinity values. - if value is None: - return None if isinstance(value, list): value = [cls.serialize_properties_or_content_blocks(item) for item in value] elif hasattr(value, "model_dump"): @@ -225,11 +220,10 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg] class MessageRead(MessageBase): id: UUID flow_id: UUID | None = Field() - session_metadata: dict | None = None class MessageCreate(MessageBase): - session_metadata: dict | None = None + pass class MessageUpdate(SQLModel): @@ -242,4 +236,3 @@ class MessageUpdate(SQLModel): edit: bool | None = None error: bool | None = None properties: Properties | None = None - session_metadata: dict | None = None diff --git a/src/backend/base/langflow/services/deps.py b/src/backend/base/langflow/services/deps.py index 14f9bcbdc9..194d4e621d 100644 --- a/src/backend/base/langflow/services/deps.py +++ b/src/backend/base/langflow/services/deps.py @@ -269,3 +269,20 @@ def get_flow_events_service(): from langflow.services.flow_events.factory import FlowEventsServiceFactory return get_service(ServiceType.FLOW_EVENTS_SERVICE, FlowEventsServiceFactory()) + + +_memory_base_service = None + + +def get_memory_base_service(): + """Returns the singleton MemoryBaseService instance. + + MemoryBaseService is stateless (all state lives in the DB) so a module-level + singleton is sufficient — no service manager or factory required. + """ + global _memory_base_service # noqa: PLW0603 + if _memory_base_service is None: + from langflow.services.memory_base.service import MemoryBaseService + + _memory_base_service = MemoryBaseService() + return _memory_base_service diff --git a/src/backend/base/langflow/services/jobs/__init__.py b/src/backend/base/langflow/services/jobs/__init__.py index 8fd4b155b2..53e37c4096 100644 --- a/src/backend/base/langflow/services/jobs/__init__.py +++ b/src/backend/base/langflow/services/jobs/__init__.py @@ -1,5 +1,6 @@ """Job service package.""" +from langflow.services.jobs.exceptions import DuplicateJobError from langflow.services.jobs.service import JobService -__all__ = ["JobService"] +__all__ = ["DuplicateJobError", "JobService"] diff --git a/src/backend/base/langflow/services/jobs/exceptions.py b/src/backend/base/langflow/services/jobs/exceptions.py new file mode 100644 index 0000000000..385a596a1f --- /dev/null +++ b/src/backend/base/langflow/services/jobs/exceptions.py @@ -0,0 +1,16 @@ +"""Domain exceptions for the jobs service.""" + +from __future__ import annotations + + +class JobError(RuntimeError): + """Base exception for job-domain errors.""" + + +class DuplicateJobError(JobError): + """Raised by create_job() when a non-retryable job with the same dedupe_key already exists. + + (QUEUED, IN_PROGRESS, or COMPLETED). + FAILED and CANCELLED are retryable and do not trigger this error. + Extends RuntimeError so existing except RuntimeError callers keep working. + """ diff --git a/src/backend/base/langflow/services/jobs/service.py b/src/backend/base/langflow/services/jobs/service.py index 914e683e83..b258aec8c9 100644 --- a/src/backend/base/langflow/services/jobs/service.py +++ b/src/backend/base/langflow/services/jobs/service.py @@ -10,15 +10,16 @@ if TYPE_CHECKING: from datetime import datetime, timezone from uuid import UUID +from sqlmodel import col, func, select + from langflow.services.base import Service from langflow.services.database.models.jobs.crud import ( - get_job_by_job_id, - get_jobs_by_flow_id, get_latest_jobs_by_asset_ids, update_job_status, ) from langflow.services.database.models.jobs.model import Job, JobStatus, JobType from langflow.services.deps import session_scope +from langflow.services.jobs.exceptions import DuplicateJobError class JobService(Service): @@ -30,11 +31,14 @@ class JobService(Service): """Initialize the job service.""" self.set_ready() - async def get_jobs_by_flow_id(self, flow_id: UUID | str, page: int = 1, page_size: int = 10) -> list[Job]: - """Get jobs for a specific flow with pagination. + async def get_jobs_by_flow_id( + self, flow_id: UUID | str, user_id: UUID, page: int = 1, page_size: int = 10 + ) -> list[Job]: + """Get jobs for a specific flow with pagination, filtered by user. Args: flow_id: The flow ID to filter jobs by + user_id: The user ID to enforce ownership page: Page number (1-indexed) page_size: Number of jobs per page @@ -45,7 +49,16 @@ class JobService(Service): flow_id = UUID(flow_id) async with session_scope() as session: - return await get_jobs_by_flow_id(session, flow_id, page=page, size=page_size) + stmt = ( + select(Job) + .where(Job.flow_id == flow_id) + .where((Job.user_id == user_id) | (Job.user_id.is_(None))) + .order_by(col(Job.created_at).desc()) + .offset((page - 1) * page_size) + .limit(page_size) + ) + result = await session.exec(stmt) + return list(result.all()) async def get_job_by_job_id(self, job_id: UUID | str, user_id: UUID | None = None) -> Job | None: """Get job for a specific job ID. @@ -62,7 +75,11 @@ class JobService(Service): job_id = UUID(job_id) async with session_scope() as session: - return await get_job_by_job_id(session, job_id, user_id=user_id) + stmt = select(Job).where(Job.job_id == job_id) + if user_id: + stmt = stmt.where((Job.user_id == user_id) | (Job.user_id.is_(None))) + result = await session.exec(stmt) + return result.first() async def create_job( self, @@ -72,16 +89,19 @@ class JobService(Service): asset_id: UUID | None = None, asset_type: str | None = None, user_id: UUID | None = None, + dedupe_key: str | None = None, ) -> Job: """Create a new job record with QUEUED status. Args: job_id: The job ID flow_id: The flow ID + user_id: The user ID job_type: The job type asset_id: The asset ID asset_type: The asset type user_id: The user ID who owns this job + dedupe_key: Optional idempotency key to prevent duplicate jobs for the same batch Returns: Created Job object @@ -93,6 +113,18 @@ class JobService(Service): flow_id = UUID(flow_id) async with session_scope() as session: + if dedupe_key is not None: + stmt = ( + select(func.count()) + .select_from(Job) + .where(Job.dedupe_key == dedupe_key) + .where(col(Job.status).in_([JobStatus.QUEUED, JobStatus.IN_PROGRESS, JobStatus.COMPLETED])) + ) + result = await session.exec(stmt) + if result.one() > 0: + msg = f"A non-retryable job with dedupe_key={dedupe_key!r} already exists" + raise DuplicateJobError(msg) + job = Job( job_id=job_id, flow_id=flow_id, @@ -101,6 +133,7 @@ class JobService(Service): asset_id=asset_id, asset_type=asset_type, user_id=user_id, + dedupe_key=dedupe_key, ) session.add(job) await session.flush() @@ -210,3 +243,18 @@ class JobService(Service): await logger.ainfo(f"Job {job_id} completed successfully") await self.update_job_status(job_id, JobStatus.COMPLETED, finished_timestamp=True) return result + + async def _validate_ownership(self, job_id: UUID, user_id: UUID) -> Job: + """Verify that a job exists and belongs to the specified user. + + Raises: + ValueError: If the job is not found or is NOT owned by the user. + """ + job = await self.get_job_by_job_id(job_id) + if job is None: + msg = f"Job {job_id} not found" + raise ValueError(msg) + if job.user_id is not None and job.user_id != user_id: + msg = f"Access denied for job {job_id}" + raise ValueError(msg) + return job diff --git a/src/backend/base/langflow/services/memory_base/__init__.py b/src/backend/base/langflow/services/memory_base/__init__.py new file mode 100644 index 0000000000..c7b1f2c756 --- /dev/null +++ b/src/backend/base/langflow/services/memory_base/__init__.py @@ -0,0 +1,3 @@ +from langflow.services.memory_base.service import MemoryBaseService + +__all__ = ["MemoryBaseService"] diff --git a/src/backend/base/langflow/services/memory_base/service.py b/src/backend/base/langflow/services/memory_base/service.py new file mode 100644 index 0000000000..0e71aef4a8 --- /dev/null +++ b/src/backend/base/langflow/services/memory_base/service.py @@ -0,0 +1,705 @@ +"""MemoryBase service - business logic for CRUD and ingestion orchestration. + +Edge cases handled: +- Name uniqueness per user: 409 if a Memory Base with the same name already exists. +- Deletion during sync: cancels active tasks before DB deletion. +- KB deletion on delete: removes the associated KB directory from disk. +- Concurrent task prevention: returns 409 if a job is already IN_PROGRESS. +- Threshold updates: deferred; does not re-evaluate pending count immediately. +- FS / Vector DB mismatch: detects and surfaces a warning flag. +- Regenerate: resets all session cursors to None and re-triggers ingestion. +""" + +from __future__ import annotations + +import json +import re +import uuid +from datetime import datetime, timezone +from typing import TYPE_CHECKING + +from lfx.log.logger import logger +from sqlmodel import col, func, select + +from langflow.api.utils.kb_helpers import KBAnalysisHelper, KBStorageHelper +from langflow.services.database.models.jobs.model import Job, JobStatus, JobType +from langflow.services.database.models.memory_base.model import ( + MemoryBase, + MemoryBaseCreate, + MemoryBaseSession, + MemoryBaseUpdate, + MemoryBaseWorkflowRun, +) +from langflow.services.deps import get_job_service, get_task_service, session_scope +from langflow.services.jobs import DuplicateJobError +from langflow.services.memory_base.task import ingest_memory_task + +if TYPE_CHECKING: + from pathlib import Path + + from sqlmodel.ext.asyncio.session import AsyncSession + +# Provider inference map — mirrors provider_patterns in KBAnalysisHelper._detect_embedding_provider +# so we can derive the provider from a model name string without filesystem access. +_MODEL_TO_PROVIDER: list[tuple[list[str], str]] = [ + (["text-embedding", "ada-", "gpt-"], "OpenAI"), + (["embed-english", "embed-multilingual"], "Cohere"), + (["sentence-transformers", "bert-", "huggingface"], "HuggingFace"), + (["palm", "gecko", "google"], "Google"), + (["ollama"], "Ollama"), + (["azure"], "Azure OpenAI"), +] + + +def _infer_embedding_provider(embedding_model: str) -> str: + """Derive embedding provider name from a model string.""" + lower = embedding_model.lower() + for patterns, provider in _MODEL_TO_PROVIDER: + if any(p in lower for p in patterns): + return provider + return "OpenAI" # Safe default — matches _resolve_embedding fallback + + +def _sanitize_kb_name(name: str) -> str: + """Lowercase, replace spaces/hyphens with underscores, strip non-alphanum.""" + sanitized = name.strip().lower() + sanitized = re.sub(r"[\s\-]+", "_", sanitized) + sanitized = re.sub(r"[^\w]", "", sanitized) + return sanitized or "memory" + + +class MemoryBaseService: + """Service layer for MemoryBase CRUD and ingestion orchestration.""" + + # ------------------------------------------------------------------ # + # CRUD # + # ------------------------------------------------------------------ # + + async def create(self, payload: MemoryBaseCreate, user_id: uuid.UUID) -> MemoryBase: + # 1. Verify that the referenced flow belongs to this user. Without this + # check a user could point a Memory Base at another user's flow and then + # read captured conversation history via /sessions + /messages. + # Raises PermissionError (→ 404 at the API layer to avoid info-leak) if + # the flow does not exist or is not owned by user_id. + async with session_scope() as db: + from langflow.services.database.models.flow.model import Flow + + flow_result = await db.exec(select(Flow).where(Flow.id == payload.flow_id).where(Flow.user_id == user_id)) + if flow_result.first() is None: + msg = f"Flow {payload.flow_id} not found" + raise PermissionError(msg) + + # 2. Resolve username — needed for the KB path, separate session so the + # connection is not held open during disk I/O below. + async with session_scope() as db: + kb_username = await self._resolve_kb_username(db, user_id) + + # 2. Auto-generate kb_name: sanitized_name_<8hex> + kb_name = f"{_sanitize_kb_name(payload.name)}_{uuid.uuid4().hex[:8]}" + + # 3. Create KB directory and embedding_metadata.json on disk (best-effort). + # KB init runs before the DB insert — if the insert later fails, the + # orphaned directory is harmless (nothing writes to it without a DB record). + embedding_provider = _infer_embedding_provider(payload.embedding_model) + await self._initialize_kb( + kb_name=kb_name, + kb_username=kb_username, + embedding_provider=embedding_provider, + embedding_model=payload.embedding_model, + ) + + # 4. Uniqueness check + insert in ONE session_scope to close the TOCTOU window. + # Both the SELECT and the INSERT see the same DB snapshot; a concurrent + # create with the same name is caught by the re-check here. + async with session_scope() as db: + existing = await db.exec( + select(MemoryBase).where(MemoryBase.user_id == user_id).where(MemoryBase.name == payload.name) + ) + if existing.first() is not None: + msg = f"A Memory Base named '{payload.name}' already exists for this user" + raise ValueError(msg) + + mb = MemoryBase( + **payload.model_dump(exclude={"user_id"}), + user_id=user_id, + kb_name=kb_name, + ) + db.add(mb) + await db.commit() + await db.refresh(mb) + + return mb + + async def _initialize_kb( + self, + *, + kb_name: str, + kb_username: str, + embedding_provider: str, + embedding_model: str, + ) -> None: + """Create KB directory, initialize Chroma, and write embedding_metadata.json. + + Mirrors the logic in knowledge_bases.py:create_knowledge_base so Memory Base + KBs are immediately visible with the correct metadata (including is_memory_base: true). + """ + import chromadb + + kb_root = KBStorageHelper.get_root_path() + if not kb_root: + await logger.awarning("KB root path not configured — Memory Base KB will not be initialized on disk.") + return + + kb_path: Path = kb_root / kb_username / kb_name + kb_path.mkdir(parents=True, exist_ok=True) + + # Initialize Chroma collection so the directory is non-empty and readable + try: + client = KBStorageHelper.get_fresh_chroma_client(kb_path) + client.create_collection(name=kb_name) + except (OSError, ValueError, chromadb.errors.ChromaError) as exc: + await logger.awarning("Initial Chroma setup for %s failed: %s", kb_name, exc) + finally: + client = None # type: ignore[assignment] + KBStorageHelper.release_chroma_resources(kb_path) + + embedding_metadata = { + "id": str(uuid.uuid4()), + "embedding_provider": embedding_provider, + "embedding_model": embedding_model, + "is_memory_base": True, + "created_at": datetime.now(timezone.utc).isoformat(), + "chunks": 0, + "words": 0, + "characters": 0, + "avg_chunk_size": 0.0, + "size": 0, + "source_types": ["memory"], + } + (kb_path / "embedding_metadata.json").write_text(json.dumps(embedding_metadata, indent=2)) + + async def list_for_user(self, user_id: uuid.UUID) -> list[MemoryBase]: + async with session_scope() as db: + stmt = select(MemoryBase).where(MemoryBase.user_id == user_id) + result = await db.exec(stmt) + return list(result.all()) + + def list_for_user_stmt(self, user_id: uuid.UUID, flow_id: uuid.UUID | None = None): # type: ignore[return] + """Return the SQLModel select statement for pagination at the API layer.""" + stmt = select(MemoryBase).where(MemoryBase.user_id == user_id) + if flow_id is not None: + stmt = stmt.where(MemoryBase.flow_id == flow_id) + return stmt + + async def get(self, memory_base_id: uuid.UUID, user_id: uuid.UUID) -> MemoryBase | None: + async with session_scope() as db: + stmt = select(MemoryBase).where(MemoryBase.id == memory_base_id).where(MemoryBase.user_id == user_id) + result = await db.exec(stmt) + return result.first() + + async def update( + self, + memory_base_id: uuid.UUID, + user_id: uuid.UUID, + patch: MemoryBaseUpdate, + ) -> MemoryBase | None: + """Update mutable fields. + + Threshold changes take effect on the NEXT auto-capture trigger; any + already-running ingestion task ignores the change (immutable args). + """ + async with session_scope() as db: + stmt = select(MemoryBase).where(MemoryBase.id == memory_base_id).where(MemoryBase.user_id == user_id) + result = await db.exec(stmt) + mb = result.first() + if mb is None: + return None + for field, value in patch.model_dump(exclude_unset=True).items(): + setattr(mb, field, value) + db.add(mb) + await db.commit() + await db.refresh(mb) + return mb + + async def delete(self, memory_base_id: uuid.UUID, user_id: uuid.UUID) -> bool: + """Delete a MemoryBase and its associated KB directory. + + Edge cases: + - If a sync task is active, cancel it BEFORE committing the DB deletion. + - KB directory deletion is best-effort after the DB commit — a failure + is logged but not re-raised so the caller always gets a clean 204. + """ + async with session_scope() as db: + stmt = select(MemoryBase).where(MemoryBase.id == memory_base_id).where(MemoryBase.user_id == user_id) + result = await db.exec(stmt) + mb = result.first() + if mb is None: + return False + + kb_name = mb.kb_name + kb_username = await self._resolve_kb_username(db, user_id) + + # Cancel active ingestion jobs before removing the DB record + await self._cancel_active_jobs(memory_base_id=memory_base_id, db=db) + + await db.delete(mb) + await db.commit() + + # Delete the corresponding KB from disk (best-effort — DB already committed) + await self._delete_kb(kb_name=kb_name, kb_username=kb_username) + + return True + + async def _delete_kb(self, *, kb_name: str, kb_username: str) -> None: + """Remove the KB directory from disk. Logs on failure, does not raise.""" + if not kb_name: + return + kb_root = KBStorageHelper.get_root_path() + if not kb_root: + return + kb_path = kb_root / kb_username / kb_name + try: + KBStorageHelper.delete_storage(kb_path, kb_name) + except (OSError, ValueError): + await logger.awarning( + "Could not delete KB '%s' from disk after Memory Base deletion.", kb_name, exc_info=True + ) + + # ------------------------------------------------------------------ # + # Sessions # + # ------------------------------------------------------------------ # + + async def verify_ownership(self, memory_base_id: uuid.UUID, user_id: uuid.UUID) -> None: + """Raise ValueError if the Memory Base does not belong to user_id.""" + async with session_scope() as db: + await self._get_mb_or_raise(db, memory_base_id, user_id) + + def sessions_stmt(self, memory_base_id: uuid.UUID): # type: ignore[return] + """Return the select statement for persisted sessions, for use with apaginate.""" + return ( + select(MemoryBaseSession) + .where(MemoryBaseSession.memory_base_id == memory_base_id) + .order_by(col(MemoryBaseSession.last_sync_at).desc()) + ) + + # ------------------------------------------------------------------ # + # Ingestion # + # ------------------------------------------------------------------ # + + async def trigger_ingestion( + self, + memory_base_id: uuid.UUID, + user_id: uuid.UUID, + session_id: str, + ) -> str: + """Manually trigger (or auto-trigger) an ingestion sync. + + Returns: + job_id string for the newly created job. + + Raises: + ValueError: If MemoryBase not found. + RuntimeError: If a job is already active (caller should return 409). + """ + async with session_scope() as db: + mb = await self._get_mb_or_raise(db, memory_base_id, user_id) + + # Ensure a session record exists + mbs = await self._get_or_create_session(db, memory_base_id, session_id) + + # Snapshot the cursor NOW (immutable arg for the task) + cursor_id_snapshot = mbs.cursor_id + + # Build dedupe_key from the latest uncovered WORKFLOW run for idempotency. + # Format: "ingestion:{memory_base_id}:{session_id}:{latest_workflow_job_id}" + # Uniqueness: job type prefix + MB scope + session scope + batch identity. + # Dedup enforcement (QUEUED/IN_PROGRESS/COMPLETED) happens inside create_job(). + # Retries are allowed when the prior job was FAILED or CANCELLED. + latest_job_id = await self._get_latest_pending_workflow_job_id(db, mb, mbs) + dedupe_key: str | None = None + if latest_job_id is not None: + dedupe_key = f"ingestion:{memory_base_id}:{session_id}:{latest_job_id}" + + kb_username = await self._resolve_kb_username(db, mb.user_id) + embedding_provider, embedding_model = self._resolve_embedding(mb.kb_name, kb_username) + + # Create tracking job + job_service = get_job_service() + job_id = uuid.uuid4() + await job_service.create_job( + job_id=job_id, + flow_id=mb.flow_id, + user_id=mb.user_id, + job_type=JobType.INGESTION, + asset_id=memory_base_id, + asset_type="memory_base", + dedupe_key=dedupe_key, + ) + + task_service = get_task_service() + await task_service.fire_and_forget_task( + job_service.execute_with_status, + job_id=job_id, + run_coro_func=ingest_memory_task, + memory_base_id=memory_base_id, + session_id=session_id, + flow_id=mb.flow_id, + kb_name=mb.kb_name, + kb_username=kb_username, + user_id=mb.user_id, + embedding_provider=embedding_provider, + embedding_model=embedding_model, + cursor_id=cursor_id_snapshot, + task_job_id=job_id, + job_service=job_service, + ) + + return str(job_id) + + # ------------------------------------------------------------------ # + # Auto-capture hook (called from flow execution engine) # + # ------------------------------------------------------------------ # + + async def on_flow_output( + self, + flow_id: uuid.UUID, + session_id: str, + job_id: uuid.UUID | None, + ) -> None: + """Called after a flow run completes. + + For every MemoryBase watching this flow with auto_capture=True: + 1. Record the workflow run in the tracking table (inside _maybe_trigger). + 2. Count uncovered WORKFLOW runs for this session. + 3. If count >= threshold, fire ingestion task. + """ + async with session_scope() as db: + stmt = ( + select(MemoryBase).where(MemoryBase.flow_id == flow_id).where(MemoryBase.auto_capture == True) # noqa: E712 + ) + result = await db.exec(stmt) + memory_bases = list(result.all()) + + for mb in memory_bases: + try: + await logger.adebug( + "Auto-capture check | memory_base=%s name=%r threshold=%s session=%s", + mb.id, + mb.name, + mb.threshold, + session_id, + ) + await self._maybe_trigger(mb=mb, session_id=session_id, job_id=job_id) + except (RuntimeError, ValueError, OSError): + await logger.aerror( + "Auto-capture failed for memory_base=%s session=%s", mb.id, session_id, exc_info=True + ) + + async def _maybe_trigger(self, *, mb: MemoryBase, session_id: str, job_id: uuid.UUID | None) -> None: + async with session_scope() as db: + mbs = await self._get_or_create_session(db, mb.id, session_id) + + # Record this workflow run before evaluating the threshold. + await self._insert_workflow_run(db, mb.id, session_id, job_id) + + pending = await self._count_pending(db, mb, mbs) + + if pending < mb.threshold: + return + + cursor_id_snapshot = mbs.cursor_id + + # Build dedupe_key from the latest pending WORKFLOW run for idempotency. + latest_wf_job_id = await self._get_latest_pending_workflow_job_id(db, mb, mbs) + dedupe_key: str | None = None + if latest_wf_job_id is not None: + dedupe_key = f"ingestion:{mb.id}:{session_id}:{latest_wf_job_id}" + + kb_username = await self._resolve_kb_username(db, mb.user_id) + + embedding_provider, embedding_model = self._resolve_embedding(mb.kb_name, kb_username) + + job_service = get_job_service() + job_id = uuid.uuid4() + try: + await job_service.create_job( + job_id=job_id, + flow_id=mb.flow_id, + user_id=mb.user_id, + job_type=JobType.INGESTION, + asset_id=mb.id, + asset_type="memory_base", + dedupe_key=dedupe_key, + ) + except DuplicateJobError: + await logger.adebug("Auto-capture: duplicate job for dedupe_key=%s - skipping.", dedupe_key) + return + + task_service = get_task_service() + await task_service.fire_and_forget_task( + job_service.execute_with_status, + job_id=job_id, + run_coro_func=ingest_memory_task, + memory_base_id=mb.id, + session_id=session_id, + flow_id=mb.flow_id, + kb_name=mb.kb_name, + kb_username=kb_username, + user_id=mb.user_id, + embedding_provider=embedding_provider, + embedding_model=embedding_model, + cursor_id=cursor_id_snapshot, + task_job_id=job_id, + job_service=job_service, + ) + + # ------------------------------------------------------------------ # + # FS / Vector DB mismatch detection # + # ------------------------------------------------------------------ # + + async def check_mismatch(self, memory_base_id: uuid.UUID, user_id: uuid.UUID) -> bool: + """Return True if metadata claims processed rows but vector store is empty. + + The UI should surface a "Mismatch Detected" warning and offer Regenerate. + """ + async with session_scope() as db: + mb = await self._get_mb_or_raise(db, memory_base_id, user_id) + stmt = select(func.sum(MemoryBaseSession.total_processed)).where( + MemoryBaseSession.memory_base_id == memory_base_id + ) + result = await db.exec(stmt) + total_processed: int = result.first() or 0 + + if total_processed == 0: + return False + + kb_username = await self._resolve_kb_username_by_user_id(user_id) + kb_root = KBStorageHelper.get_root_path() + if not kb_root: + return False + kb_path = kb_root / kb_username / mb.kb_name + if not kb_path.exists(): + return True + + metadata = KBAnalysisHelper.get_metadata(kb_path, fast=True) + return int(metadata.get("chunks", 0)) == 0 + + async def regenerate(self, memory_base_id: uuid.UUID, user_id: uuid.UUID) -> list[str]: + """Reset all session cursors to None and re-trigger ingestion per session. + + Used to recover from FS / Vector DB mismatch (Chroma dir deleted externally). + Returns list of newly created job IDs. + Also deletes all MessageIngestionRecord rows for this memory base atomically + with the cursor reset so that re-ingestion starts clean without hitting the + unique constraint. + """ + from sqlalchemy import delete as sa_delete + + from langflow.services.database.models.memory_base.model import MessageIngestionRecord + + async with session_scope() as db: + await self._get_mb_or_raise(db, memory_base_id, user_id) + + stmt = select(MemoryBaseSession).where(MemoryBaseSession.memory_base_id == memory_base_id) + result = await db.exec(stmt) + sessions = list(result.all()) + + for s in sessions: + s.cursor_id = None + db.add(s) + + # Delete existing ingestion records so re-ingestion inserts fresh rows + await db.exec( # type: ignore[call-overload] + sa_delete(MessageIngestionRecord).where(MessageIngestionRecord.memory_base_id == memory_base_id) + ) + await db.commit() + + job_ids: list[str] = [] + for s in sessions: + try: + jid = await self.trigger_ingestion(memory_base_id, user_id, s.session_id) + job_ids.append(jid) + except DuplicateJobError: + await logger.awarning( + "Regenerate: duplicate batch already ingested for session %s - skipped.", s.session_id + ) + except RuntimeError: + await logger.awarning( + "Regenerate: active job exists for session %s - reset cursor but skipped trigger.", s.session_id + ) + return job_ids + + # ------------------------------------------------------------------ # + # Internal helpers # + # ------------------------------------------------------------------ # + + async def _get_mb_or_raise(self, db: AsyncSession, memory_base_id: uuid.UUID, user_id: uuid.UUID) -> MemoryBase: + stmt = select(MemoryBase).where(MemoryBase.id == memory_base_id).where(MemoryBase.user_id == user_id) + result = await db.exec(stmt) + mb = result.first() + if mb is None: + msg = f"MemoryBase {memory_base_id} not found" + raise ValueError(msg) + return mb + + async def _get_or_create_session( + self, db: AsyncSession, memory_base_id: uuid.UUID, session_id: str + ) -> MemoryBaseSession: + stmt = ( + select(MemoryBaseSession) + .where(MemoryBaseSession.memory_base_id == memory_base_id) + .where(MemoryBaseSession.session_id == session_id) + ) + result = await db.exec(stmt) + mbs = result.first() + if mbs is None: + mbs = MemoryBaseSession(memory_base_id=memory_base_id, session_id=session_id) + db.add(mbs) + await db.commit() + await db.refresh(mbs) + return mbs + + async def _insert_workflow_run( + self, + db: AsyncSession, + memory_base_id: uuid.UUID, + session_id: str, + job_id: uuid.UUID | None, + ) -> None: + """Record a WORKFLOW job run for (memory_base_id, session_id). + + Verifies that job_id refers to a WORKFLOW type job before inserting. + Uses dialect-specific INSERT ... ON CONFLICT DO NOTHING for idempotency — + safe to call multiple times with the same arguments. + Skips silently if job_id is None or the job is not of WORKFLOW type. + """ + if job_id is None: + await logger.awarning( + "on_flow_output called with no job_id for memory_base=%s session=%s — run not recorded.", + memory_base_id, + session_id, + ) + return + + job_result = await db.exec(select(Job).where(Job.job_id == job_id).where(Job.type == JobType.WORKFLOW)) + if job_result.first() is None: + await logger.awarning( + "job_id=%s is not a WORKFLOW job — skipping workflow run record for memory_base=%s session=%s.", + job_id, + memory_base_id, + session_id, + ) + return + + row = { + "id": uuid.uuid4(), + "memory_base_id": memory_base_id, + "session_id": session_id, + "workflow_job_id": job_id, + "ingestion_job_id": None, + "recorded_at": datetime.now(timezone.utc), + } + conn = await db.connection() + if conn.dialect.name == "postgresql": + from sqlalchemy.dialects.postgresql import insert as pg_insert + + stmt = pg_insert(MemoryBaseWorkflowRun).values([row]).on_conflict_do_nothing() + else: + from sqlalchemy.dialects.sqlite import insert as sqlite_insert + + stmt = sqlite_insert(MemoryBaseWorkflowRun).values([row]).on_conflict_do_nothing() + await db.exec(stmt) # type: ignore[call-overload] + await db.commit() + + async def _count_pending(self, db: AsyncSession, mb: MemoryBase, mbs: MemoryBaseSession) -> int: + """Count WORKFLOW runs for this (memory_base, session) not yet covered by a completed ingestion. + + A row in memory_base_workflow_run with ingestion_job_id IS NULL means the run + has not been processed by any ingestion job. Count pending = number of such rows. + This is session-scoped and time-independent; job failures leave rows NULL so they + are correctly re-counted on the next threshold check. + """ + stmt = ( + select(func.count()) + .select_from(MemoryBaseWorkflowRun) + .where(MemoryBaseWorkflowRun.memory_base_id == mb.id) + .where(MemoryBaseWorkflowRun.session_id == mbs.session_id) + .where(MemoryBaseWorkflowRun.ingestion_job_id == None) # noqa: E711 + ) + try: + result = await db.exec(stmt) + row = result.first() + if row is None: + return 0 + return int(row) + except (TypeError, ValueError, OSError) as e: + await logger.aerror("Error counting pending workflow runs: %s", e) + return 0 + + async def _get_latest_pending_workflow_job_id( + self, db: AsyncSession, mb: MemoryBase, mbs: MemoryBaseSession + ) -> uuid.UUID | None: + """Return the workflow_job_id of the most recent uncovered workflow run for this session. + + Used to build a stable dedupe key for the ingestion job so that the same + batch of runs cannot trigger duplicate ingestion jobs. + """ + stmt = ( + select(MemoryBaseWorkflowRun.workflow_job_id) + .where(MemoryBaseWorkflowRun.memory_base_id == mb.id) + .where(MemoryBaseWorkflowRun.session_id == mbs.session_id) + .where(MemoryBaseWorkflowRun.ingestion_job_id == None) # noqa: E711 + .order_by(col(MemoryBaseWorkflowRun.recorded_at).desc()) + .limit(1) + ) + result = await db.exec(stmt) + return result.first() + + async def _cancel_active_jobs(self, *, memory_base_id: uuid.UUID, db: AsyncSession) -> None: + """Cancel all IN_PROGRESS or QUEUED jobs for this memory base.""" + stmt = ( + select(Job) + .where(Job.asset_id == memory_base_id) + .where(Job.asset_type == "memory_base") + .where(col(Job.status).in_([JobStatus.IN_PROGRESS, JobStatus.QUEUED])) + ) + result = await db.exec(stmt) + active_jobs = list(result.all()) + + task_service = get_task_service() + job_service = get_job_service() + for job in active_jobs: + try: + await task_service.revoke_task(job.job_id) + await job_service.update_job_status(job.job_id, JobStatus.CANCELLED) + await logger.ainfo("Cancelled job %s for memory_base %s", job.job_id, memory_base_id) + except (RuntimeError, ValueError, OSError): + await logger.awarning( + "Could not cancel job %s for memory_base %s", job.job_id, memory_base_id, exc_info=True + ) + + async def _resolve_kb_username(self, db: AsyncSession, user_id: uuid.UUID) -> str: + from langflow.services.database.models.user.model import User + + stmt = select(User.username).where(User.id == user_id) + result = await db.exec(stmt) + username = result.first() + if not username: + msg = f"User {user_id} not found" + raise ValueError(msg) + return username + + async def _resolve_kb_username_by_user_id(self, user_id: uuid.UUID) -> str: + async with session_scope() as db: + return await self._resolve_kb_username(db, user_id) + + def _resolve_embedding(self, kb_name: str, kb_username: str) -> tuple[str, str]: + """Read embedding provider/model from KB metadata.json, with sane defaults.""" + kb_root = KBStorageHelper.get_root_path() + if not kb_root: + return "OpenAI", "text-embedding-3-small" + kb_path: Path = kb_root / kb_username / kb_name + metadata = KBAnalysisHelper.get_metadata(kb_path, fast=True) + provider = metadata.get("embedding_provider") or "OpenAI" + model = metadata.get("embedding_model") or "text-embedding-3-small" + return provider, model diff --git a/src/backend/base/langflow/services/memory_base/task.py b/src/backend/base/langflow/services/memory_base/task.py new file mode 100644 index 0000000000..0b2ed86e43 --- /dev/null +++ b/src/backend/base/langflow/services/memory_base/task.py @@ -0,0 +1,518 @@ +"""Background task for Memory Base ingestion. + +Design principles enforced here: +- Cursor atomicity: cursor_id is NEVER updated before ingestion confirms success. +- Retry safety: If a job fails, cursor_id remains at the last known good position. +- Serialization: A per-(memory_base_id, session_id) lock prevents concurrent jobs from + racing to write the same messages into Chroma. The lock is acquired before any DB or + Chroma access and released in a finally block. +- Live cursor: After acquiring the lock, the current cursor_id is re-read from the DB + (not the dispatch-time snapshot) so the pending message fetch always starts from the + true latest position, even if a prior job advanced the cursor while this job waited. + +The actual Chroma write logic is shared with KB file ingestion via +``KBIngestionHelper.write_documents_to_chroma`` — no duplicate batching/retry code here. +""" + +from __future__ import annotations + +import asyncio +import json +import types +import weakref +from datetime import datetime, timezone +from typing import TYPE_CHECKING + +from langchain_chroma import Chroma +from langchain_core.documents import Document +from langchain_text_splitters import RecursiveCharacterTextSplitter +from lfx.log.logger import logger +from sqlmodel import col, select + +from langflow.api.utils.kb_helpers import KBAnalysisHelper, KBIngestionHelper, KBStorageHelper +from langflow.services.database.models.memory_base.model import MemoryBaseSession, MemoryBaseWorkflowRun +from langflow.services.database.models.message.model import MessageTable +from langflow.services.deps import session_scope + +if TYPE_CHECKING: + import uuid + from pathlib import Path + + from langflow.services.jobs.service import JobService + +# Chunk size for splitting long messages before embedding +_MESSAGE_CHUNK_SIZE = 1000 +_MESSAGE_CHUNK_OVERLAP = 100 + +# Per-(memory_base_id, session_id) lock registry — serializes concurrent ingestion jobs +# for the same session so that two jobs dispatched before either completes cannot race to +# write overlapping messages into Chroma. Pattern follows api/v2/mcp.py:_update_server_locks. +# WeakValueDictionary: entries are GC'd automatically once no coroutine holds a strong +# reference to the lock (i.e. after the task releases it and no other task is waiting). +_session_ingestion_locks: weakref.WeakValueDictionary[tuple, asyncio.Lock] = weakref.WeakValueDictionary() + + +def _get_or_create_session_lock(key: tuple) -> asyncio.Lock: + """Return the asyncio.Lock for the given (memory_base_id, session_id) key. + + Creates a new lock if none exists. The caller must hold a strong reference + to the returned lock for the duration of its use so the WeakValueDictionary + entry is not collected prematurely. The Python GIL guarantees that the + dict read + conditional write is effectively atomic in the async event loop. + """ + lock = _session_ingestion_locks.get(key) + if lock is None: + lock = asyncio.Lock() + _session_ingestion_locks[key] = lock + return lock + + +# How long a job waits to acquire the session lock before timing out. When the timeout +# expires, asyncio.TimeoutError is re-raised so execute_with_status records JobStatus.TIMED_OUT. +_LOCK_WAIT_TIMEOUT_SECS: int = 600 # 10 minutes + + +async def _read_live_cursor(memory_base_id: uuid.UUID, session_id: str) -> uuid.UUID | None: + """Read the current cursor_id from the DB inside the serialization lock. + + Returns the live cursor — not the dispatch-time snapshot — so the pending + message fetch always starts from the true latest position. Returns None if + the session does not exist or no messages have been ingested yet. + """ + async with session_scope() as db: + stmt = ( + select(MemoryBaseSession.cursor_id) + .where(MemoryBaseSession.memory_base_id == memory_base_id) + .where(MemoryBaseSession.session_id == session_id) + ) + result = await db.exec(stmt) + return result.first() + + +async def ingest_memory_task( + *, + memory_base_id: uuid.UUID, + session_id: str, + flow_id: uuid.UUID, + kb_name: str, + kb_username: str, + user_id: uuid.UUID, + embedding_provider: str, + embedding_model: str, + cursor_id: uuid.UUID | None, + task_job_id: uuid.UUID, + job_service: JobService, +) -> dict: + """Ingest pending output messages from a session into the target Knowledge Base. + + Serialization: acquires a per-(memory_base_id, session_id) asyncio.Lock before any + DB or Chroma access. Concurrent jobs for the same session wait up to + _LOCK_WAIT_TIMEOUT_SECS; if the lock cannot be acquired in time, asyncio.TimeoutError + is re-raised so execute_with_status records JobStatus.TIMED_OUT. + + Live cursor: after acquiring the lock, the current cursor_id is re-read from the DB. + ``cursor_id`` (the argument) is the dispatch-time snapshot kept only for logging. + + Note: ``task_job_id`` (not ``job_id``) is used to avoid colliding with the + ``job_id`` kwarg consumed by ``JobService.execute_with_status`` at the call site. + + Args: + memory_base_id: ID of the MemoryBase configuration. + session_id: Conversation/session identifier. + flow_id: The flow whose outputs are being captured. + kb_name: Target Knowledge Base directory name. + kb_username: Username (filesystem path component for the KB). + user_id: Owner UUID — passed to ``_build_embeddings`` for API-key resolution. + embedding_provider: Embedding provider name (e.g. "OpenAI"). + embedding_model: Embedding model identifier. + cursor_id: Dispatch-time snapshot of the last known cursor (for logging only). + The live cursor is re-read from the DB after lock acquisition. + task_job_id: Job ID for cancellation checking. + job_service: Service for checking cancellation. + + Returns: + Dict with ingestion summary. + + Raises: + asyncio.TimeoutError: If the session lock cannot be acquired within the timeout. + execute_with_status catches this and records JobStatus.TIMED_OUT. + Exception: Re-raises any other failure; cursor is NOT advanced on failure. + """ + await logger.adebug( + "Ingestion job started | memory_base=%s session=%s dispatch_cursor=%s job=%s", + memory_base_id, + session_id, + cursor_id, + task_job_id, + ) + kb_root = KBStorageHelper.get_root_path() + if not kb_root: + msg = "Knowledge base root path is not configured" + raise RuntimeError(msg) + + kb_path: Path = kb_root / kb_username / kb_name + + # ---- 0. Acquire per-session serialization lock ---- + # asyncio.TimeoutError is NOT caught here — it propagates to execute_with_status + # which records JobStatus.TIMED_OUT. The lock is never held when TimeoutError fires, + # so there is nothing to release. + lock = _get_or_create_session_lock((memory_base_id, session_id)) + try: + await asyncio.wait_for(lock.acquire(), timeout=_LOCK_WAIT_TIMEOUT_SECS) + except asyncio.TimeoutError: + await logger.awarning( + "Ingestion lock wait timeout | memory_base=%s session=%s job=%s — re-raising for TIMED_OUT status.", + memory_base_id, + session_id, + task_job_id, + ) + raise + + try: + # ---- 0b. Re-read live cursor inside the lock ---- + live_cursor_id = await _read_live_cursor(memory_base_id, session_id) + await logger.adebug( + "Ingestion lock acquired | memory_base=%s session=%s dispatch_cursor=%s live_cursor=%s job=%s", + memory_base_id, + session_id, + cursor_id, + live_cursor_id, + task_job_id, + ) + + # ---- 1. Fetch pending output messages for this session ---- + messages = await _fetch_pending_messages( + flow_id=flow_id, + session_id=session_id, + cursor_id=live_cursor_id, + ) + if not messages: + await logger.ainfo("MemoryBase %s / session %s: no pending messages, skipping.", memory_base_id, session_id) + return {"message": "No pending messages", "ingested": 0} + + # ---- 2. Build documents from messages ---- + documents = _build_documents_from_messages(messages, session_id=session_id, flow_id=str(flow_id)) + + if not documents: + return {"message": "No non-empty messages to ingest", "ingested": 0} + + # ---- 3. Check cancellation before touching the vector store ---- + if await KBIngestionHelper.is_job_cancelled(job_service, task_job_id): + return {"message": "Job cancelled before ingestion", "ingested": 0} + + # ---- 4. Open Chroma, write, then sync KB metadata — all before closing ---- + # user_stub carries only .id so that build_embeddings can resolve API keys. + user_stub = types.SimpleNamespace(id=user_id) + embeddings = await KBIngestionHelper.build_embeddings(embedding_provider, embedding_model, user_stub) + + client = KBStorageHelper.get_fresh_chroma_client(kb_path) + written = 0 + try: + chroma = Chroma(client=client, embedding_function=embeddings, collection_name=kb_name) + + written = await KBIngestionHelper.write_documents_to_chroma( + documents=documents, + chroma=chroma, + task_job_id=task_job_id, + job_service=job_service, + ) + + if written == len(documents): + # Sync embedding_metadata.json while the collection is still open, + # matching the pattern used by perform_ingestion. + _sync_kb_metadata(kb_path=kb_path, chroma=chroma) + finally: + client = None + chroma = None # type: ignore[assignment] + KBStorageHelper.release_chroma_resources(kb_path) + + if written < len(documents): + # Job was cancelled mid-write; cursor must NOT advance. + return {"message": "Job cancelled during ingestion", "ingested": 0} + + # ---- 5. Bulk-stamp ingestion metadata on every ingested message ---- + await _mark_messages_ingested(messages=messages, job_id=task_job_id, memory_base_id=memory_base_id) + + # ---- 6. Update cursor atomically ONLY after confirmed success ---- + last_message_id = messages[-1].id + ingested_count = len(messages) + await _advance_cursor( + memory_base_id=memory_base_id, + session_id=session_id, + new_cursor_id=last_message_id, + ingested_count=ingested_count, + task_job_id=task_job_id, + ) + + await logger.ainfo( + "Ingestion job finished | memory_base=%s session=%s job=%s ingested=%d new_cursor=%s", + memory_base_id, + session_id, + task_job_id, + ingested_count, + last_message_id, + ) + return {"message": "Success", "ingested": ingested_count} + + finally: + lock.release() + + +async def _fetch_pending_messages( + *, + flow_id: uuid.UUID, + session_id: str, + cursor_id: uuid.UUID | None, +) -> list[MessageTable]: + """Fetch all messages for this session that come after cursor_id. + + Ordering is (timestamp ASC, id ASC) — a deterministic, stable sort that + handles same-timestamp ties correctly. The cursor is a compound + (cursor_ts, cursor_id) position: messages are included when + ``timestamp > cursor_ts`` OR ``(timestamp == cursor_ts AND id > cursor_id)``. + + This prevents silent data loss when two messages share the same timestamp: + any message with ts == cursor_ts but a UUID that sorts *after* the cursor id + is correctly included in the next batch. UUID ordering within a tie is + arbitrary but consistent, so the partition is always correct. + + is_output filtering is intentionally omitted so the full conversation + batch (user turns + model turns) is ingested into the KB. + """ + from sqlalchemy import and_, or_ + + async with session_scope() as db: + stmt = ( + select(MessageTable) + .where(MessageTable.flow_id == flow_id) + .where(MessageTable.session_id == session_id) + .order_by(col(MessageTable.timestamp).asc(), col(MessageTable.id).asc()) + ) + if cursor_id is not None: + cursor_stmt = select(MessageTable.timestamp, MessageTable.id).where(MessageTable.id == cursor_id) + result = await db.exec(cursor_stmt) + cursor_row = result.first() + if cursor_row: + cursor_ts, c_id = cursor_row + stmt = stmt.where( + or_( + col(MessageTable.timestamp) > cursor_ts, + and_( + col(MessageTable.timestamp) == cursor_ts, + col(MessageTable.id) > c_id, + ), + ) + ) + + result = await db.exec(stmt) + return list(result.all()) + + +async def _mark_messages_ingested( + *, + messages: list[MessageTable], + job_id: uuid.UUID, + memory_base_id: uuid.UUID, +) -> None: + """Batch-insert ingestion records for all successfully ingested messages. + + Uses dialect-specific INSERT ... ON CONFLICT DO NOTHING for idempotency: + if a job retries after Chroma write succeeds but before cursor advance, + re-inserting the same rows is a safe no-op. + Called only after a confirmed successful Chroma write. + """ + from uuid import uuid4 as _uuid4 + + from langflow.services.database.models.memory_base.model import MessageIngestionRecord + + ingested_at = datetime.now(timezone.utc) + rows = [ + { + "id": _uuid4(), + "message_id": msg.id, + "memory_base_id": memory_base_id, + "job_id": job_id, + "session_id": msg.session_id, + "ingested_at": ingested_at, + } + for msg in messages + ] + async with session_scope() as db: + conn = await db.connection() + if conn.dialect.name == "postgresql": + from sqlalchemy.dialects.postgresql import insert as pg_insert + + stmt = pg_insert(MessageIngestionRecord).values(rows).on_conflict_do_nothing() + else: + from sqlalchemy.dialects.sqlite import insert as sqlite_insert + + stmt = sqlite_insert(MessageIngestionRecord).values(rows).on_conflict_do_nothing() + await db.exec(stmt) # type: ignore[call-overload] + await db.commit() + + +def _extract_content_block_text(content_blocks: list) -> str: + """Extract embeddable text from content blocks of type text, code, and json. + + Blocks of any other type (tool_use, error, media, etc.) are skipped. + Each extracted piece is separated by a blank line so chunk boundaries + remain readable in the vector store. + """ + parts: list[str] = [] + for block in content_blocks: + # content_blocks are stored as JSON; each block is a dict at runtime. + contents: list = block.get("contents", []) if isinstance(block, dict) else [] + for entry in contents: + if not isinstance(entry, dict): + continue + entry_type = entry.get("type") + if entry_type == "text": + fragment = (entry.get("text") or "").strip() + elif entry_type == "code": + lang = entry.get("language") or "" + code = (entry.get("code") or "").strip() + fragment = f"```{lang}\n{code}\n```" if code else "" + elif entry_type == "json": + data = entry.get("data") + fragment = json.dumps(data, ensure_ascii=False) if data is not None else "" + else: + continue + if fragment: + parts.append(fragment) + return "\n\n".join(parts) + + +def _build_documents_from_messages( + messages: list[MessageTable], + *, + session_id: str, + flow_id: str, +) -> list[Document]: + """Convert MessageTable rows into LangChain Documents. + + Each message's embeddable text is the concatenation of msg.text and any + content-block fragments whose type is text, code, or json. Other block + types (tool_use, error, media, …) are ignored. Long combined texts are + split by RecursiveCharacterTextSplitter before embedding. + """ + splitter = RecursiveCharacterTextSplitter( + chunk_size=_MESSAGE_CHUNK_SIZE, + chunk_overlap=_MESSAGE_CHUNK_OVERLAP, + ) + docs: list[Document] = [] + for msg in messages: + parts: list[str] = [] + if msg.text and msg.text.strip(): + parts.append(msg.text.strip()) + cb_text = _extract_content_block_text(msg.content_blocks or []) + if cb_text: + parts.append(cb_text) + + text = "\n\n".join(parts) + if not text: + continue + chunks = splitter.split_text(text) + for i, chunk in enumerate(chunks): + docs.append( + Document( + page_content=chunk, + metadata={ + "message_id": str(msg.id), + "session_id": session_id, + "flow_id": flow_id, + "sender": msg.sender, + "sender_name": msg.sender_name, + "timestamp": msg.timestamp.isoformat() if msg.timestamp else "", + "run_id": str(msg.run_id) if msg.run_id else "", + "chunk_index": i, + "total_chunks": len(chunks), + "source": f"memory_base/{session_id}", + }, + ) + ) + return docs + + +def _sync_kb_metadata(*, kb_path: Path, chroma: Chroma) -> None: + """Update embedding_metadata.json after a successful Memory Base ingestion. + + Mirrors the post-write metadata sync in ``KBIngestionHelper.perform_ingestion``: + - Refreshes chunk / word / character counts from the live Chroma collection. + - Updates on-disk size. + - Stamps ``is_memory_base: true`` (required for Knowledge Retrieval filtering). + - Sets ``source_types: ["memory"]`` to distinguish from file-based KBs. + + Called while the Chroma client is still open so that ``update_text_metrics`` + can query the collection directly without opening a second client. + """ + try: + metadata = KBAnalysisHelper.get_metadata(kb_path, fast=True) + KBAnalysisHelper.update_text_metrics(kb_path, metadata, chroma=chroma) + metadata["size"] = KBStorageHelper.get_directory_size(kb_path) + metadata["is_memory_base"] = True + # Preserve any existing source_types but always include "memory" + existing = set(metadata.get("source_types") or []) + existing.add("memory") + metadata["source_types"] = sorted(existing) + (kb_path / "embedding_metadata.json").write_text(json.dumps(metadata, indent=2)) + except (OSError, json.JSONDecodeError, ValueError): + # Metadata sync is best-effort; a failure here must not block the cursor advance. + import logging + + logging.getLogger(__name__).warning("KB metadata sync failed for %s", kb_path, exc_info=True) + + +async def _advance_cursor( + *, + memory_base_id: uuid.UUID, + session_id: str, + new_cursor_id: uuid.UUID, + ingested_count: int, + task_job_id: uuid.UUID, +) -> None: + """Atomically advance the cursor, update session stats, and stamp workflow run records. + + This is the FINAL step in ``ingest_memory_task``. It must only be called + after ``write_documents_to_chroma`` confirms all documents were successfully + persisted. + + In addition to updating the message cursor and stats on MemoryBaseSession, this + stamps all pending MemoryBaseWorkflowRun rows (ingestion_job_id IS NULL) for this + session with ``task_job_id``. This marks them as accounted-for so they are not + re-counted toward the threshold on the next on_flow_output call. + If ingestion fails and this function is never called, those rows stay NULL and are + correctly re-counted on the next threshold check. + """ + from sqlalchemy import update as sa_update + + async with session_scope() as db: + stmt = ( + select(MemoryBaseSession) + .where(MemoryBaseSession.memory_base_id == memory_base_id) + .where(MemoryBaseSession.session_id == session_id) + ) + result = await db.exec(stmt) + mbs = result.first() + if mbs is None: + await logger.awarning( + "MemoryBaseSession for (%s, %s) vanished before cursor update - skipping.", + memory_base_id, + session_id, + ) + return + + mbs.cursor_id = new_cursor_id + mbs.total_processed += ingested_count + mbs.last_sync_at = datetime.now(timezone.utc) + db.add(mbs) + + # Stamp all pending workflow run rows for this session as covered by this ingestion. + await db.exec( # type: ignore[call-overload] + sa_update(MemoryBaseWorkflowRun) + .where(MemoryBaseWorkflowRun.memory_base_id == memory_base_id) + .where(MemoryBaseWorkflowRun.session_id == session_id) + .where(MemoryBaseWorkflowRun.ingestion_job_id == None) # noqa: E711 + .values(ingestion_job_id=task_job_id) + ) + + await db.commit() diff --git a/src/backend/tests/unit/test_knowledge_bases_api.py b/src/backend/tests/unit/test_knowledge_bases_api.py index 548eb55351..35f6699b3c 100644 --- a/src/backend/tests/unit/test_knowledge_bases_api.py +++ b/src/backend/tests/unit/test_knowledge_bases_api.py @@ -714,7 +714,7 @@ class TestPerformIngestionTask: @patch("langflow.api.utils.kb_helpers.KBStorageHelper.get_fresh_chroma_client") @patch("langflow.api.utils.kb_helpers.Chroma") - @patch("langflow.api.utils.kb_helpers.KBIngestionHelper._build_embeddings", new_callable=AsyncMock) + @patch("langflow.api.utils.kb_helpers.KBIngestionHelper.build_embeddings", new_callable=AsyncMock) @patch("langflow.api.utils.kb_helpers.KBAnalysisHelper.get_metadata") @patch("langflow.api.utils.kb_helpers.KBStorageHelper.get_directory_size") @patch("langflow.api.utils.kb_helpers.KBAnalysisHelper.update_text_metrics") @@ -764,7 +764,7 @@ class TestPerformIngestionTask: @patch("langflow.api.utils.kb_helpers.KBStorageHelper.get_fresh_chroma_client") @patch("langflow.api.utils.kb_helpers.Chroma") - @patch("langflow.api.utils.kb_helpers.KBIngestionHelper._build_embeddings", new_callable=AsyncMock) + @patch("langflow.api.utils.kb_helpers.KBIngestionHelper.build_embeddings", new_callable=AsyncMock) @patch("langflow.api.utils.kb_helpers.KBIngestionHelper.cleanup_chroma_chunks_by_job", new_callable=AsyncMock) async def test_perform_ingestion_rollback( self, mock_cleanup, mock_build, mock_chroma, mock_fresh_client, mock_kb_path diff --git a/src/backend/tests/unit/test_memory_base_task.py b/src/backend/tests/unit/test_memory_base_task.py new file mode 100644 index 0000000000..e08adc7236 --- /dev/null +++ b/src/backend/tests/unit/test_memory_base_task.py @@ -0,0 +1,886 @@ +"""Unit tests for langflow.services.memory_base.task. + +Covers the gaps not addressed by TestIngestMemoryTask in test_memory_bases.py: +- ingest_memory_task: missing kb_root, pre-ingestion cancel, zero-document early-out +- _extract_content_block_text: all block types, edge cases +- _build_documents_from_messages: chunking, content-block text, missing fields +- _sync_kb_metadata: source_types merge +- _advance_cursor: vanished session, normal update +- _mark_messages_ingested: correct DB update shape +""" + +from __future__ import annotations + +import asyncio +import json +import uuid +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from langflow.services.database.models.message.model import MessageTable + +# ------------------------------------------------------------------ # +# Shared helpers # +# ------------------------------------------------------------------ # + + +def _make_message( + *, + flow_id: uuid.UUID | None = None, + session_id: str = "sess-1", + text: str = "hello", + run_id: uuid.UUID | None = None, + timestamp: datetime | None = None, + content_blocks: list | None = None, +) -> MessageTable: + return MessageTable( + id=uuid.uuid4(), + sender="AI", + sender_name="Bot", + session_id=session_id, + text=text, + flow_id=flow_id or uuid.uuid4(), + timestamp=timestamp or datetime.now(timezone.utc), + run_id=run_id, + content_blocks=content_blocks or [], + ) + + +def _fake_scope(mock_db): + class _FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + scope = MagicMock() + scope.return_value = _FakeCtx() + return scope + + +# ------------------------------------------------------------------ # +# ingest_memory_task — orchestrator edge cases # +# ------------------------------------------------------------------ # + + +class TestIngestMemoryTaskEdgeCases: + @pytest.mark.asyncio + async def test_raises_when_kb_root_not_configured(self): + from langflow.services.memory_base.task import ingest_memory_task + + with ( + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=None, + ), + pytest.raises(RuntimeError, match="Knowledge base root path is not configured"), + ): + await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=uuid.uuid4(), + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + ) + + @pytest.mark.asyncio + async def test_returns_early_when_job_cancelled_before_write(self, tmp_path): + """is_job_cancelled=True after fetch must return without touching Chroma.""" + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id) + chroma_client_called = False + + def fake_get_client(_path): + nonlocal chroma_client_called + chroma_client_called = True + return MagicMock() + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task._build_documents_from_messages", + return_value=[MagicMock()], + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.is_job_cancelled", + AsyncMock(return_value=True), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_fresh_chroma_client", + side_effect=fake_get_client, + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + ): + result = await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + ) + + assert result == {"message": "Job cancelled before ingestion", "ingested": 0} + assert not chroma_client_called + + @pytest.mark.asyncio + async def test_returns_early_when_documents_list_is_empty(self, tmp_path): + """All-whitespace messages produce zero documents — early exit before Chroma.""" + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id, text=" ") # whitespace only + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + ): + result = await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + ) + + assert result == {"message": "No non-empty messages to ingest", "ingested": 0} + + @pytest.mark.asyncio + async def test_mark_messages_ingested_called_on_success(self, tmp_path): + """_mark_messages_ingested must be called exactly once on a successful run.""" + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id) + + mark_ingested_mock = AsyncMock() + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task._build_documents_from_messages", + return_value=[MagicMock()], + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.is_job_cancelled", + AsyncMock(return_value=False), + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.build_embeddings", + AsyncMock(return_value=MagicMock()), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_fresh_chroma_client", + return_value=MagicMock(), + ), + patch("langflow.services.memory_base.task.Chroma"), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.write_documents_to_chroma", + AsyncMock(return_value=1), + ), + patch("langflow.services.memory_base.task._sync_kb_metadata"), + patch("langflow.services.memory_base.task._mark_messages_ingested", mark_ingested_mock), + patch("langflow.services.memory_base.task._advance_cursor", AsyncMock()), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + patch("langflow.services.memory_base.task.KBStorageHelper.release_chroma_resources"), + ): + await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + ) + + mark_ingested_mock.assert_awaited_once() + + @pytest.mark.asyncio + async def test_mark_messages_ingested_not_called_when_cancelled(self, tmp_path): + """When write returns fewer docs than sent, messages must NOT be stamped.""" + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id) + + mark_ingested_mock = AsyncMock() + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task._build_documents_from_messages", + return_value=[MagicMock()], + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.is_job_cancelled", + AsyncMock(return_value=False), + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.build_embeddings", + AsyncMock(return_value=MagicMock()), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_fresh_chroma_client", + return_value=MagicMock(), + ), + patch("langflow.services.memory_base.task.Chroma"), + # Partial write simulates mid-run cancellation + patch( + "langflow.services.memory_base.task.KBIngestionHelper.write_documents_to_chroma", + AsyncMock(return_value=0), + ), + patch("langflow.services.memory_base.task._sync_kb_metadata"), + patch("langflow.services.memory_base.task._mark_messages_ingested", mark_ingested_mock), + patch("langflow.services.memory_base.task._advance_cursor", AsyncMock()), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + patch("langflow.services.memory_base.task.KBStorageHelper.release_chroma_resources"), + ): + result = await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + ) + + mark_ingested_mock.assert_not_awaited() + assert "cancelled" in result["message"].lower() + + +# ------------------------------------------------------------------ # +# _extract_content_block_text # +# ------------------------------------------------------------------ # + + +class TestExtractContentBlockText: + def _call(self, blocks): + from langflow.services.memory_base.task import _extract_content_block_text + + return _extract_content_block_text(blocks) + + def test_empty_list_returns_empty_string(self): + assert self._call([]) == "" + + def test_text_block_extracted(self): + blocks = [{"contents": [{"type": "text", "text": "hello world"}]}] + result = self._call(blocks) + assert result == "hello world" + + def test_text_block_whitespace_only_skipped(self): + blocks = [{"contents": [{"type": "text", "text": " "}]}] + result = self._call(blocks) + assert result == "" + + def test_code_block_with_language(self): + blocks = [{"contents": [{"type": "code", "language": "python", "code": "print('hi')"}]}] + result = self._call(blocks) + assert result == "```python\nprint('hi')\n```" + + def test_code_block_without_language(self): + blocks = [{"contents": [{"type": "code", "language": "", "code": "x = 1"}]}] + result = self._call(blocks) + assert result == "```\nx = 1\n```" + + def test_code_block_empty_code_skipped(self): + blocks = [{"contents": [{"type": "code", "language": "python", "code": ""}]}] + result = self._call(blocks) + assert result == "" + + def test_json_block_serialized(self): + data = {"key": "value", "num": 42} + blocks = [{"contents": [{"type": "json", "data": data}]}] + result = self._call(blocks) + assert result == json.dumps(data, ensure_ascii=False) + + def test_json_block_none_data_skipped(self): + blocks = [{"contents": [{"type": "json", "data": None}]}] + result = self._call(blocks) + assert result == "" + + def test_unknown_block_type_skipped(self): + blocks = [ + { + "contents": [ + {"type": "tool_use", "tool": "search"}, + {"type": "error", "message": "failed"}, + {"type": "media", "url": "http://x.com/img.png"}, + {"type": "text", "text": "kept"}, + ] + } + ] + result = self._call(blocks) + assert result == "kept" + + def test_non_dict_entry_skipped(self): + # entries that aren't dicts should be silently skipped + blocks = [{"contents": ["not a dict", 42, None, {"type": "text", "text": "ok"}]}] + result = self._call(blocks) + assert result == "ok" + + def test_non_dict_block_skipped(self): + # top-level block that isn't a dict + blocks = ["string block", {"contents": [{"type": "text", "text": "valid"}]}] + result = self._call(blocks) + assert result == "valid" + + def test_multiple_blocks_joined_with_double_newline(self): + blocks = [ + {"contents": [{"type": "text", "text": "first"}]}, + {"contents": [{"type": "text", "text": "second"}]}, + ] + result = self._call(blocks) + assert result == "first\n\nsecond" + + def test_multiple_entries_in_same_block_joined(self): + blocks = [ + { + "contents": [ + {"type": "text", "text": "a"}, + {"type": "text", "text": "b"}, + ] + } + ] + result = self._call(blocks) + assert result == "a\n\nb" + + +# ------------------------------------------------------------------ # +# _build_documents_from_messages # +# ------------------------------------------------------------------ # + + +class TestBuildDocumentsFromMessages: + def _call(self, messages, *, session_id="s1", flow_id=None): + from langflow.services.memory_base.task import _build_documents_from_messages + + return _build_documents_from_messages( + messages, + session_id=session_id, + flow_id=flow_id or str(uuid.uuid4()), + ) + + def test_content_blocks_contribute_to_doc_text(self): + flow_id = uuid.uuid4() + msg = _make_message( + flow_id=flow_id, + text="", + content_blocks=[{"contents": [{"type": "text", "text": "from block"}]}], + ) + docs = self._call([msg], flow_id=str(flow_id)) + assert len(docs) == 1 + assert "from block" in docs[0].page_content + + def test_text_and_content_blocks_combined(self): + flow_id = uuid.uuid4() + msg = _make_message( + flow_id=flow_id, + text="msg text", + content_blocks=[{"contents": [{"type": "text", "text": "block text"}]}], + ) + docs = self._call([msg], flow_id=str(flow_id)) + assert len(docs) == 1 + assert "msg text" in docs[0].page_content + assert "block text" in docs[0].page_content + + def test_long_message_split_into_multiple_chunks(self): + from langflow.services.memory_base.task import _MESSAGE_CHUNK_SIZE + + flow_id = uuid.uuid4() + # Craft text longer than chunk size + long_text = "x " * (_MESSAGE_CHUNK_SIZE + 100) + msg = _make_message(flow_id=flow_id, text=long_text) + docs = self._call([msg], flow_id=str(flow_id)) + assert len(docs) > 1 + + def test_chunk_index_and_total_chunks_metadata(self): + from langflow.services.memory_base.task import _MESSAGE_CHUNK_SIZE + + flow_id = uuid.uuid4() + long_text = "word " * (_MESSAGE_CHUNK_SIZE // 4) + msg = _make_message(flow_id=flow_id, text=long_text) + docs = self._call([msg], flow_id=str(flow_id)) + for i, doc in enumerate(docs): + assert doc.metadata["chunk_index"] == i + assert doc.metadata["total_chunks"] == len(docs) + + def test_missing_run_id_stored_as_empty_string(self): + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id, run_id=None) + docs = self._call([msg], flow_id=str(flow_id)) + assert docs[0].metadata["run_id"] == "" + + def test_run_id_stored_as_string(self): + flow_id = uuid.uuid4() + run_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id, run_id=run_id) + docs = self._call([msg], flow_id=str(flow_id)) + assert docs[0].metadata["run_id"] == str(run_id) + + def test_missing_timestamp_stored_as_empty_string(self): + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id) + # validate_assignment=True on the model prevents setting timestamp=None + # directly, so bypass Pydantic validation via object.__setattr__. + object.__setattr__(msg, "timestamp", None) + docs = self._call([msg], flow_id=str(flow_id)) + assert docs[0].metadata["timestamp"] == "" + + def test_source_metadata_uses_session_id(self): + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id, session_id="my-session") + docs = self._call([msg], session_id="my-session", flow_id=str(flow_id)) + assert docs[0].metadata["source"] == "memory_base/my-session" + + def test_multiple_messages_produce_separate_docs(self): + flow_id = uuid.uuid4() + msgs = [_make_message(flow_id=flow_id, text=f"msg {i}") for i in range(3)] + docs = self._call(msgs, flow_id=str(flow_id)) + assert len(docs) == 3 + message_ids = [d.metadata["message_id"] for d in docs] + assert len(set(message_ids)) == 3 + + +# ------------------------------------------------------------------ # +# _sync_kb_metadata # +# ------------------------------------------------------------------ # + + +class TestSyncKbMetadata: + def test_preserves_existing_source_types(self, tmp_path): + from langflow.services.memory_base.task import _sync_kb_metadata + + kb_path = tmp_path / "kb" + kb_path.mkdir() + + with ( + patch( + "langflow.services.memory_base.task.KBAnalysisHelper.get_metadata", + return_value={"chunks": 5, "source_types": ["file"]}, + ), + patch("langflow.services.memory_base.task.KBAnalysisHelper.update_text_metrics"), + patch("langflow.services.memory_base.task.KBStorageHelper.get_directory_size", return_value=2048), + ): + _sync_kb_metadata(kb_path=kb_path, chroma=MagicMock()) + + written = json.loads((kb_path / "embedding_metadata.json").read_text()) + assert "file" in written["source_types"] + assert "memory" in written["source_types"] + + def test_source_types_sorted(self, tmp_path): + from langflow.services.memory_base.task import _sync_kb_metadata + + kb_path = tmp_path / "kb" + kb_path.mkdir() + + with ( + patch( + "langflow.services.memory_base.task.KBAnalysisHelper.get_metadata", + return_value={"chunks": 0, "source_types": ["zzz", "aaa"]}, + ), + patch("langflow.services.memory_base.task.KBAnalysisHelper.update_text_metrics"), + patch("langflow.services.memory_base.task.KBStorageHelper.get_directory_size", return_value=0), + ): + _sync_kb_metadata(kb_path=kb_path, chroma=MagicMock()) + + written = json.loads((kb_path / "embedding_metadata.json").read_text()) + assert written["source_types"] == sorted(written["source_types"]) + + def test_json_decode_error_swallowed(self, tmp_path): + from langflow.services.memory_base.task import _sync_kb_metadata + + kb_path = tmp_path / "kb" + kb_path.mkdir() + + with patch( + "langflow.services.memory_base.task.KBAnalysisHelper.get_metadata", + side_effect=json.JSONDecodeError("bad", "", 0), + ): + # Must not raise + _sync_kb_metadata(kb_path=kb_path, chroma=MagicMock()) + + def test_value_error_swallowed(self, tmp_path): + from langflow.services.memory_base.task import _sync_kb_metadata + + kb_path = tmp_path / "kb" + kb_path.mkdir() + + with ( + patch( + "langflow.services.memory_base.task.KBAnalysisHelper.get_metadata", + return_value={}, + ), + patch( + "langflow.services.memory_base.task.KBAnalysisHelper.update_text_metrics", + side_effect=ValueError("bad metric"), + ), + ): + _sync_kb_metadata(kb_path=kb_path, chroma=MagicMock()) + + +# ------------------------------------------------------------------ # +# _advance_cursor # +# ------------------------------------------------------------------ # + + +class TestAdvanceCursor: + @pytest.mark.asyncio + async def test_normal_update(self): + from langflow.services.database.models.memory_base.model import MemoryBaseSession + from langflow.services.memory_base.task import _advance_cursor + + mb_id = uuid.uuid4() + new_cursor = uuid.uuid4() + task_job_id = uuid.uuid4() + + mbs = MemoryBaseSession( + id=uuid.uuid4(), + memory_base_id=mb_id, + session_id="s1", + cursor_id=None, + total_processed=5, + ) + + mock_db = AsyncMock() + mock_select_result = MagicMock() + mock_select_result.first = MagicMock(return_value=mbs) + # First exec = SELECT, second exec = UPDATE MemoryBaseWorkflowRun + mock_db.exec = AsyncMock(side_effect=[mock_select_result, MagicMock()]) + + with patch("langflow.services.memory_base.task.session_scope", _fake_scope(mock_db)): + await _advance_cursor( + memory_base_id=mb_id, + session_id="s1", + new_cursor_id=new_cursor, + ingested_count=3, + task_job_id=task_job_id, + ) + + assert mbs.cursor_id == new_cursor + assert mbs.total_processed == 8 # 5 + 3 + assert mbs.last_sync_at is not None + mock_db.add.assert_called_once_with(mbs) + mock_db.commit.assert_awaited_once() + + @pytest.mark.asyncio + async def test_vanished_session_does_not_raise(self): + """If MemoryBaseSession is gone, _advance_cursor must log a warning and return.""" + from langflow.services.memory_base.task import _advance_cursor + + mock_db = AsyncMock() + mock_result = MagicMock() + mock_result.first = MagicMock(return_value=None) # session vanished + mock_db.exec = AsyncMock(return_value=mock_result) + + with patch("langflow.services.memory_base.task.session_scope", _fake_scope(mock_db)): + # Must not raise + await _advance_cursor( + memory_base_id=uuid.uuid4(), + session_id="gone", + new_cursor_id=uuid.uuid4(), + ingested_count=1, + task_job_id=uuid.uuid4(), + ) + + mock_db.add.assert_not_called() + mock_db.commit.assert_not_awaited() + + +# ------------------------------------------------------------------ # +# _mark_messages_ingested # +# ------------------------------------------------------------------ # + + +class TestMarkMessagesIngested: + @pytest.mark.asyncio + async def test_executes_bulk_update(self): + from langflow.services.memory_base.task import _mark_messages_ingested + + flow_id = uuid.uuid4() + messages = [_make_message(flow_id=flow_id) for _ in range(3)] + job_id = uuid.uuid4() + memory_base_id = uuid.uuid4() + + mock_conn = MagicMock() + mock_conn.dialect.name = "sqlite" + + mock_db = AsyncMock() + mock_db.exec = AsyncMock() + mock_db.connection = AsyncMock(return_value=mock_conn) + + with patch("langflow.services.memory_base.task.session_scope", _fake_scope(mock_db)): + await _mark_messages_ingested(messages=messages, job_id=job_id, memory_base_id=memory_base_id) + + mock_db.exec.assert_awaited_once() + mock_db.commit.assert_awaited_once() + + @pytest.mark.asyncio + async def test_update_sets_ingestion_job_id_and_timestamp(self): + """The INSERT statement must include job_id and ingested_at for each message.""" + from langflow.services.memory_base.task import _mark_messages_ingested + + flow_id = uuid.uuid4() + messages = [_make_message(flow_id=flow_id)] + job_id = uuid.uuid4() + memory_base_id = uuid.uuid4() + + captured_stmt = {} + + mock_conn = MagicMock() + mock_conn.dialect.name = "sqlite" + + async def capture_exec(stmt): + captured_stmt["stmt"] = stmt + return MagicMock() + + mock_db = AsyncMock() + mock_db.exec = capture_exec + mock_db.connection = AsyncMock(return_value=mock_conn) + + with patch("langflow.services.memory_base.task.session_scope", _fake_scope(mock_db)): + await _mark_messages_ingested(messages=messages, job_id=job_id, memory_base_id=memory_base_id) + + assert "stmt" in captured_stmt, "db.exec was not called" + stmt = captured_stmt["stmt"] + # The INSERT statement must reference job_id and ingested_at columns + stmt_str = str(stmt.compile()) + assert "job_id" in stmt_str + assert "ingested_at" in stmt_str + + +# ------------------------------------------------------------------ # +# TestIngestionLocking — serialization and cursor re-read # +# ------------------------------------------------------------------ # + + +class TestIngestionLocking: + """Tests for the per-session lock, live cursor re-read, and graceful early-exits.""" + + _BASE_KWARGS: dict = { + "session_id": "s1", + "kb_name": "kb", + "kb_username": "user", + "embedding_provider": "OpenAI", + "embedding_model": "text-embedding-3-small", + } + + @pytest.mark.asyncio + async def test_live_cursor_used_not_dispatch_snapshot(self, tmp_path): + """_fetch_pending_messages must receive the live cursor, not the dispatch-time one.""" + import langflow.services.memory_base.task as task_module + + memory_base_id = uuid.uuid4() + flow_id = uuid.uuid4() + dispatch_cursor = uuid.uuid4() # what was captured at dispatch time + live_cursor = uuid.uuid4() # what the DB currently says + + fetch_calls: list = [] + + async def _recording_fetch(*, flow_id, session_id, cursor_id): # noqa: ARG001 + fetch_calls.append(cursor_id) + return [] # empty → early exit; no Chroma setup needed + + task_module._session_ingestion_locks.pop((memory_base_id, "s1"), None) + + with ( + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + patch( + "langflow.services.memory_base.task._read_live_cursor", + AsyncMock(return_value=live_cursor), + ), + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + side_effect=_recording_fetch, + ), + ): + result = await task_module.ingest_memory_task( + memory_base_id=memory_base_id, + flow_id=flow_id, + user_id=uuid.uuid4(), + cursor_id=dispatch_cursor, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + **self._BASE_KWARGS, + ) + + assert result == {"message": "No pending messages", "ingested": 0} + assert len(fetch_calls) == 1 + assert fetch_calls[0] == live_cursor, ( + f"Expected fetch called with live_cursor={live_cursor!r}, got {fetch_calls[0]!r}" + ) + + @pytest.mark.asyncio + async def test_lock_released_on_task_exception(self, tmp_path): + """Lock must be released via finally even when the task raises inside the lock body.""" + import langflow.services.memory_base.task as task_module + + memory_base_id = uuid.uuid4() + + task_module._session_ingestion_locks.pop((memory_base_id, "s1"), None) + + with ( + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + patch( + "langflow.services.memory_base.task._read_live_cursor", + AsyncMock(return_value=None), + ), + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(side_effect=RuntimeError("DB exploded inside lock")), + ), + pytest.raises(RuntimeError, match="DB exploded inside lock"), + ): + await task_module.ingest_memory_task( + memory_base_id=memory_base_id, + flow_id=uuid.uuid4(), + user_id=uuid.uuid4(), + cursor_id=None, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + **self._BASE_KWARGS, + ) + + lock = task_module._session_ingestion_locks[(memory_base_id, "s1")] + assert not lock.locked(), "Lock must be released after exception (finally block must have run)" + + @pytest.mark.asyncio + async def test_lock_timeout_raises_asyncio_timeout_error(self, tmp_path): + """When the lock cannot be acquired within the timeout, asyncio.TimeoutError is raised. + + This allows execute_with_status to record JobStatus.TIMED_OUT for an accurate audit trail. + """ + import langflow.services.memory_base.task as task_module + + memory_base_id = uuid.uuid4() + + task_module._session_ingestion_locks.pop((memory_base_id, "s1"), None) + # Use the factory so the lock is inserted into the WeakValueDictionary; + # holding blocking_lock as a strong reference keeps it alive there so the + # task finds the same lock object and blocks on acquire. + blocking_lock = task_module._get_or_create_session_lock((memory_base_id, "s1")) + await blocking_lock.acquire() # hold it — task will timeout waiting + + try: + with ( + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + patch.object(task_module, "_LOCK_WAIT_TIMEOUT_SECS", 0.01), + pytest.raises(asyncio.TimeoutError), + ): + await task_module.ingest_memory_task( + memory_base_id=memory_base_id, + flow_id=uuid.uuid4(), + user_id=uuid.uuid4(), + cursor_id=None, + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + **self._BASE_KWARGS, + ) + finally: + blocking_lock.release() + + @pytest.mark.asyncio + async def test_noop_when_cursor_advanced_by_prior_job(self, tmp_path): + """If a prior job already advanced the cursor to msg3, fetch from msg3 finds nothing — graceful exit.""" + import langflow.services.memory_base.task as task_module + + memory_base_id = uuid.uuid4() + msg3_id = uuid.uuid4() + + advance_cursor_mock = AsyncMock() + + task_module._session_ingestion_locks.pop((memory_base_id, "s1"), None) + + with ( + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path, + ), + patch( + "langflow.services.memory_base.task._read_live_cursor", + AsyncMock(return_value=msg3_id), # prior job advanced to msg3 + ), + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[]), # nothing after msg3 + ), + patch( + "langflow.services.memory_base.task._advance_cursor", + advance_cursor_mock, + ), + ): + result = await task_module.ingest_memory_task( + memory_base_id=memory_base_id, + flow_id=uuid.uuid4(), + user_id=uuid.uuid4(), + cursor_id=None, # dispatch-time snapshot before prior job ran + task_job_id=uuid.uuid4(), + job_service=MagicMock(), + **self._BASE_KWARGS, + ) + + assert result == {"message": "No pending messages", "ingested": 0} + advance_cursor_mock.assert_not_awaited() diff --git a/src/backend/tests/unit/test_memory_bases.py b/src/backend/tests/unit/test_memory_bases.py new file mode 100644 index 0000000000..04b9263773 --- /dev/null +++ b/src/backend/tests/unit/test_memory_bases.py @@ -0,0 +1,1545 @@ +"""Unit tests for the MemoryBase feature. + +Coverage areas: +- DB model creation and field defaults +- MemoryBaseService CRUD operations +- Concurrency guard (409 on duplicate active job) +- Cursor atomicity (cursor not advanced on ingestion failure) +- Threshold-change deferral +- FS/VectorDB mismatch detection +- Regenerate: cursor reset + re-trigger +- API endpoint routing (happy path + error paths) +- ingest_memory_task: pending message fetch, document building, cursor advance +""" + +from __future__ import annotations + +import contextlib +import uuid +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from langflow.services.database.models.memory_base.model import ( + MemoryBase, + MemoryBaseCreate, + MemoryBaseSession, + MemoryBaseUpdate, +) +from langflow.services.database.models.message.model import MessageTable + +# ------------------------------------------------------------------ # +# Helpers # +# ------------------------------------------------------------------ # + + +def _make_mb( + *, + user_id: uuid.UUID | None = None, + flow_id: uuid.UUID | None = None, + threshold: int = 10, + auto_capture: bool = True, +) -> MemoryBase: + return MemoryBase( + id=uuid.uuid4(), + name="test_mb", + flow_id=flow_id or uuid.uuid4(), + user_id=user_id or uuid.uuid4(), + threshold=threshold, + kb_name="test_kb", + auto_capture=auto_capture, + created_at=datetime.now(timezone.utc), + ) + + +def _make_session( + *, + memory_base_id: uuid.UUID | None = None, + session_id: str = "sess-1", + cursor_id: uuid.UUID | None = None, + total_processed: int = 0, +) -> MemoryBaseSession: + return MemoryBaseSession( + id=uuid.uuid4(), + memory_base_id=memory_base_id or uuid.uuid4(), + session_id=session_id, + cursor_id=cursor_id, + total_processed=total_processed, + ) + + +def _make_message( + *, + flow_id: uuid.UUID, + session_id: str, + is_output: bool = True, + text: str = "Hello from the bot", + run_id: uuid.UUID | None = None, +) -> MessageTable: + return MessageTable( + id=uuid.uuid4(), + sender="AI", + sender_name="Bot", + session_id=session_id, + text=text, + flow_id=flow_id, + is_output=is_output, + run_id=run_id, + timestamp=datetime.now(timezone.utc), + ) + + +# ------------------------------------------------------------------ # +# Model tests # +# ------------------------------------------------------------------ # + + +class TestMemoryBaseModel: + def test_defaults(self): + mb = MemoryBase( + name="mb", + flow_id=uuid.uuid4(), + user_id=uuid.uuid4(), + kb_name="kb", + ) + assert mb.threshold == 50 + assert mb.auto_capture is True + + def test_create_schema(self): + payload = MemoryBaseCreate( + name="mb", + flow_id=uuid.uuid4(), + user_id=uuid.uuid4(), + threshold=25, + kb_name="kb", + ) + assert payload.threshold == 25 + + def test_update_schema_partial(self): + patch = MemoryBaseUpdate(threshold=100) + dumped = patch.model_dump(exclude_unset=True) + assert "threshold" in dumped + assert "name" not in dumped + + def test_memory_base_session_defaults(self): + mbs = MemoryBaseSession( + memory_base_id=uuid.uuid4(), + session_id="s1", + ) + assert mbs.cursor_id is None + assert mbs.total_processed == 0 + assert mbs.last_sync_at is None + + +class TestMessageExtensions: + """Ensure the new fields exist on MessageTable.""" + + def test_run_id_field_exists(self): + msg = _make_message(flow_id=uuid.uuid4(), session_id="s1") + assert hasattr(msg, "run_id") + assert msg.run_id is None + + def test_is_output_field_defaults_false(self): + msg = MessageTable( + sender="Human", + sender_name="User", + session_id="s1", + text="hi", + ) + assert msg.is_output is False + + def test_is_output_can_be_set(self): + msg = _make_message(flow_id=uuid.uuid4(), session_id="s1", is_output=True) + assert msg.is_output is True + + +# ------------------------------------------------------------------ # +# Service tests (mock DB) # +# ------------------------------------------------------------------ # + + +class TestMemoryBaseServiceCRUD: + @pytest.fixture + def service(self): + from langflow.services.memory_base.service import MemoryBaseService + + return MemoryBaseService() + + @pytest.mark.asyncio + async def test_create_stores_user_id(self, service): + user_id = uuid.uuid4() + payload = MemoryBaseCreate( + name="mb", + flow_id=uuid.uuid4(), + user_id=user_id, + kb_name="kb", + ) + + created_mb = _make_mb(user_id=user_id) + + with patch.object(service, "create", AsyncMock(return_value=created_mb)): + result = await service.create(payload, user_id=user_id) + + assert result.user_id == user_id + + @pytest.mark.asyncio + async def test_get_returns_none_for_wrong_user(self, service): + with patch.object(service, "get", AsyncMock(return_value=None)): + result = await service.get(uuid.uuid4(), user_id=uuid.uuid4()) + assert result is None + + @pytest.mark.asyncio + async def test_update_returns_none_for_missing(self, service): + with patch.object(service, "update", AsyncMock(return_value=None)): + result = await service.update(uuid.uuid4(), uuid.uuid4(), MemoryBaseUpdate(threshold=5)) + assert result is None + + @pytest.mark.asyncio + async def test_delete_returns_false_for_missing(self, service): + with patch.object(service, "delete", AsyncMock(return_value=False)): + result = await service.delete(uuid.uuid4(), user_id=uuid.uuid4()) + assert result is False + + +class TestMemoryBaseCreateFlowOwnership: + """Regression tests for cross-user data leak via flow_id at creation time. + + A user must not be able to create a Memory Base pointed at another user's + flow. Without the ownership check, on_flow_output() would capture that + flow's conversation history into the attacker's Memory Base, which they + could then read via /sessions + /messages. + """ + + @pytest.fixture + def service(self): + from langflow.services.memory_base.service import MemoryBaseService + + return MemoryBaseService() + + def _fake_scope(self, mock_db): + class _FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + scope = MagicMock() + scope.return_value = _FakeCtx() + return scope + + @pytest.mark.asyncio + async def test_create_rejects_unowned_flow(self, service): + """PermissionError raised when flow_id belongs to a different user.""" + user_id = uuid.uuid4() + payload = MemoryBaseCreate(name="mb", flow_id=uuid.uuid4()) + + mock_db = AsyncMock() + # exec() is async, but .first() on the result is synchronous — use MagicMock + # so that flow_result.first() returns None without returning a coroutine. + exec_result = MagicMock() + exec_result.first.return_value = None + mock_db.exec = AsyncMock(return_value=exec_result) + + with ( + patch("langflow.services.memory_base.service.session_scope", self._fake_scope(mock_db)), + pytest.raises(PermissionError, match="not found"), + ): + await service.create(payload, user_id=user_id) + + @pytest.mark.asyncio + async def test_create_allows_owned_flow(self, service): + """No PermissionError when flow_id is owned by the requesting user.""" + from langflow.services.database.models.flow.model import Flow + + user_id = uuid.uuid4() + flow_id = uuid.uuid4() + payload = MemoryBaseCreate(name="mb", flow_id=flow_id) + + owned_flow = Flow(id=flow_id, user_id=user_id, name="my flow") + created_mb = _make_mb(user_id=user_id, flow_id=flow_id) + + # exec() is async; .first() on the result is synchronous — use MagicMock. + # First call: flow ownership check → returns owned_flow (passes). + # Second call (different session_scope): name-uniqueness check → None. + first_exec = MagicMock() + first_exec.first.return_value = owned_flow + second_exec = MagicMock() + second_exec.first.return_value = None + + mock_db = AsyncMock() + mock_db.exec = AsyncMock(side_effect=[first_exec, second_exec]) + mock_db.refresh = AsyncMock(return_value=created_mb) + + with ( + patch("langflow.services.memory_base.service.session_scope", self._fake_scope(mock_db)), + patch.object(service, "_resolve_kb_username", AsyncMock(return_value="testuser")), + patch.object(service, "_initialize_kb", AsyncMock()), + contextlib.suppress(Exception), + ): + # Should not raise PermissionError — ownership check passes + await service.create(payload, user_id=user_id) + + # The flow ownership query must have been executed + mock_db.exec.assert_called() + + @pytest.mark.asyncio + async def test_create_endpoint_returns_404_for_unowned_flow(self): + """POST /memories returns 404 (not 403/409) when flow_id is unowned. + + 404 is intentional: returning 403 would reveal that the flow exists, + which is an information leak in itself. + Tests the try/except mapping in the route handler directly. + """ + from fastapi import HTTPException + from langflow.api.v1.memories import create_memory_base + from langflow.services.database.models.user.model import User + + fake_user = User(id=uuid.uuid4(), username="alice") + mock_service = MagicMock() + mock_service.create = AsyncMock(side_effect=PermissionError("Flow abc not found")) + + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=mock_service), + pytest.raises(HTTPException) as exc_info, + ): + await create_memory_base( + current_user=fake_user, + payload=MemoryBaseCreate(name="mb", flow_id=uuid.uuid4()), + ) + + assert exc_info.value.status_code == 404 + + +class TestMemoryBaseServiceConcurrency: + """409 guard: only one active ingestion per (memory_base_id, session_id).""" + + @pytest.fixture + def service(self): + from langflow.services.memory_base.service import MemoryBaseService + + return MemoryBaseService() + + def _fake_scope(self, mock_db): + class _FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + scope = MagicMock() + scope.return_value = _FakeCtx() + return scope + + @pytest.mark.asyncio + async def test_trigger_raises_when_job_active(self, service): + """DuplicateJobError from create_job propagates out of trigger_ingestion.""" + from langflow.services.jobs import DuplicateJobError + + mb = _make_mb() + mbs = _make_session(memory_base_id=mb.id) + mock_db = AsyncMock() + + mock_job_svc = MagicMock() + mock_job_svc.create_job = AsyncMock(side_effect=DuplicateJobError("already running")) + + with ( + patch("langflow.services.memory_base.service.session_scope", self._fake_scope(mock_db)), + patch.object(service, "_get_mb_or_raise", AsyncMock(return_value=mb)), + patch.object(service, "_get_or_create_session", AsyncMock(return_value=mbs)), + patch.object(service, "_get_latest_pending_workflow_job_id", AsyncMock(return_value=uuid.uuid4())), + patch.object(service, "_resolve_kb_username", AsyncMock(return_value="testuser")), + patch.object(service, "_resolve_embedding", return_value=("OpenAI", "text-embedding-3-small")), + patch("langflow.services.memory_base.service.get_job_service", return_value=mock_job_svc), + pytest.raises(DuplicateJobError), + ): + await service.trigger_ingestion(mb.id, mb.user_id, "sess-1") + + @pytest.mark.asyncio + async def test_trigger_succeeds_when_no_active_job(self, service): + mb = _make_mb() + mbs = _make_session(memory_base_id=mb.id) + mock_db = AsyncMock() + + mock_job_svc = MagicMock() + mock_job_svc.create_job = AsyncMock() + mock_task_svc = MagicMock() + mock_task_svc.fire_and_forget_task = AsyncMock() + + with ( + patch("langflow.services.memory_base.service.session_scope", self._fake_scope(mock_db)), + patch.object(service, "_get_mb_or_raise", AsyncMock(return_value=mb)), + patch.object(service, "_get_or_create_session", AsyncMock(return_value=mbs)), + patch.object(service, "_get_latest_pending_workflow_job_id", AsyncMock(return_value=uuid.uuid4())), + patch.object(service, "_resolve_kb_username", AsyncMock(return_value="testuser")), + patch.object(service, "_resolve_embedding", return_value=("OpenAI", "text-embedding-3-small")), + patch("langflow.services.memory_base.service.get_job_service", return_value=mock_job_svc), + patch("langflow.services.memory_base.service.get_task_service", return_value=mock_task_svc), + ): + job_id = await service.trigger_ingestion(mb.id, mb.user_id, "sess-1") + + assert isinstance(job_id, str) + mock_job_svc.create_job.assert_awaited_once() + mock_task_svc.fire_and_forget_task.assert_awaited_once() + + +class TestMemoryBaseServiceThreshold: + """Threshold update should NOT immediately re-evaluate pending count.""" + + @pytest.fixture + def service(self): + from langflow.services.memory_base.service import MemoryBaseService + + return MemoryBaseService() + + @pytest.mark.asyncio + async def test_threshold_update_does_not_trigger_ingestion(self, service): + """Updating threshold via PATCH should never fire a task.""" + mb_updated = _make_mb(threshold=5) + + with patch.object(service, "update", AsyncMock(return_value=mb_updated)): + result = await service.update(mb_updated.id, mb_updated.user_id, MemoryBaseUpdate(threshold=5)) + + assert result.threshold == 5 + # No ingestion task should have been triggered as a side effect + + +class TestMemoryBaseServiceMismatch: + @pytest.fixture + def service(self): + from langflow.services.memory_base.service import MemoryBaseService + + return MemoryBaseService() + + @pytest.mark.asyncio + async def test_mismatch_detected_when_processed_but_empty_store(self, service, tmp_path): + mb = _make_mb() + + with ( + patch.object(service, "_get_mb_or_raise", AsyncMock(return_value=mb)), + patch.object(service, "_resolve_kb_username_by_user_id", AsyncMock(return_value="testuser")), + patch("langflow.services.memory_base.service.session_scope") as mock_scope, + patch("langflow.services.memory_base.service.KBStorageHelper.get_root_path", return_value=tmp_path), + patch( + "langflow.services.memory_base.service.KBAnalysisHelper.get_metadata", + return_value={"chunks": 0}, + ), + ): + # Simulate session_scope returns total_processed=10 + mock_db = AsyncMock() + mock_db.exec = AsyncMock(return_value=MagicMock(first=MagicMock(return_value=10))) + + class FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + mock_scope.return_value = FakeCtx() + + # Create KB path dir so path.exists() is True + kb_path = tmp_path / "testuser" / mb.kb_name + kb_path.mkdir(parents=True) + + result = await service.check_mismatch(mb.id, mb.user_id) + + assert result is True + + async def test_no_mismatch_when_nothing_processed(self, service): + mb = _make_mb() + + with ( + patch.object(service, "_get_mb_or_raise", AsyncMock(return_value=mb)), + patch("langflow.services.memory_base.service.session_scope") as mock_scope, + ): + mock_db = AsyncMock() + mock_db.exec = AsyncMock(return_value=MagicMock(first=MagicMock(return_value=0))) + + class FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + mock_scope.return_value = FakeCtx() + + result = await service.check_mismatch(mb.id, mb.user_id) + + assert result is False + + +class TestMemoryBaseServiceRegenerate: + @pytest.fixture + def service(self): + from langflow.services.memory_base.service import MemoryBaseService + + return MemoryBaseService() + + @pytest.mark.asyncio + async def test_regenerate_resets_cursors_and_triggers(self, service): + mb = _make_mb() + mbs1 = _make_session(memory_base_id=mb.id, session_id="s1", cursor_id=uuid.uuid4()) + mbs2 = _make_session(memory_base_id=mb.id, session_id="s2", cursor_id=uuid.uuid4()) + + triggered_sessions: list[str] = [] + + async def fake_trigger(_mb_id, _user_id, session_id): + triggered_sessions.append(session_id) + return str(uuid.uuid4()) + + with ( + patch("langflow.services.memory_base.service.session_scope") as mock_scope, + patch.object(service, "trigger_ingestion", side_effect=fake_trigger), + ): + mock_db = AsyncMock() + mock_mb_result = MagicMock() + mock_mb_result.first = MagicMock(return_value=mb) + mock_session_result = MagicMock() + mock_session_result.all = MagicMock(return_value=[mbs1, mbs2]) + mock_db.exec = AsyncMock(side_effect=[mock_mb_result, mock_session_result, MagicMock()]) + mock_db.add = MagicMock() + mock_db.commit = AsyncMock() + + class FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + mock_scope.return_value = FakeCtx() + + job_ids = await service.regenerate(mb.id, mb.user_id) + + assert len(job_ids) == 2 + assert set(triggered_sessions) == {"s1", "s2"} + # Verify cursors were reset + assert mbs1.cursor_id is None + assert mbs2.cursor_id is None + + +# ------------------------------------------------------------------ # +# Task tests # +# ------------------------------------------------------------------ # + + +class TestIngestMemoryTask: + async def test_no_op_when_no_pending_messages(self, tmp_path): + from langflow.services.memory_base.task import ingest_memory_task + + job_service = MagicMock() + job_id = uuid.uuid4() + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[]), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path / "kb", + ), + ): + result = await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=uuid.uuid4(), + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=job_id, + job_service=job_service, + ) + + assert result["ingested"] == 0 + + @pytest.mark.asyncio + async def test_cursor_not_advanced_on_ingestion_failure(self, tmp_path): + """Critical: cursor_id must stay unchanged if ingestion fails.""" + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + mb_id = uuid.uuid4() + old_cursor = uuid.uuid4() + + msg = _make_message(flow_id=flow_id, session_id="s1") + job_service = MagicMock() + job_id = uuid.uuid4() + + advance_cursor_called = False + + async def fake_advance_cursor(**_kwargs): + nonlocal advance_cursor_called + advance_cursor_called = True + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task._build_documents_from_messages", + return_value=[MagicMock()], + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.is_job_cancelled", + AsyncMock(return_value=False), + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.build_embeddings", + AsyncMock(return_value=MagicMock()), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_fresh_chroma_client", + return_value=MagicMock(), + ), + patch("langflow.services.memory_base.task.Chroma"), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.write_documents_to_chroma", + AsyncMock(side_effect=RuntimeError("Chroma exploded")), + ), + patch( + "langflow.services.memory_base.task._advance_cursor", + side_effect=fake_advance_cursor, + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path / "kb", + ), + patch("langflow.services.memory_base.task.KBStorageHelper.release_chroma_resources"), + pytest.raises(RuntimeError, match="Chroma exploded"), + ): + await ingest_memory_task( + memory_base_id=mb_id, + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=old_cursor, + task_job_id=job_id, + job_service=job_service, + ) + + # Cursor must NOT have been advanced + assert not advance_cursor_called, "cursor_id must not advance when ingestion fails" + + @pytest.mark.asyncio + async def test_metadata_synced_on_success(self, tmp_path): + """embedding_metadata.json must be updated after a successful ingestion.""" + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id, session_id="s1") + job_service = MagicMock() + job_id = uuid.uuid4() + + sync_called_with: dict = {} + + def fake_sync_kb_metadata(*, kb_path, chroma): + sync_called_with["kb_path"] = kb_path + sync_called_with["chroma"] = chroma + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task._build_documents_from_messages", + return_value=[MagicMock()], + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.is_job_cancelled", + AsyncMock(return_value=False), + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.build_embeddings", + AsyncMock(return_value=MagicMock()), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_fresh_chroma_client", + return_value=MagicMock(), + ), + patch("langflow.services.memory_base.task.Chroma"), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.write_documents_to_chroma", + AsyncMock(return_value=1), + ), + patch("langflow.services.memory_base.task._sync_kb_metadata", side_effect=fake_sync_kb_metadata), + patch("langflow.services.memory_base.task._mark_messages_ingested", AsyncMock()), + patch("langflow.services.memory_base.task._advance_cursor", AsyncMock()), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path / "kb", + ), + patch("langflow.services.memory_base.task.KBStorageHelper.release_chroma_resources"), + ): + await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=job_id, + job_service=job_service, + ) + + assert "kb_path" in sync_called_with, "_sync_kb_metadata was not called on success" + + @pytest.mark.asyncio + async def test_metadata_not_synced_when_cancelled(self, tmp_path): + """embedding_metadata.json must NOT be updated when ingestion is cancelled.""" + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id, session_id="s1") + job_service = MagicMock() + job_id = uuid.uuid4() + sync_called = False + + def fake_sync(*_args, **_kwargs): + nonlocal sync_called + sync_called = True + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task._build_documents_from_messages", + return_value=[MagicMock()], + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.is_job_cancelled", + AsyncMock(return_value=False), + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.build_embeddings", + AsyncMock(return_value=MagicMock()), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_fresh_chroma_client", + return_value=MagicMock(), + ), + patch("langflow.services.memory_base.task.Chroma"), + # write_documents_to_chroma returns fewer docs than sent → cancelled + patch( + "langflow.services.memory_base.task.KBIngestionHelper.write_documents_to_chroma", + AsyncMock(return_value=0), + ), + patch("langflow.services.memory_base.task._sync_kb_metadata", side_effect=fake_sync), + patch("langflow.services.memory_base.task._advance_cursor", AsyncMock()), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path / "kb", + ), + patch("langflow.services.memory_base.task.KBStorageHelper.release_chroma_resources"), + ): + result = await ingest_memory_task( + memory_base_id=uuid.uuid4(), + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=job_id, + job_service=job_service, + ) + + assert not sync_called, "_sync_kb_metadata must not be called when ingestion is cancelled" + assert "cancelled" in result["message"].lower() + + @pytest.mark.asyncio + async def test_cursor_advanced_on_success(self, tmp_path): + from langflow.services.memory_base.task import ingest_memory_task + + flow_id = uuid.uuid4() + mb_id = uuid.uuid4() + + msg = _make_message(flow_id=flow_id, session_id="s1") + job_service = MagicMock() + job_id = uuid.uuid4() + + advance_kwargs: dict = {} + + async def fake_advance_cursor(**kwargs): + advance_kwargs.update(kwargs) + + with ( + patch( + "langflow.services.memory_base.task._fetch_pending_messages", + AsyncMock(return_value=[msg]), + ), + patch( + "langflow.services.memory_base.task._build_documents_from_messages", + return_value=[MagicMock()], + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.is_job_cancelled", + AsyncMock(return_value=False), + ), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.build_embeddings", + AsyncMock(return_value=MagicMock()), + ), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_fresh_chroma_client", + return_value=MagicMock(), + ), + patch("langflow.services.memory_base.task.Chroma"), + patch( + "langflow.services.memory_base.task.KBIngestionHelper.write_documents_to_chroma", + AsyncMock(return_value=1), + ), + patch("langflow.services.memory_base.task._sync_kb_metadata"), + patch("langflow.services.memory_base.task._mark_messages_ingested", AsyncMock()), + patch("langflow.services.memory_base.task._advance_cursor", side_effect=fake_advance_cursor), + patch( + "langflow.services.memory_base.task.KBStorageHelper.get_root_path", + return_value=tmp_path / "kb", + ), + patch("langflow.services.memory_base.task.KBStorageHelper.release_chroma_resources"), + ): + result = await ingest_memory_task( + memory_base_id=mb_id, + session_id="s1", + flow_id=flow_id, + kb_name="kb", + kb_username="user", + user_id=uuid.uuid4(), + embedding_provider="OpenAI", + embedding_model="text-embedding-3-small", + cursor_id=None, + task_job_id=job_id, + job_service=job_service, + ) + + assert result["ingested"] == 1 + assert advance_kwargs["new_cursor_id"] == msg.id + assert advance_kwargs["ingested_count"] == 1 + + def test_sync_kb_metadata_stamps_is_memory_base(self, tmp_path): + """_sync_kb_metadata must write is_memory_base: true to the metadata file.""" + import json + + from langflow.services.memory_base.task import _sync_kb_metadata + + kb_path = tmp_path / "test_kb" + kb_path.mkdir() + + mock_chroma = MagicMock() + + with ( + patch( + "langflow.services.memory_base.task.KBAnalysisHelper.get_metadata", + return_value={"chunks": 0, "embedding_provider": "OpenAI"}, + ), + patch("langflow.services.memory_base.task.KBAnalysisHelper.update_text_metrics"), + patch("langflow.services.memory_base.task.KBStorageHelper.get_directory_size", return_value=1024), + ): + _sync_kb_metadata(kb_path=kb_path, chroma=mock_chroma) + + written = json.loads((kb_path / "embedding_metadata.json").read_text()) + assert written["is_memory_base"] is True + assert "memory" in written.get("source_types", []) + + def test_sync_kb_metadata_failure_does_not_raise(self, tmp_path): + """Metadata sync errors must be swallowed so the cursor can still advance.""" + from langflow.services.memory_base.task import _sync_kb_metadata + + kb_path = tmp_path / "no_such_dir" # does not exist + + with patch( + "langflow.services.memory_base.task.KBAnalysisHelper.get_metadata", + side_effect=OSError("disk full"), + ): + # Must not raise + _sync_kb_metadata(kb_path=kb_path, chroma=MagicMock()) + + def test_build_documents_skips_empty_messages(self): + from langflow.services.memory_base.task import _build_documents_from_messages + + flow_id = uuid.uuid4() + messages = [ + _make_message(flow_id=flow_id, session_id="s1", text=""), + _make_message(flow_id=flow_id, session_id="s1", text=" "), + _make_message(flow_id=flow_id, session_id="s1", text="Valid content here."), + ] + docs = _build_documents_from_messages(messages, session_id="s1", flow_id=str(flow_id)) + assert len(docs) == 1 + assert docs[0].page_content == "Valid content here." + + def test_build_documents_metadata(self): + from langflow.services.memory_base.task import _build_documents_from_messages + + flow_id = uuid.uuid4() + run_id = uuid.uuid4() + msg = _make_message(flow_id=flow_id, session_id="s1", text="Test output.", run_id=run_id) + docs = _build_documents_from_messages([msg], session_id="s1", flow_id=str(flow_id)) + assert docs[0].metadata["message_id"] == str(msg.id) + assert docs[0].metadata["run_id"] == str(run_id) + assert docs[0].metadata["session_id"] == "s1" + + +# ------------------------------------------------------------------ # +# on_flow_output hook and threshold-trigger tests # +# ------------------------------------------------------------------ # + + +class TestOnFlowOutputHook: + """Tests for on_flow_output, _maybe_trigger threshold logic, and hook wiring.""" + + @pytest.fixture + def service(self): + from langflow.services.memory_base.service import MemoryBaseService + + return MemoryBaseService() + + def _fake_scope(self, mock_db): + """Return a mock session_scope context manager backed by mock_db.""" + + class _FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + scope = MagicMock() + scope.return_value = _FakeCtx() + return scope + + @pytest.mark.asyncio + async def test_on_flow_output_skips_when_below_threshold(self, service): + """No job must be created when pending message count is below the threshold.""" + mb = _make_mb(threshold=5) + mbs = _make_session(memory_base_id=mb.id) + mock_db = AsyncMock() + + with ( + patch("langflow.services.memory_base.service.session_scope", self._fake_scope(mock_db)), + patch.object(service, "_get_or_create_session", AsyncMock(return_value=mbs)), + patch.object(service, "_insert_workflow_run", AsyncMock()), + patch.object(service, "_count_pending", AsyncMock(return_value=3)), # 3 < threshold 5 + patch("langflow.services.memory_base.service.get_job_service") as mock_jsc, + ): + await service._maybe_trigger(mb=mb, session_id="s1", job_id=None) + + mock_jsc.return_value.create_job.assert_not_called() + + @pytest.mark.asyncio + async def test_on_flow_output_triggers_when_threshold_met(self, service): + """A job must be created and dispatched when pending message count meets threshold.""" + mb = _make_mb(threshold=3) + mbs = _make_session(memory_base_id=mb.id) + mock_db = AsyncMock() + + mock_job_svc = MagicMock() + mock_job_svc.create_job = AsyncMock() + mock_task_svc = MagicMock() + mock_task_svc.fire_and_forget_task = AsyncMock() + + with ( + patch("langflow.services.memory_base.service.session_scope", self._fake_scope(mock_db)), + patch.object(service, "_get_or_create_session", AsyncMock(return_value=mbs)), + patch.object(service, "_insert_workflow_run", AsyncMock()), + patch.object(service, "_count_pending", AsyncMock(return_value=5)), # 5 >= threshold 3 + patch.object(service, "_get_latest_pending_workflow_job_id", AsyncMock(return_value=uuid.uuid4())), + patch.object(service, "_resolve_kb_username", AsyncMock(return_value="testuser")), + patch.object(service, "_resolve_embedding", return_value=("OpenAI", "text-embedding-3-small")), + patch("langflow.services.memory_base.service.get_job_service", return_value=mock_job_svc), + patch("langflow.services.memory_base.service.get_task_service", return_value=mock_task_svc), + ): + await service._maybe_trigger(mb=mb, session_id="s1", job_id=None) + + mock_job_svc.create_job.assert_awaited_once() + mock_task_svc.fire_and_forget_task.assert_awaited_once() + + @pytest.mark.asyncio + async def test_on_flow_output_is_silent_on_error(self, service): + """on_flow_output must swallow _maybe_trigger exceptions without propagating them. + + This guarantees memory-base failures never cause regressions in flow execution. + """ + flow_id = uuid.uuid4() + mb = _make_mb(flow_id=flow_id, auto_capture=True) + mock_db = AsyncMock() + result_mock = MagicMock() + result_mock.all = MagicMock(return_value=[mb]) + mock_db.exec = AsyncMock(return_value=result_mock) + + with ( + patch("langflow.services.memory_base.service.session_scope", self._fake_scope(mock_db)), + patch.object(service, "_maybe_trigger", AsyncMock(side_effect=RuntimeError("boom"))), + ): + # Must not raise even though _maybe_trigger blows up + await service.on_flow_output(flow_id=flow_id, session_id="s1", job_id=uuid.uuid4()) + + @pytest.mark.asyncio + async def test_hook_wiring_playground(self): + """Playground path: background_tasks.add_task dispatches on_flow_output with correct kwargs. + + Verifies the contract of the hook-dispatch block added to generate_flow_events + in api/build.py after end_all_traces(). + """ + from starlette.background import BackgroundTasks + + flow_id = uuid.uuid4() + run_id = uuid.uuid4() + + mb_service = MagicMock() + mb_service.on_flow_output = AsyncMock() + + bg_tasks = MagicMock(spec=BackgroundTasks) + mock_graph = MagicMock() + mock_graph.run_id = str(run_id) # graph.run_id is always a str + mock_graph.session_id = "test-session" + + with patch("langflow.api.build.get_memory_base_service", return_value=mb_service): + import langflow.api.build as build_module + + # Confirm the import is wired at module level + assert hasattr(build_module, "get_memory_base_service") + + # Execute the same hook-dispatch block as in generate_flow_events + _run_id_uuid = uuid.UUID(mock_graph.run_id) if mock_graph.run_id else None + bg_tasks.add_task( + mb_service.on_flow_output, + flow_id=flow_id, + session_id=mock_graph.session_id or str(flow_id), + run_id=_run_id_uuid, + ) + + bg_tasks.add_task.assert_called_once_with( + mb_service.on_flow_output, + flow_id=flow_id, + session_id="test-session", + run_id=run_id, # UUID, not str — type-cast from graph.run_id + ) + + @pytest.mark.asyncio + async def test_hook_wiring_v2_async_wrapper(self): + """V2 async path: _run_and_notify preserves run_graph_internal result and dispatches hook. + + Verifies the behavioral contract of the closure added to execute_workflow_background + in api/v2/workflow.py: the wrapper must be transparent to execute_with_status + (return value unchanged) while also firing the memory-base hook. + """ + expected_result = (MagicMock(), "effective-session-42") + run_graph_mock = AsyncMock(return_value=expected_result) + hook_mock = AsyncMock() + task_service_mock = MagicMock() + task_service_mock.fire_and_forget_task = AsyncMock() + + hook_flow_id = uuid.uuid4() + hook_run_id = uuid.uuid4() + + # Mirror the _run_and_notify closure from workflow.py execute_workflow_background + async def _run_and_notify(**kwargs): + result = await run_graph_mock(**kwargs) + _, _effective_session_id = result + with contextlib.suppress(Exception): + await task_service_mock.fire_and_forget_task( + hook_mock, + flow_id=hook_flow_id, + session_id=_effective_session_id, + run_id=hook_run_id, + ) + return result + + result = await _run_and_notify(graph=MagicMock()) + + # Return value must be identical — execute_with_status depends on this + assert result == expected_result + # Hook must be dispatched with the session_id extracted from run_graph_internal + task_service_mock.fire_and_forget_task.assert_awaited_once_with( + hook_mock, + flow_id=hook_flow_id, + session_id="effective-session-42", + run_id=hook_run_id, + ) + + @pytest.mark.asyncio + async def test_hook_failure_does_not_affect_wrapper_return(self): + """If fire_and_forget_task raises inside _run_and_notify, the return value is still correct.""" + expected_result = (MagicMock(), "some-session") + run_graph_mock = AsyncMock(return_value=expected_result) + task_service_mock = MagicMock() + task_service_mock.fire_and_forget_task = AsyncMock(side_effect=RuntimeError("dispatch failed")) + + async def _run_and_notify(**kwargs): + result = await run_graph_mock(**kwargs) + _, _effective_session_id = result + with contextlib.suppress(Exception): + await task_service_mock.fire_and_forget_task( + AsyncMock(), + flow_id=uuid.uuid4(), + session_id=_effective_session_id, + run_id=uuid.uuid4(), + ) + return result + + result = await _run_and_notify(graph=MagicMock()) + assert result == expected_result + + +# ------------------------------------------------------------------ # +# API endpoint routing tests # +# ------------------------------------------------------------------ # + + +class TestMemoriesAPIRouting: + """Verify routing and response codes without hitting the DB.""" + + @pytest.fixture + def patched_service(self): + """Patch get_memory_base_service in memories.py.""" + mock_svc = MagicMock() + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=mock_svc): + yield mock_svc + + @pytest.mark.asyncio + async def test_get_not_found_returns_404(self, patched_service): + """Handler returns 404 when service.get returns None (covered via direct call).""" + from fastapi import HTTPException + from langflow.api.v1.memories import get_memory_base + + patched_service.get = AsyncMock(return_value=None) + + mock_user = MagicMock() + mock_user.id = uuid.uuid4() + + with pytest.raises(HTTPException) as exc_info: + await get_memory_base(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert exc_info.value.status_code == 404 + + @pytest.mark.asyncio + async def test_flush_conflict_returns_409(self, patched_service): + """trigger_ingestion raising RuntimeError should map to HTTP 409.""" + from langflow.api.v1.memories import flush_memory_base + + patched_service.trigger_ingestion = AsyncMock(side_effect=RuntimeError("already in progress")) + + # We call the handler directly to test the error mapping + mock_user = MagicMock() + mock_user.id = uuid.uuid4() + + from fastapi import HTTPException + from langflow.api.v1.memories import FlushRequest + + with pytest.raises(HTTPException) as exc_info: + await flush_memory_base( + memory_base_id=uuid.uuid4(), + body=FlushRequest(session_id="s1"), + current_user=mock_user, + ) + from fastapi import HTTPException + + assert isinstance(exc_info.value, HTTPException) + assert exc_info.value.status_code == 409 + + +# ------------------------------------------------------------------ # +# API handler unit tests (direct invocation, no HTTP stack) # +# ------------------------------------------------------------------ # + + +class TestMemoriesAPIHandlers: + """Call endpoint handlers directly, mocking _service, to test all status-code branches.""" + + @pytest.fixture + def mock_user(self): + user = MagicMock() + user.id = uuid.uuid4() + return user + + # ---------------------------------------------------------------- # + # create_memory_base # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_create_success_returns_memory_base_read(self, mock_user): + from langflow.api.v1.memories import create_memory_base + + mb = _make_mb(user_id=mock_user.id) + payload = MemoryBaseCreate(name="mb", flow_id=mb.flow_id, user_id=mock_user.id, kb_name="kb") + + svc = MagicMock() + svc.create = AsyncMock(return_value=mb) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await create_memory_base(current_user=mock_user, payload=payload) + + assert result.id == mb.id + assert result.user_id == mock_user.id + + @pytest.mark.asyncio + async def test_create_duplicate_name_returns_409(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import create_memory_base + + payload = MemoryBaseCreate(name="dup", flow_id=uuid.uuid4(), user_id=mock_user.id, kb_name="kb") + + svc = MagicMock() + svc.create = AsyncMock(side_effect=ValueError("name already in use")) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await create_memory_base(current_user=mock_user, payload=payload) + + assert exc_info.value.status_code == 409 + + # ---------------------------------------------------------------- # + # get_memory_base # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_get_success_returns_memory_base_read(self, mock_user): + from langflow.api.v1.memories import get_memory_base + + mb = _make_mb(user_id=mock_user.id) + + svc = MagicMock() + svc.get = AsyncMock(return_value=mb) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await get_memory_base(memory_base_id=mb.id, current_user=mock_user) + + assert result.id == mb.id + + @pytest.mark.asyncio + async def test_get_not_found_raises_404(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import get_memory_base + + svc = MagicMock() + svc.get = AsyncMock(return_value=None) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await get_memory_base(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert exc_info.value.status_code == 404 + + # ---------------------------------------------------------------- # + # update_memory_base # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_update_success_returns_updated_record(self, mock_user): + from langflow.api.v1.memories import update_memory_base + + mb = _make_mb(user_id=mock_user.id, threshold=99) + + svc = MagicMock() + svc.update = AsyncMock(return_value=mb) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await update_memory_base( + memory_base_id=mb.id, + current_user=mock_user, + patch=MemoryBaseUpdate(threshold=99), + ) + + assert result.threshold == 99 + + @pytest.mark.asyncio + async def test_update_not_found_raises_404(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import update_memory_base + + svc = MagicMock() + svc.update = AsyncMock(return_value=None) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await update_memory_base( + memory_base_id=uuid.uuid4(), + current_user=mock_user, + patch=MemoryBaseUpdate(threshold=5), + ) + + assert exc_info.value.status_code == 404 + + # ---------------------------------------------------------------- # + # delete_memory_base # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_delete_success_returns_none(self, mock_user): + from langflow.api.v1.memories import delete_memory_base + + svc = MagicMock() + svc.delete = AsyncMock(return_value=True) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await delete_memory_base(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert result is None + + @pytest.mark.asyncio + async def test_delete_not_found_raises_404(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import delete_memory_base + + svc = MagicMock() + svc.delete = AsyncMock(return_value=False) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await delete_memory_base(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert exc_info.value.status_code == 404 + + # ---------------------------------------------------------------- # + # flush_memory_base # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_flush_success_returns_job_id(self, mock_user): + from langflow.api.v1.memories import FlushRequest, flush_memory_base + + job_id = str(uuid.uuid4()) + + svc = MagicMock() + svc.trigger_ingestion = AsyncMock(return_value=job_id) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await flush_memory_base( + memory_base_id=uuid.uuid4(), + current_user=mock_user, + body=FlushRequest(session_id="s1"), + ) + + assert result == {"job_id": job_id} + + @pytest.mark.asyncio + async def test_flush_value_error_raises_404(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import FlushRequest, flush_memory_base + + svc = MagicMock() + svc.trigger_ingestion = AsyncMock(side_effect=ValueError("memory base not found")) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await flush_memory_base( + memory_base_id=uuid.uuid4(), + current_user=mock_user, + body=FlushRequest(session_id="s1"), + ) + + assert exc_info.value.status_code == 404 + + @pytest.mark.asyncio + async def test_flush_duplicate_job_error_raises_409(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import FlushRequest, flush_memory_base + from langflow.services.jobs import DuplicateJobError + + svc = MagicMock() + svc.trigger_ingestion = AsyncMock(side_effect=DuplicateJobError("already running")) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await flush_memory_base( + memory_base_id=uuid.uuid4(), + current_user=mock_user, + body=FlushRequest(session_id="s1"), + ) + + assert exc_info.value.status_code == 409 + + # ---------------------------------------------------------------- # + # check_mismatch # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_check_mismatch_detected_returns_true(self, mock_user): + from langflow.api.v1.memories import check_mismatch + + svc = MagicMock() + svc.check_mismatch = AsyncMock(return_value=True) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await check_mismatch(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert result.mismatch_detected is True + + @pytest.mark.asyncio + async def test_check_mismatch_not_detected_returns_false(self, mock_user): + from langflow.api.v1.memories import check_mismatch + + svc = MagicMock() + svc.check_mismatch = AsyncMock(return_value=False) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await check_mismatch(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert result.mismatch_detected is False + + @pytest.mark.asyncio + async def test_check_mismatch_not_found_raises_404(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import check_mismatch + + svc = MagicMock() + svc.check_mismatch = AsyncMock(side_effect=ValueError("not found")) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await check_mismatch(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert exc_info.value.status_code == 404 + + # ---------------------------------------------------------------- # + # regenerate_memory_base # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_regenerate_success_returns_job_ids(self, mock_user): + from langflow.api.v1.memories import regenerate_memory_base + + job_ids = [str(uuid.uuid4()), str(uuid.uuid4())] + + svc = MagicMock() + svc.regenerate = AsyncMock(return_value=job_ids) + with patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc): + result = await regenerate_memory_base(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert result.job_ids == job_ids + + @pytest.mark.asyncio + async def test_regenerate_not_found_raises_404(self, mock_user): + from fastapi import HTTPException + from langflow.api.v1.memories import regenerate_memory_base + + svc = MagicMock() + svc.regenerate = AsyncMock(side_effect=ValueError("not found")) + with ( + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await regenerate_memory_base(memory_base_id=uuid.uuid4(), current_user=mock_user) + + assert exc_info.value.status_code == 404 + + # ---------------------------------------------------------------- # + # list_sessions # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_list_sessions_ownership_failure_raises_404(self, mock_user): + from fastapi import HTTPException + from fastapi_pagination import Params + from langflow.api.v1.memories import list_sessions + + mock_db = AsyncMock() + + class FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + svc = MagicMock() + svc._get_mb_or_raise = AsyncMock(side_effect=ValueError("not found")) + with ( + patch("langflow.api.v1.memories.session_scope", return_value=FakeCtx()), + patch("langflow.api.v1.memories.get_memory_base_service", return_value=svc), + pytest.raises(HTTPException) as exc_info, + ): + await list_sessions( + memory_base_id=uuid.uuid4(), + current_user=mock_user, + params=Params(), + ) + + assert exc_info.value.status_code == 404 + + # ---------------------------------------------------------------- # + # list_session_messages # + # ---------------------------------------------------------------- # + + @pytest.mark.asyncio + async def test_list_session_messages_not_found_raises_404(self, mock_user): + from fastapi import HTTPException + from fastapi_pagination import Params + from langflow.api.v1.memories import list_session_messages + + mock_db = AsyncMock() + result_mock = MagicMock() + result_mock.first = MagicMock(return_value=None) + mock_db.exec = AsyncMock(return_value=result_mock) + + class FakeCtx: + async def __aenter__(self): + return mock_db + + async def __aexit__(self, *a): + pass + + with ( + patch("langflow.api.v1.memories.session_scope", return_value=FakeCtx()), + pytest.raises(HTTPException) as exc_info, + ): + await list_session_messages( + memory_base_id=uuid.uuid4(), + session_id="s1", + current_user=mock_user, + params=Params(), + ) + + assert exc_info.value.status_code == 404 + + # ---------------------------------------------------------------- # + # MessageReadResponse schema # + # ---------------------------------------------------------------- # + + def test_message_read_response_from_attributes(self): + from langflow.api.v1.memories import MessageReadResponse + + msg = _make_message(flow_id=uuid.uuid4(), session_id="s1", text="hello") + response = MessageReadResponse.model_validate(msg, from_attributes=True) + + assert response.text == "hello" + assert response.session_id == "s1" + assert response.sender == "AI" + assert response.content_blocks == [] + + def test_flush_request_schema(self): + from langflow.api.v1.memories import FlushRequest + + req = FlushRequest(session_id="my-session") + assert req.session_id == "my-session" + + def test_mismatch_response_schema(self): + from langflow.api.v1.memories import MismatchResponse + + assert MismatchResponse(mismatch_detected=True).mismatch_detected is True + assert MismatchResponse(mismatch_detected=False).mismatch_detected is False + + def test_regenerate_response_schema(self): + from langflow.api.v1.memories import RegenerateResponse + + ids = ["a", "b", "c"] + assert RegenerateResponse(job_ids=ids).job_ids == ids