mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-30 01:52:29 +00:00
up
This commit is contained in:
parent
ca40f26455
commit
17919ebffa
5 changed files with 1041 additions and 667 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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所有测试通过!")
|
||||
931
tests4/unittest/test_keyword_index.py
Normal file
931
tests4/unittest/test_keyword_index.py
Normal file
|
|
@ -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)} 项测试通过!")
|
||||
Loading…
Add table
Reference in a new issue