From d72ebbc554cadb4bbc30d73424af3d876c69be41 Mon Sep 17 00:00:00 2001 From: ump45nose <52391318+ump45nose@users.noreply.github.com> Date: Tue, 11 Aug 2026 16:18:46 +0800 Subject: [PATCH] fix: isolate citation registries per pipeline --- examples/demos/LightResearch.yaml | 1 + .../demos/server/LightResearch_server.yaml | 5 ++ servers/custom/src/custom.py | 61 +++++++++++++------ tests/test_citation_registry.py | 39 ++++++++++++ 4 files changed, 88 insertions(+), 18 deletions(-) create mode 100644 tests/test_citation_registry.py diff --git a/examples/demos/LightResearch.yaml b/examples/demos/LightResearch.yaml index 5d1618d1..f8cbd3d9 100644 --- a/examples/demos/LightResearch.yaml +++ b/examples/demos/LightResearch.yaml @@ -52,3 +52,4 @@ pipeline: complete: [] - prompt.webnote_gen_answer - generation.generate +- custom.clear_citation_registry diff --git a/examples/demos/server/LightResearch_server.yaml b/examples/demos/server/LightResearch_server.yaml index 4c8af5d8..553900d4 100644 --- a/examples/demos/server/LightResearch_server.yaml +++ b/examples/demos/server/LightResearch_server.yaml @@ -56,11 +56,16 @@ custom: q_ls: q_ls output: - q_ls + - citation_registry_id assign_citation_ids_stateful: input: ret_psg: ret_psg + citation_registry_id: citation_registry_id output: - ret_psg + clear_citation_registry: + input: + citation_registry_id: citation_registry_id parameter: servers/custom/parameter.yaml path: servers/custom/src/custom.py prompt: diff --git a/servers/custom/src/custom.py b/servers/custom/src/custom.py index f7e61a86..266ad158 100644 --- a/servers/custom/src/custom.py +++ b/servers/custom/src/custom.py @@ -1,7 +1,8 @@ -import re -import json import copy -from typing import List, Dict, Any +import json +import re +from typing import Any, Dict, List +from uuid import uuid4 from ultrarag.server import UltraRAG_MCP_Server @@ -400,21 +401,30 @@ def assign_citation_ids( class CitationRegistry: - _instances: Dict[int, Dict[str, Any]] = {} + _instances: Dict[str, Dict[int, Dict[str, Any]]] = {} @classmethod - def reset(cls): - cls._instances = {} + def create(cls) -> str: + registry_id = uuid4().hex + cls._instances[registry_id] = {} + return registry_id @classmethod - def get_or_create(cls, query_index: int) -> Dict[str, Any]: - if query_index not in cls._instances: - cls._instances[query_index] = {"registry": {}, "counter": 0} - return cls._instances[query_index] + def clear(cls, registry_id: str) -> None: + cls._instances.pop(registry_id, None) @classmethod - def assign_id(cls, query_index: int, doc_text: str) -> int: - state = cls.get_or_create(query_index) + def get_or_create(cls, registry_id: str, query_index: int) -> Dict[str, Any]: + if registry_id not in cls._instances: + raise ValueError(f"Unknown citation registry: {registry_id}") + registry = cls._instances[registry_id] + if query_index not in registry: + registry[query_index] = {"registry": {}, "counter": 0} + return registry[query_index] + + @classmethod + def assign_id(cls, registry_id: str, query_index: int, doc_text: str) -> int: + state = cls.get_or_create(registry_id, query_index) doc_hash = doc_text.strip() if doc_hash in state["registry"]: @@ -425,7 +435,7 @@ def assign_id(cls, query_index: int, doc_text: str) -> int: return state["counter"] -@app.tool(output="q_ls->q_ls") +@app.tool(output="q_ls->q_ls,citation_registry_id") def init_citation_registry(q_ls: List[str]) -> Dict[str, Any]: """Initialize citation registry for stateful citation assignment. @@ -433,20 +443,24 @@ def init_citation_registry(q_ls: List[str]) -> Dict[str, Any]: q_ls: List of queries Returns: - Dictionary with 'q_ls' (pass-through) + Dictionary with 'q_ls' (pass-through) and an isolated registry ID """ - CitationRegistry.reset() - return {"q_ls": q_ls} + return { + "q_ls": q_ls, + "citation_registry_id": CitationRegistry.create(), + } -@app.tool(output="ret_psg->ret_psg") +@app.tool(output="ret_psg,citation_registry_id->ret_psg") def assign_citation_ids_stateful( ret_psg: List[List[str]], + citation_registry_id: str, ) -> Dict[str, Any]: """Assign unique citation IDs to passages using stateful registry. Args: ret_psg: List of lists of document strings + citation_registry_id: Registry ID returned by init_citation_registry Returns: Dictionary with 'ret_psg' containing passages with unique citation IDs @@ -457,7 +471,11 @@ def assign_citation_ids_stateful( cited_docs = [] for doc in docs_list: doc_text = str(doc).strip() - doc_id = CitationRegistry.assign_id(i, doc_text) + doc_id = CitationRegistry.assign_id( + citation_registry_id, + i, + doc_text, + ) cited_docs.append(f"[{doc_id}] {doc_text}") result_psg.append(cited_docs) @@ -466,6 +484,13 @@ def assign_citation_ids_stateful( } +@app.tool(output="citation_registry_id->None") +def clear_citation_registry(citation_registry_id: str) -> Dict[str, Any]: + """Release citation state after a pipeline finishes.""" + CitationRegistry.clear(citation_registry_id) + return {} + + # ==================== SurveyCPM Citation Tools ==================== diff --git a/tests/test_citation_registry.py b/tests/test_citation_registry.py new file mode 100644 index 00000000..9b339dfe --- /dev/null +++ b/tests/test_citation_registry.py @@ -0,0 +1,39 @@ +from importlib.util import module_from_spec, spec_from_file_location +from pathlib import Path + +import pytest + +CUSTOM_MODULE_PATH = ( + Path(__file__).resolve().parents[1] / "servers" / "custom" / "src" / "custom.py" +) + + +def _load_custom_module(): + spec = spec_from_file_location("ultrarag_custom", CUSTOM_MODULE_PATH) + assert spec and spec.loader + module = module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_citation_registries_are_isolated_between_pipeline_runs() -> None: + custom = _load_custom_module() + + registry_a = custom.init_citation_registry(["request-a"])["citation_registry_id"] + first_a = custom.assign_citation_ids_stateful([["doc-a"]], registry_a) + + registry_b = custom.init_citation_registry(["request-b"])["citation_registry_id"] + first_b = custom.assign_citation_ids_stateful([["doc-b"]], registry_b) + continued_a = custom.assign_citation_ids_stateful( + [["doc-a", "doc-c"]], + registry_a, + ) + + assert first_a["ret_psg"] == [["[1] doc-a"]] + assert first_b["ret_psg"] == [["[1] doc-b"]] + assert continued_a["ret_psg"] == [["[1] doc-a", "[2] doc-c"]] + + custom.clear_citation_registry(registry_a) + custom.clear_citation_registry(registry_b) + with pytest.raises(ValueError, match="Unknown citation registry"): + custom.assign_citation_ids_stateful([["doc-a"]], registry_a)