From 17919ebffad7607891d7e3b3d68b3c484d048e71 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 26 May 2026 17:26:37 +0800 Subject: [PATCH] up --- reme4/components/service/base_service.py | 41 +- reme4/components/service/http_service.py | 116 +-- reme4/components/service/mcp_service.py | 34 +- tests4/unittest/test_bm25_lite.py | 586 -------------- tests4/unittest/test_keyword_index.py | 931 +++++++++++++++++++++++ 5 files changed, 1041 insertions(+), 667 deletions(-) delete mode 100644 tests4/unittest/test_bm25_lite.py create mode 100644 tests4/unittest/test_keyword_index.py diff --git a/reme4/components/service/base_service.py b/reme4/components/service/base_service.py index 1c50abd3..8d4d2402 100644 --- a/reme4/components/service/base_service.py +++ b/reme4/components/service/base_service.py @@ -1,10 +1,14 @@ -"""Base service class for exposing jobs via HTTP, MCP, etc.""" +"""Base class for services that expose jobs over a network protocol.""" +import json +import os from abc import abstractmethod +from contextlib import asynccontextmanager from typing import TYPE_CHECKING from ..base_component import BaseComponent from ..job.base_job import BaseJob +from ...constants import REME_SERVICE_INFO from ...enumeration import ComponentEnum if TYPE_CHECKING: @@ -12,28 +16,51 @@ if TYPE_CHECKING: class BaseService(BaseComponent): - """Base class for services that expose jobs via HTTP, MCP, etc.""" + """Skeleton for services (HTTP, MCP, ...) that turn jobs into endpoints or tools.""" component_type = ComponentEnum.SERVICE def __init__(self, **kwargs): super().__init__(**kwargs) + # Underlying framework instance (FastAPI, FastMCP, ...); populated by build_service(). self.service = None + # ----- Subclass contract --------------------------------------------- + @abstractmethod def build_service(self, app: "Application") -> None: - """Initialize the underlying service framework.""" + """Instantiate and configure the underlying server framework.""" @abstractmethod def add_job(self, job: BaseJob) -> None: - """Register a single job with the service.""" + """Register a single job as a callable endpoint or tool.""" @abstractmethod def start_service(self, app: "Application") -> None: - """Start serving requests.""" + """Block on serving requests until shutdown.""" + + # ----- Shared helpers ------------------------------------------------ + + def _lifespan(self, app: "Application", host: str, port: int): + """Build an async-context lifespan that brackets the server with app start/close. + + Publishes the bound address via the REME_SERVICE_INFO environment variable so + in-process clients can discover where this service is listening. + """ + + @asynccontextmanager + async def lifespan(_): + await app.start() + service_info = json.dumps({"host": host, "port": port}) + os.environ[REME_SERVICE_INFO] = service_info + self.logger.info(f"{self.name} started: {REME_SERVICE_INFO}={service_info}") + yield + await app.close() + + return lifespan def add_jobs(self, app: "Application") -> None: - """Register all non-background jobs from the application context.""" + """Register every job from the app context except background-only ones.""" for name, job in app.context.jobs.items(): if job.backend == "background": continue @@ -44,7 +71,7 @@ class BaseService(BaseComponent): self.logger.error(f"Failed to add job {name}: {e}") def run_app(self, app: "Application") -> None: - """Build, populate, and start the service.""" + """Build the service, register jobs, then start serving (blocking).""" self.build_service(app) self.add_jobs(app) self.start_service(app) diff --git a/reme4/components/service/http_service.py b/reme4/components/service/http_service.py index a1092947..95ed5e72 100644 --- a/reme4/components/service/http_service.py +++ b/reme4/components/service/http_service.py @@ -1,11 +1,8 @@ -"""HTTP service implementation for ReMe.""" +"""HTTP service: exposes jobs as FastAPI endpoints (JSON, or SSE for stream jobs).""" import asyncio -import json -import os import warnings from collections.abc import AsyncGenerator -from contextlib import asynccontextmanager from typing import TYPE_CHECKING import uvicorn @@ -16,7 +13,7 @@ from fastapi.responses import StreamingResponse from .base_service import BaseService from ..component_registry import R from ..job import BaseJob, StreamJob -from ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT, REME_SERVICE_INFO +from ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT from ...schema import Request, Response from ...utils import execute_stream_task @@ -24,27 +21,76 @@ if TYPE_CHECKING: from ...application import Application +# uvicorn 0.41 still imports these deprecated websockets symbols on startup, +# even though we don't use WebSocket. Silence just those specific warnings. +_WEBSOCKET_DEPRECATION_PATTERNS = ( + r".*websockets\.legacy is deprecated.*", + r".*WebSocketServerProtocol is deprecated.*", +) + + @R.register("http") class HttpService(BaseService): - """HTTP service: normal jobs -> JSON endpoints, stream jobs -> SSE endpoints.""" + """Map non-stream jobs to JSON POST endpoints and StreamJobs to SSE endpoints.""" def __init__(self, host: str = REME_DEFAULT_HOST, port: int = REME_DEFAULT_PORT, **kwargs): super().__init__(**kwargs) self.host: str = host self.port: int = port - def _add_job(self, job: BaseJob) -> None: - async def execute_endpoint(request: Request) -> Response: + # ----- BaseService contract ------------------------------------------ + + def build_service(self, app: "Application") -> None: + """Create the FastAPI app with permissive CORS and an app-managed lifespan.""" + self.service = FastAPI( + title=app.config.app_name, + lifespan=self._lifespan(app, self.host, self.port), + ) + self.service.add_middleware( + CORSMiddleware, # type: ignore[arg-type] + allow_origins=["*"], + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], + ) + + def add_job(self, job: BaseJob) -> None: + """Dispatch to streaming or non-streaming registration based on job type.""" + if isinstance(job, StreamJob): + self._add_stream_job(job) + else: + self._add_json_job(job) + + def start_service(self, app: "Application") -> None: + """Run uvicorn, suppressing unrelated websocket deprecation noise.""" + for pattern in _WEBSOCKET_DEPRECATION_PATTERNS: + warnings.filterwarnings("ignore", category=DeprecationWarning, message=pattern) + uvicorn.run(self.service, host=self.host, port=self.port, **self.kwargs) + + # ----- Endpoint factories -------------------------------------------- + + def _add_json_job(self, job: BaseJob) -> None: + """Register a job as POST /{job.name} returning a JSON Response.""" + + async def endpoint(request: Request) -> Response: return await job(**request.model_dump(exclude_none=True)) - self.service.post(path=f"/{job.name}", response_model=Response, description=job.description)(execute_endpoint) + self.service.post( + f"/{job.name}", + response_model=Response, + description=job.description, + )(endpoint) def _add_stream_job(self, job: StreamJob) -> None: - async def execute_stream_endpoint(request: Request) -> StreamingResponse: - stream_queue = asyncio.Queue() - task = asyncio.create_task(job(stream_queue=stream_queue, **request.model_dump(exclude_none=True))) + """Register a StreamJob as POST /{job.name} streaming chunks as text/event-stream.""" - async def generate_stream() -> AsyncGenerator[bytes, None]: + async def endpoint(request: Request) -> StreamingResponse: + stream_queue: asyncio.Queue = asyncio.Queue() + task = asyncio.create_task( + job(stream_queue=stream_queue, **request.model_dump(exclude_none=True)), + ) + + async def body() -> AsyncGenerator[bytes, None]: async for chunk in execute_stream_task( stream_queue=stream_queue, task=task, @@ -54,46 +100,6 @@ class HttpService(BaseService): assert isinstance(chunk, bytes) yield chunk - return StreamingResponse(generate_stream(), media_type="text/event-stream") + return StreamingResponse(body(), media_type="text/event-stream") - self.service.post(f"/{job.name}")(execute_stream_endpoint) - - def add_job(self, job: BaseJob) -> None: - if isinstance(job, StreamJob): - self._add_stream_job(job) - else: - self._add_job(job) - - def build_service(self, app: "Application") -> None: - @asynccontextmanager - async def lifespan(_: FastAPI): - await app.start() - service_info = json.dumps({"host": self.host, "port": self.port}) - os.environ[REME_SERVICE_INFO] = service_info - self.logger.info(f"ReMe Service started: {REME_SERVICE_INFO}={service_info}") - yield - await app.close() - - self.service = FastAPI(title=app.config.app_name, lifespan=lifespan) - self.service.add_middleware( - CORSMiddleware, # type: ignore[arg-type] - allow_origins=["*"], - allow_credentials=True, - allow_methods=["*"], - allow_headers=["*"], - ) - - def start_service(self, app: "Application") -> None: - # uvicorn 0.41 still imports websockets.legacy / WebSocketServerProtocol - # on startup; silence those specific lines since we don't use WebSocket. - warnings.filterwarnings( - "ignore", - category=DeprecationWarning, - message=r".*websockets\.legacy is deprecated.*", - ) - warnings.filterwarnings( - "ignore", - category=DeprecationWarning, - message=r".*WebSocketServerProtocol is deprecated.*", - ) - uvicorn.run(self.service, host=self.host, port=self.port, **self.kwargs) + self.service.post(f"/{job.name}")(endpoint) diff --git a/reme4/components/service/mcp_service.py b/reme4/components/service/mcp_service.py index 8f22450d..9203aff9 100644 --- a/reme4/components/service/mcp_service.py +++ b/reme4/components/service/mcp_service.py @@ -1,8 +1,5 @@ -"""MCP (Model Context Protocol) service implementation.""" +"""MCP (Model Context Protocol) service: exposes jobs as MCP tools.""" -import json -import os -from contextlib import asynccontextmanager from typing import TYPE_CHECKING from fastmcp import FastMCP @@ -11,8 +8,8 @@ from fastmcp.tools import FunctionTool from .base_service import BaseService from ..component_registry import R -from ..job import StreamJob, BaseJob -from ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT, REME_SERVICE_INFO +from ..job import BaseJob, StreamJob +from ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT if TYPE_CHECKING: from ...application import Application @@ -20,7 +17,7 @@ if TYPE_CHECKING: @R.register("mcp") class MCPService(BaseService): - """Expose jobs as MCP (Model Context Protocol) tools.""" + """Expose non-stream jobs as MCP tools over stdio, SSE, or other supported transports.""" def __init__( self, @@ -34,19 +31,17 @@ class MCPService(BaseService): self.host: str = host self.port: int = port + # ----- BaseService contract ------------------------------------------ + def build_service(self, app: "Application") -> None: - @asynccontextmanager - async def lifespan(_: FastMCP): - await app.start() - service_info = json.dumps({"host": self.host, "port": self.port}) - os.environ[REME_SERVICE_INFO] = service_info - self.logger.info(f"ReMe MCP Service started: {REME_SERVICE_INFO}={service_info}") - yield - await app.close() + """Create the FastMCP server with an app-managed lifespan.""" + self.service = FastMCP( + name=app.config.app_name, + lifespan=self._lifespan(app, self.host, self.port), + ) - self.service = FastMCP(name=app.config.app_name, lifespan=lifespan) - - def add_job(self, job: "BaseJob") -> None: + def add_job(self, job: BaseJob) -> None: + """Register a non-stream job as an MCP tool; StreamJobs are skipped (not supported).""" if isinstance(job, StreamJob): return @@ -64,7 +59,8 @@ class MCPService(BaseService): ) def start_service(self, app: "Application") -> None: - transport_kwargs = {} + """Run the MCP server; bind host/port only when the transport is network-based.""" + transport_kwargs: dict = {} if self.transport != "stdio": transport_kwargs["host"] = self.host transport_kwargs["port"] = self.port diff --git a/tests4/unittest/test_bm25_lite.py b/tests4/unittest/test_bm25_lite.py deleted file mode 100644 index 75865797..00000000 --- a/tests4/unittest/test_bm25_lite.py +++ /dev/null @@ -1,586 +0,0 @@ -"""Tests for BM25Index search engine.""" - -# pylint: disable=protected-access - -import asyncio -import os -import tempfile -import warnings - -from reme4.components.keyword_index import BM25Index -from reme4.components.tokenizer import RegexTokenizer - -# Filter jieba/pkg_resources deprecation warnings -warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba") -warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources") - - -class temp_chdir: - """Context manager to temporarily chdir into a path and restore on exit.""" - - def __init__(self, path): - self.path = path - self.old = None - - def __enter__(self): - self.old = os.getcwd() - os.chdir(self.path) - return self - - def __exit__(self, *exc): - os.chdir(self.old) - - -async def create_bm25(k1: float = 1.5, b: float = 0.75) -> BM25Index: - """Create and start a BM25Index in cwd with a non-filtering tokenizer. - - The non-filtering tokenizer keeps short test texts (e.g. "hello world") visible, - since several common test words ("hello", "我", "的") are in the default stopwords. - """ - bm25 = BM25Index(k1=k1, b=b) - # Replace the unresolved Dependency placeholder with a real tokenizer instance. - tokenizer = RegexTokenizer(filter_stopwords=False) - bm25.tokenizer = tokenizer - bm25._owned.append(tokenizer) - await bm25.start() - return bm25 - - -def test_basic_init(): - """Test BM25Index initialization.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = BM25Index() - assert bm25.k1 == 1.5 - assert bm25.b == 0.75 - assert bm25.vocab == {} - assert bm25.inverted_index == {} - assert bm25.doc_meta == {} - assert bm25.n_docs == 0 - assert bm25.avg_len == 0.0 - print("✓ test_basic_init passed") - - asyncio.run(run()) - - -def test_start_with_tokenizer(): - """Test BM25Index starts and initializes tokenizer.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - assert bm25.tokenizer is not None - assert bm25.is_started - - await bm25.close() - assert not bm25.is_started - print("✓ test_start_with_tokenizer passed") - - asyncio.run(run()) - - -def test_add_single_doc(): - """Test adding a single document.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs({"doc1": "hello world"}) - - assert bm25.n_docs == 1 - assert bm25.total_len > 0 - assert "doc1" in bm25.doc_meta - - await bm25.close() - print("✓ test_add_single_doc passed") - - asyncio.run(run()) - - -def test_add_multiple_docs(): - """Test adding multiple documents.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = { - "doc1": "hello world", - "doc2": "hello python", - "doc3": "world python", - } - await bm25.add_docs(docs) - - assert bm25.n_docs == 3 - assert len(bm25.vocab) > 0 - assert len(bm25.inverted_index) > 0 - - await bm25.close() - print("✓ test_add_multiple_docs passed") - - asyncio.run(run()) - - -def test_retrieve_basic(): - """Test basic retrieval functionality.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = { - "doc1": "python programming language", - "doc2": "java programming language", - "doc3": "python data analysis", - } - await bm25.add_docs(docs) - - results = await bm25.retrieve("python", limit=3) - assert len(results) <= 3 - assert "doc1" in results or "doc3" in results - - await bm25.close() - print("✓ test_retrieve_basic passed") - - asyncio.run(run()) - - -def test_retrieve_with_limit(): - """Test retrieval with result limit.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = {f"doc{i}": f"python programming {i}" for i in range(10)} - await bm25.add_docs(docs) - - results = await bm25.retrieve("python", limit=3) - assert len(results) == 3 - - results = await bm25.retrieve("python", limit=5) - assert len(results) == 5 - - await bm25.close() - print("✓ test_retrieve_with_limit passed") - - asyncio.run(run()) - - -def test_retrieve_empty_query(): - """Test retrieval with empty or unknown query.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = {"doc1": "hello world"} - await bm25.add_docs(docs) - - results = await bm25.retrieve("", limit=3) - assert results == {} - - results = await bm25.retrieve("unknownxyz", limit=3) - assert results == {} - - await bm25.close() - print("✓ test_retrieve_empty_query passed") - - asyncio.run(run()) - - -def test_retrieve_empty_index(): - """Test retrieval from empty index.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - results = await bm25.retrieve("python", limit=3) - assert results == {} - - await bm25.close() - print("✓ test_retrieve_empty_index passed") - - asyncio.run(run()) - - -def test_update_doc(): - """Test updating an existing document.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs({"doc1": "hello world python"}) - old_len = bm25.total_len - - await bm25.add_docs({"doc1": "java"}) - assert bm25.n_docs == 1 - assert bm25.total_len != old_len - - results = await bm25.retrieve("java", limit=1) - assert "doc1" in results - - await bm25.close() - print("✓ test_update_doc passed") - - asyncio.run(run()) - - -def test_remove_doc(): - """Test removing a document.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = { - "doc1": "hello world", - "doc2": "hello python", - } - await bm25.add_docs(docs) - assert bm25.n_docs == 2 - - bm25._remove_doc("doc1") - assert bm25.n_docs == 1 - assert "doc1" not in bm25.doc_meta - - results = await bm25.retrieve("hello", limit=2) - assert "doc1" not in results - assert "doc2" in results - - await bm25.close() - print("✓ test_remove_doc passed") - - asyncio.run(run()) - - -def test_remove_nonexistent_doc(): - """Test removing a nonexistent document (should be no-op).""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs({"doc1": "hello world"}) - bm25._remove_doc("nonexistent") - assert bm25.n_docs == 1 - - await bm25.close() - print("✓ test_remove_nonexistent_doc passed") - - asyncio.run(run()) - - -def test_clear(): - """Test clearing the index.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs( - { - "doc1": "hello world", - "doc2": "hello python", - }, - ) - assert bm25.n_docs == 2 - - await bm25.clear() - assert bm25.n_docs == 0 - assert bm25.vocab == {} - assert bm25.inverted_index == {} - assert bm25.doc_meta == {} - assert bm25.total_len == 0 - assert bm25._idf_cache == {} - - await bm25.close() - print("✓ test_clear passed") - - asyncio.run(run()) - - -def test_optimize_index(): - """Test optimize_index functionality to compact vocab.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs({"doc1": "hello world"}) - bm25._remove_doc("doc1") - - assert bm25.n_docs == 0 - assert len(bm25.vocab) > 0 - - await bm25.optimize_index() - assert bm25.vocab == {} - assert bm25.inverted_index == {} - - await bm25.close() - print("✓ test_optimize_index passed") - - asyncio.run(run()) - - -def test_optimize_index_with_docs(): - """Test optimize_index with remaining documents.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs( - { - "doc1": "hello world", - "doc2": "hello python", - }, - ) - - old_vocab = bm25.vocab.copy() - bm25._remove_doc("doc1") - - await bm25.optimize_index() - - assert bm25.n_docs == 1 - assert "doc2" in bm25.doc_meta - assert len(bm25.vocab) < len(old_vocab) - - results = await bm25.retrieve("hello", limit=1) - assert "doc2" in results - - await bm25.close() - print("✓ test_optimize_index_with_docs passed") - - asyncio.run(run()) - - -def test_persistence(): - """Test dump and load persistence.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - docs = { - "doc1": "hello world", - "doc2": "hello python", - "doc3": "programming language", - } - await bm25.add_docs(docs) - - old_vocab = bm25.vocab.copy() - old_doc_meta = {k: dict(v) for k, v in bm25.doc_meta.items()} - - await bm25.dump() - await bm25.close() - - bm25_new = await create_bm25() - - assert bm25_new.vocab == old_vocab - assert bm25_new.n_docs == 3 - for doc_id in old_doc_meta: - assert doc_id in bm25_new.doc_meta - - results = await bm25_new.retrieve("hello", limit=2) - assert "doc1" in results or "doc2" in results - - await bm25_new.close() - print("✓ test_persistence passed") - - asyncio.run(run()) - - -def test_custom_params(): - """Test custom k1 and b parameters.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25(k1=2.0, b=0.5) - - assert bm25.k1 == 2.0 - assert bm25.b == 0.5 - - await bm25.add_docs({"doc1": "test document"}) - results = await bm25.retrieve("test", limit=1) - assert "doc1" in results - - await bm25.close() - print("✓ test_custom_params passed") - - asyncio.run(run()) - - -def test_chinese_text(): - """Test with Chinese text.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = { - "doc1": "我爱北京天安门", - "doc2": "北京是中国的首都", - "doc3": "上海的天气很好", - } - await bm25.add_docs(docs) - - results = await bm25.retrieve("北", limit=2) - assert len(results) <= 2 - assert "doc1" in results or "doc2" in results - - await bm25.close() - print("✓ test_chinese_text passed") - - asyncio.run(run()) - - -def test_mixed_chinese_english(): - """Test with mixed Chinese and English text.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = { - "doc1": "Python 是一种编程语言", - "doc2": "Java 编程语言", - "doc3": "Python 数据分析", - } - await bm25.add_docs(docs) - - results = await bm25.retrieve("Python", limit=3) - assert len(results) > 0 - - results = await bm25.retrieve("编", limit=2) - assert len(results) > 0 - - await bm25.close() - print("✓ test_mixed_chinese_english passed") - - asyncio.run(run()) - - -def test_idf_cache(): - """Test IDF cache functionality.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs( - { - "doc1": "hello world", - "doc2": "hello python", - }, - ) - - token = "hello" - if token in bm25.vocab: - tid = bm25.vocab[token] - idf1 = bm25._get_idf(tid) - assert tid in bm25._idf_cache - idf2 = bm25._get_idf(tid) - assert idf1 == idf2 - - await bm25.close() - print("✓ test_idf_cache passed") - - asyncio.run(run()) - - -def test_avg_len(): - """Test average document length calculation.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - assert bm25.avg_len == 0.0 - - await bm25.add_docs({"doc1": "hello world python"}) - assert bm25.avg_len > 0 - - await bm25.add_docs({"doc2": "test"}) - new_avg = bm25.avg_len - assert new_avg > 0 - - await bm25.close() - print("✓ test_avg_len passed") - - asyncio.run(run()) - - -def test_score_ordering(): - """Test that results are ordered by score descending.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - docs = { - "doc1": "python python python", - "doc2": "python python", - "doc3": "python", - } - await bm25.add_docs(docs) - - results = await bm25.retrieve("python", limit=3) - scores = list(results.values()) - - for i in range(len(scores) - 1): - assert scores[i] >= scores[i + 1] - - await bm25.close() - print("✓ test_score_ordering passed") - - asyncio.run(run()) - - -def test_empty_doc(): - """Test adding empty document.""" - - async def run(): - with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir): - bm25 = await create_bm25() - - await bm25.add_docs({"doc1": ""}) - assert bm25.n_docs == 0 - - await bm25.add_docs({"doc2": " "}) - assert bm25.n_docs == 0 - - await bm25.close() - print("✓ test_empty_doc passed") - - asyncio.run(run()) - - -if __name__ == "__main__": - print("\n=== BM25Index Tests ===") - test_basic_init() - test_start_with_tokenizer() - test_add_single_doc() - test_add_multiple_docs() - test_retrieve_basic() - test_retrieve_with_limit() - test_retrieve_empty_query() - test_retrieve_empty_index() - test_update_doc() - test_remove_doc() - test_remove_nonexistent_doc() - test_clear() - test_optimize_index() - test_optimize_index_with_docs() - test_persistence() - test_custom_params() - test_chinese_text() - test_mixed_chinese_english() - test_idf_cache() - test_avg_len() - test_score_ordering() - test_empty_doc() - print("\n所有测试通过!") diff --git a/tests4/unittest/test_keyword_index.py b/tests4/unittest/test_keyword_index.py new file mode 100644 index 00000000..be374198 --- /dev/null +++ b/tests4/unittest/test_keyword_index.py @@ -0,0 +1,931 @@ +"""Tests for BaseKeywordIndex implementations (currently: BM25Index). + +Covers full lifecycle, CRUD, retrieval, persistence, optimize and — as the focus +of this file — Chinese / English / mixed-language behaviour driven by the +default RegexTokenizer (Chinese split per char, English words lowercased, +single-char ASCII words dropped). +""" + +# pylint: disable=protected-access + +import asyncio +import os +import tempfile +import warnings + +from reme4.components.keyword_index import BM25Index +from reme4.components.tokenizer import RegexTokenizer + +warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba") +warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources") + + +# --------------------------------------------------------------------------- # +# Helpers # +# --------------------------------------------------------------------------- # + + +class temp_chdir: + """Context manager to temporarily chdir into a path and restore on exit.""" + + def __init__(self, path): + self.path = path + self.old = None + + def __enter__(self): + self.old = os.getcwd() + os.chdir(self.path) + return self + + def __exit__(self, *exc): + os.chdir(self.old) + + +async def create_bm25(k1: float = 1.5, b: float = 0.75, + filter_stopwords: bool = False) -> BM25Index: + """Create and start a BM25Index in cwd with a non-filtering RegexTokenizer. + + Stopword filtering is off so short test words ("hello", "我", "的") survive. + """ + bm25 = BM25Index(k1=k1, b=b) + tokenizer = RegexTokenizer(filter_stopwords=filter_stopwords) + bm25.tokenizer = tokenizer + bm25._owned.append(tokenizer) + await bm25.start() + return bm25 + + +def run(coro): + """Tiny shorthand to avoid repeating asyncio.run wrappers.""" + return asyncio.run(coro) + + +# --------------------------------------------------------------------------- # +# Initialisation & lifecycle # +# --------------------------------------------------------------------------- # + + +def test_basic_init(): + """Default constructor produces empty BM25 state.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = BM25Index() + assert bm25.k1 == 1.5 + assert bm25.b == 0.75 + assert bm25.index_version == "v1" + assert bm25.vocab == {} + assert bm25.inverted_index == {} + assert bm25.doc_meta == {} + assert bm25.n_docs == 0 + assert bm25.total_len == 0 + assert bm25.avg_len == 0.0 + + run(go()) + + +def test_custom_params(): + """k1, b and index_version are honoured.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = BM25Index(k1=2.0, b=0.5, index_version="v2") + assert bm25.k1 == 2.0 + assert bm25.b == 0.5 + assert bm25.index_version == "v2" + + run(go()) + + +def test_index_file_raises_when_tokenizer_is_none(): + """index_file must raise when tokenizer is explicitly None.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = BM25Index() + bm25.tokenizer = None + try: + _ = bm25.index_file + except RuntimeError: + return + raise AssertionError("expected RuntimeError when tokenizer is None") + + run(go()) + + +def test_start_close_lifecycle(): + """start/close toggles is_started and runs underlying tokenizer.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + assert bm25.is_started + assert bm25.tokenizer is not None + await bm25.close() + assert not bm25.is_started + + run(go()) + + +def test_index_file_path_includes_tokenizer_and_version(): + """index_file path embeds tokenizer name + index_version.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + path = str(bm25.index_file) + assert "bm25_regex_v1.pkl" in path + await bm25.close() + + run(go()) + + +# --------------------------------------------------------------------------- # +# add_docs / delete_docs / update # +# --------------------------------------------------------------------------- # + + +def test_add_empty_dict_noop(): + """Adding an empty dict must not touch state.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({}) + assert bm25.n_docs == 0 + assert bm25.vocab == {} + await bm25.close() + + run(go()) + + +def test_add_single_doc(): + """A single doc populates length, vocab and metadata.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello world"}) + assert bm25.n_docs == 1 + assert bm25.total_len == 2 # 'hello', 'world' + assert bm25.avg_len == 2.0 + assert set(bm25.vocab) == {"hello", "world"} + assert "d1" in bm25.doc_meta + assert bm25.doc_meta["d1"]["len"] == 2 + await bm25.close() + + run(go()) + + +def test_add_multiple_docs_and_inverted_index(): + """Inverted index lists postings for every term across multiple docs.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "hello world", + "d2": "hello python", + "d3": "world python", + }) + assert bm25.n_docs == 3 + inv = bm25.inverted_index + tid_hello = bm25.vocab["hello"] + tid_world = bm25.vocab["world"] + tid_python = bm25.vocab["python"] + assert set(inv[tid_hello]) == {"d1", "d2"} + assert set(inv[tid_world]) == {"d1", "d3"} + assert set(inv[tid_python]) == {"d2", "d3"} + await bm25.close() + + run(go()) + + +def test_add_doc_empty_or_whitespace_is_skipped(): + """Empty / whitespace-only content yields no tokens and is silently dropped.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "", "d2": " ", "d3": "\n\t"}) + assert bm25.n_docs == 0 + await bm25.close() + + run(go()) + + +def test_update_existing_doc_swaps_content(): + """Re-adding same doc_id replaces tokens; old terms no longer match it.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello world python"}) + old_len = bm25.total_len + await bm25.add_docs({"d1": "java"}) + assert bm25.n_docs == 1 + assert bm25.total_len != old_len + assert bm25.doc_meta["d1"]["len"] == 1 + + # Old term must no longer return d1. + assert "d1" not in await bm25.retrieve("hello", limit=5) + # New term does. + assert "d1" in await bm25.retrieve("java", limit=5) + await bm25.close() + + run(go()) + + +def test_delete_single_doc(): + """delete_docs removes a single doc from retrieval and meta.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello world", "d2": "hello python"}) + assert bm25.n_docs == 2 + + await bm25.delete_docs(["d1"]) + assert bm25.n_docs == 1 + assert "d1" not in bm25.doc_meta + assert "d2" in bm25.doc_meta + + results = await bm25.retrieve("hello", limit=2) + assert "d1" not in results + assert "d2" in results + await bm25.close() + + run(go()) + + +def test_delete_multiple_docs(): + """delete_docs handles a batch list.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({f"d{i}": "hello world" for i in range(5)}) + assert bm25.n_docs == 5 + + await bm25.delete_docs(["d0", "d2", "d4"]) + assert bm25.n_docs == 2 + assert set(bm25.doc_meta) == {"d1", "d3"} + await bm25.close() + + run(go()) + + +def test_delete_nonexistent_is_noop(): + """Deleting unknown doc_ids must not raise.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello"}) + await bm25.delete_docs(["nope", "still_nope"]) + assert bm25.n_docs == 1 + await bm25.close() + + run(go()) + + +def test_re_add_after_delete(): + """Adding a doc_id back after deletion yields a fresh idx and is retrievable.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello"}) + await bm25.delete_docs(["d1"]) + assert bm25.n_docs == 0 + await bm25.add_docs({"d1": "world"}) + assert bm25.n_docs == 1 + assert "d1" in await bm25.retrieve("world", limit=1) + await bm25.close() + + run(go()) + + +# --------------------------------------------------------------------------- # +# Retrieval # +# --------------------------------------------------------------------------- # + + +def test_retrieve_empty_index(): + """Retrieving from an empty index returns {}.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + assert await bm25.retrieve("python", limit=3) == {} + await bm25.close() + + run(go()) + + +def test_retrieve_empty_or_unknown_query(): + """Empty queries and out-of-vocab queries both return {}.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello world"}) + assert await bm25.retrieve("", limit=3) == {} + assert await bm25.retrieve("zzzunknownxyz", limit=3) == {} + await bm25.close() + + run(go()) + + +def test_retrieve_limit_caps_results(): + """retrieve honours `limit`.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({f"d{i}": f"python lang {i}" for i in range(10)}) + assert len(await bm25.retrieve("python", limit=3)) == 3 + assert len(await bm25.retrieve("python", limit=5)) == 5 + # limit greater than matches: bounded by positive matches. + assert len(await bm25.retrieve("python", limit=99)) == 10 + await bm25.close() + + run(go()) + + +def test_retrieve_score_ordering_by_tf(): + """A doc with higher term frequency for the query token outranks others.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "high": "python python python", + "mid": "python python other", + "low": "python alpha beta", + }) + results = await bm25.retrieve("python", limit=3) + assert list(results.keys()) == ["high", "mid", "low"] + scores = list(results.values()) + assert scores[0] >= scores[1] >= scores[2] + await bm25.close() + + run(go()) + + +def test_retrieve_idf_favours_rare_terms(): + """In a query of {common, rare}, the doc containing the rare term wins.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + # 'common' appears everywhere → low IDF. + # 'rare' appears in only one doc → high IDF. + docs = {f"d{i}": "common filler text" for i in range(10)} + docs["target"] = "common rare term" + await bm25.add_docs(docs) + + results = await bm25.retrieve("common rare", limit=3) + assert next(iter(results)) == "target" + await bm25.close() + + run(go()) + + +def test_retrieve_length_normalization(): + """With b=0.75 (default), a much longer doc with same tf scores lower.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "short": "python", + "long": "python " + " ".join(f"w{i}" for i in range(50)), + }) + results = await bm25.retrieve("python", limit=2) + assert results["short"] > results["long"] + await bm25.close() + + run(go()) + + +def test_retrieve_duplicate_query_tokens_dont_double_count(): + """Repeating the same query token should not boost its contribution.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "python rocks"}) + once = await bm25.retrieve("python", limit=1) + many = await bm25.retrieve("python python python", limit=1) + assert once["d1"] == many["d1"] + await bm25.close() + + run(go()) + + +# --------------------------------------------------------------------------- # +# Chinese / English / mixed-language behaviour (focus) # +# --------------------------------------------------------------------------- # + + +def test_chinese_only_corpus(): + """Pure Chinese corpus indexes per-character and retrieves correctly.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "我爱北京天安门", + "d2": "北京是中国的首都", + "d3": "上海的天气很好", + }) + # Regex tokenizer splits Chinese per character. + assert "北" in bm25.vocab + assert "京" in bm25.vocab + + # Query "北京" → two tokens, both d1 and d2 match; d3 does not. + results = await bm25.retrieve("北京", limit=3) + assert set(results) == {"d1", "d2"} + await bm25.close() + + run(go()) + + +def test_english_only_corpus_is_lowercased(): + """English tokens are lowercased so case-insensitive retrieval works.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "Python Programming Language", + "d2": "Java Programming Language", + }) + assert "python" in bm25.vocab + assert "Python" not in bm25.vocab + + r_upper = await bm25.retrieve("PYTHON", limit=2) + r_lower = await bm25.retrieve("python", limit=2) + assert r_upper == r_lower + assert "d1" in r_upper + await bm25.close() + + run(go()) + + +def test_single_char_english_dropped(): + """RegexTokenizer's \\w\\w+ pattern drops single-letter ASCII tokens like 'I'.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "I love Beijing"}) + # 'I' must not appear; 'love' and 'beijing' must. + assert "i" not in bm25.vocab + assert "love" in bm25.vocab + assert "beijing" in bm25.vocab + # Querying with just "I" returns nothing. + assert await bm25.retrieve("I", limit=1) == {} + await bm25.close() + + run(go()) + + +def test_mixed_doc_chinese_query_matches(): + """A Chinese query hits docs containing those Chinese chars even when mixed.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "Python 是一种编程语言", + "d2": "Java 编程语言", + "d3": "Python 数据分析", + }) + results = await bm25.retrieve("编程", limit=3) + assert set(results) >= {"d1", "d2"} + assert "d3" not in results + await bm25.close() + + run(go()) + + +def test_mixed_doc_english_query_matches(): + """An English query hits the right mixed-language docs.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "Python 是一种编程语言", + "d2": "Java 编程语言", + "d3": "Python 数据分析", + }) + results = await bm25.retrieve("python", limit=3) + assert set(results) == {"d1", "d3"} + assert "d2" not in results + await bm25.close() + + run(go()) + + +def test_mixed_query_combines_chinese_and_english_signal(): + """A query mixing English and Chinese aggregates IDF×tf contributions.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "py_cn": "Python 编程", # matches both 'python' and '编','程' + "py_only": "Python tutorial", # matches only 'python' + "cn_only": "编程入门", # matches only '编','程' + # Avoid Chinese chars that the query splits into ('编','程') — '教程' would + # leak '程' into 'other' and pollute IDF, so use unrelated chars only. + "other": "Java 教学", + }) + results = await bm25.retrieve("Python 编程", limit=4) + # py_cn should rank highest because it matches both branches. + assert next(iter(results)) == "py_cn" + # 'other' should not appear. + assert "other" not in results + # Both unimodal matches should still appear. + assert "py_only" in results and "cn_only" in results + await bm25.close() + + run(go()) + + +def test_mixed_doc_more_matches_outrank_fewer(): + """Doc covering more query tokens (Chinese+English) outranks partial matches.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "full": "machine learning 机器 学习", + "en_only": "machine learning algorithm", + "cn_only": "机器 学习 算法", + }) + results = await bm25.retrieve("machine 机器", limit=3) + # full has both English and Chinese hits → highest score. + assert next(iter(results)) == "full" + await bm25.close() + + run(go()) + + +def test_unicode_word_with_digits_preserved(): + """Alphanumeric tokens like 'iphone15' stay whole; trailing Chinese still split.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "iPhone15 Pro 售价 9999 元", + "d2": "Android 旗舰 999 元", + }) + assert "iphone15" in bm25.vocab + assert "9999" in bm25.vocab + assert "元" in bm25.vocab + + r1 = await bm25.retrieve("iphone15", limit=2) + assert list(r1) == ["d1"] + r2 = await bm25.retrieve("元", limit=2) + assert set(r2) == {"d1", "d2"} + await bm25.close() + + run(go()) + + +def test_chinese_punctuation_ignored(): + """CJK punctuation should not produce tokens.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "你好,世界!这是 Python。"}) + for sym in [",", "!", "。"]: + assert sym not in bm25.vocab + assert "你" in bm25.vocab + assert "python" in bm25.vocab + await bm25.close() + + run(go()) + + +def test_mixed_persistence_roundtrip(): + """A mixed-language index round-trips through dump/load with identical retrieval.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "Python 编程语言", + "d2": "Java 编程", + "d3": "数据分析 with Python", + }) + before = await bm25.retrieve("Python 编程", limit=3) + await bm25.close() # close triggers dump + + bm25_2 = await create_bm25() # start triggers load + assert bm25_2.n_docs == 3 + after = await bm25_2.retrieve("Python 编程", limit=3) + assert before == after + await bm25_2.close() + + run(go()) + + +# --------------------------------------------------------------------------- # +# Persistence # +# --------------------------------------------------------------------------- # + + +def test_dump_load_roundtrip_preserves_state(): + """dump → fresh instance → load reconstructs vocab, postings and params.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25(k1=2.0, b=0.4) + await bm25.add_docs({ + "d1": "hello world", + "d2": "hello python", + "d3": "programming language", + }) + old_vocab = dict(bm25.vocab) + old_meta = {k: dict(v) for k, v in bm25.doc_meta.items()} + await bm25.dump() + await bm25.close() + + bm25_2 = await create_bm25() # default k1/b — load must overwrite + assert bm25_2.vocab == old_vocab + assert bm25_2.n_docs == 3 + assert set(bm25_2.doc_meta) == set(old_meta) + assert bm25_2.k1 == 2.0 + assert bm25_2.b == 0.4 + await bm25_2.close() + + run(go()) + + +def test_load_missing_file_keeps_empty_state(): + """Calling load() with no file on disk is a no-op.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + # No add_docs, nothing persisted. + assert not bm25.index_file.exists() + await bm25.load() + assert bm25.n_docs == 0 + await bm25.close() + + run(go()) + + +def test_load_corrupt_file_resets_index(): + """A corrupt pickle on disk is reported, deleted, and the index is cleared.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello"}) + await bm25.dump() + + # Corrupt the file. + bm25.index_file.write_bytes(b"not a pickle") + await bm25.load() + assert bm25.n_docs == 0 + assert bm25.vocab == {} + assert not bm25.index_file.exists() + await bm25.close() + + run(go()) + + +# --------------------------------------------------------------------------- # +# clear / optimize / reset_index # +# --------------------------------------------------------------------------- # + + +def test_clear_wipes_everything(): + """clear() empties state and removes the index file.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello", "d2": "world"}) + await bm25.dump() + assert bm25.index_file.exists() + + await bm25.clear() + assert bm25.n_docs == 0 + assert bm25.vocab == {} + assert bm25.inverted_index == {} + assert bm25.doc_meta == {} + assert bm25.total_len == 0 + assert bm25._idf_cache == {} + assert not bm25.index_file.exists() + await bm25.close() + + run(go()) + + +def test_optimize_drops_deleted_only_terms(): + """After deleting, optimize_index drops vocab entries that no live doc uses.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({ + "d1": "alpha unique_to_d1", + "d2": "alpha beta", + }) + assert "unique_to_d1" in bm25.vocab + + await bm25.delete_docs(["d1"]) + await bm25.optimize_index() + + assert bm25.n_docs == 1 + assert "d2" in bm25.doc_meta + # Term that only existed in d1 is gone. + assert "unique_to_d1" not in bm25.vocab + # Shared/own terms of d2 survive. + assert "alpha" in bm25.vocab and "beta" in bm25.vocab + # Retrieval still works correctly. + assert "d2" in await bm25.retrieve("alpha", limit=1) + await bm25.close() + + run(go()) + + +def test_optimize_when_all_deleted_clears_index(): + """optimize_index on a fully-deleted state collapses to empty index.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello world"}) + await bm25.delete_docs(["d1"]) + await bm25.optimize_index() + assert bm25.n_docs == 0 + assert bm25.vocab == {} + assert bm25.inverted_index == {} + await bm25.close() + + run(go()) + + +def test_optimize_noop_when_no_deletions(): + """With nothing deleted, optimize_index leaves state intact.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello world"}) + vocab_before = dict(bm25.vocab) + await bm25.optimize_index() + assert bm25.vocab == vocab_before + assert bm25.n_docs == 1 + await bm25.close() + + run(go()) + + +def test_reset_index_replaces_all_docs(): + """reset_index (inherited from base) wipes and re-adds in one call.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "old content"}) + await bm25.reset_index({"d2": "new content"}) + assert bm25.n_docs == 1 + assert "d2" in bm25.doc_meta + assert "d1" not in bm25.doc_meta + await bm25.close() + + run(go()) + + +# --------------------------------------------------------------------------- # +# Internal invariants # +# --------------------------------------------------------------------------- # + + +def test_idf_cache_populates_and_invalidates(): + """_get_idf caches results; add/delete clear the cache.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "hello world", "d2": "hello python"}) + tid_hello = bm25.vocab["hello"] + idf1 = bm25._get_idf(tid_hello) + assert tid_hello in bm25._idf_cache + assert bm25._get_idf(tid_hello) == idf1 + + # Mutating the index must invalidate the cache. + await bm25.add_docs({"d3": "hello there"}) + assert bm25._idf_cache == {} + + await bm25.delete_docs(["d1"]) + # delete_docs also clears cache; populate again then trigger via add. + _ = bm25._get_idf(bm25.vocab["hello"]) + assert bm25._idf_cache # non-empty now + await bm25.add_docs({"d4": "x y z"}) + assert bm25._idf_cache == {} + await bm25.close() + + run(go()) + + +def test_avg_len_tracks_live_docs_only(): + """avg_len excludes deleted docs. + + Note: RegexTokenizer's `\\w\\w+` pattern drops 1-letter words, so we + use multi-letter tokens to keep length math predictable. + """ + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + assert bm25.avg_len == 0.0 + await bm25.add_docs({"d1": "alpha beta gamma delta"}) # 4 tokens + await bm25.add_docs({"d2": "alpha beta"}) # 2 tokens + assert bm25.avg_len == 3.0 + + await bm25.delete_docs(["d1"]) + assert bm25.avg_len == 2.0 + await bm25.close() + + run(go()) + + +def test_deleted_docs_excluded_from_scoring(): + """A deleted doc must score 0 and never appear in retrieve().""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "python", "d2": "python", "d3": "python"}) + await bm25.delete_docs(["d2"]) + + results = await bm25.retrieve("python", limit=10) + assert set(results) == {"d1", "d3"} + assert all(s > 0 for s in results.values()) + await bm25.close() + + run(go()) + + +def test_inverted_index_hides_deleted_postings(): + """inverted_index view skips postings whose doc is deleted.""" + + async def go(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + bm25 = await create_bm25() + await bm25.add_docs({"d1": "alpha", "d2": "alpha beta"}) + tid_alpha = bm25.vocab["alpha"] + await bm25.delete_docs(["d1"]) + + inv = bm25.inverted_index + # 'alpha' posting now contains only the live doc. + assert tid_alpha in inv + assert set(inv[tid_alpha]) == {"d2"} + await bm25.close() + + run(go()) + + +# --------------------------------------------------------------------------- # +# Manual runner # +# --------------------------------------------------------------------------- # + + +if __name__ == "__main__": + import inspect + import sys + + mod = sys.modules[__name__] + tests = [ + (name, obj) for name, obj in inspect.getmembers(mod, inspect.isfunction) + if name.startswith("test_") + ] + print(f"\n=== BaseKeywordIndex / BM25Index Tests ({len(tests)}) ===\n") + failed = [] + for name, fn in tests: + try: + fn() + print(f"✓ {name}") + except Exception as exc: # noqa: BLE001 + print(f"✗ {name}: {exc!r}") + failed.append(name) + print() + if failed: + print(f"FAILED: {len(failed)} / {len(tests)}") + for n in failed: + print(f" - {n}") + sys.exit(1) + print(f"所有 {len(tests)} 项测试通过!")