Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions examples/demos/LightResearch.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -52,3 +52,4 @@ pipeline:
complete: []
- prompt.webnote_gen_answer
- generation.generate
- custom.clear_citation_registry
5 changes: 5 additions & 0 deletions examples/demos/server/LightResearch_server.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
61 changes: 43 additions & 18 deletions servers/custom/src/custom.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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"]:
Expand All @@ -425,28 +435,32 @@ 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.

Args:
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
Expand All @@ -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)

Expand All @@ -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 ====================


Expand Down
39 changes: 39 additions & 0 deletions tests/test_citation_registry.py
Original file line number Diff line number Diff line change
@@ -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)