ReMe/tests/unit/test_keyword_index.py
jinliyl e7d44f6f3b
Some checks are pending
Pre-commit / run (ubuntu-latest) (push) Waiting to run
Tests ReMe / Unit Tests - py3.11 (push) Waiting to run
Tests ReMe / Unit Tests - py3.12 (push) Waiting to run
Tests ReMe / Unit Tests - py3.13 (push) Waiting to run
Windows Smoke / CLI smoke - py3.11 (push) Waiting to run
refactor(agent): unify agent subprocess env, sessions, skills, and MCP/service jobs (#382)
* feat(config): add environment variable configuration for agent subprocesses

- Add environment field to ApplicationConfig to store variables for agent subprocesses
- Remove dynamic loading of .env files in agent wrappers
- Introduce subprocess_environment property in base agent wrapper
- Pass application-level environment variables to Claude Code and Codex agents
- Load environment variables once at startup and pass to ReMe application
- Remove dependency on load_env utility in agent wrapper implementations
- Update tests to use configured environment instead of dynamic loading
- Remove unused environment loading utilities and related test cases

* refactor(mcp): remove channel notification system and related components

- Removed channel notification step implementation
- Removed claim channel step implementation
- Removed ChannelSink class from MCP service
- Removed channel-related documentation from AGENTS.md
- Removed channel instruction text from MCP service
- Removed all channel-related tests
- Updated application context metadata comment to remove channel sink reference
- Removed channel module initialization and imports

* feat(service): add job whitelisting capability to BaseService

- Add optional jobs parameter to BaseService.__init__ to configure job whitelist
- Store jobs as set in self.jobs attribute for efficient lookup operations
- Modify add_jobs method to filter jobs based on whitelist configuration
- Update documentation in both English and Chinese to describe new feature
- Add comprehensive unit tests for job whitelisting behavior
- Implement flowchart update showing new filtering logic
- Preserve existing enable_serve flag behavior alongside new whitelisting

* refactor(service): enhance service job validation and MCP tool injection

- Add strict validation for service jobs whitelist with detailed error messages
- Implement injected job arguments support for MCP services with conflict detection
- Add tool error handling for unsuccessful responses in MCP services
- Remove duplicate job names in Codex agent wrapper using dict.fromkeys
- Update MCP server argument format from single JSON array to repeated --job flags
- Add comprehensive test coverage for job injection and error handling scenarios
- Update documentation to reflect service job validation and MCP features
- Ensure application cleanup occurs even when service lifespan encounters errors

* feat(agent): update skill handling to preserve existing Claude skills

- Change skills parameter processing to use 'all' instead of filtered list
- Add logic to select project skills without restricting Claude's existing skills
- Update variable naming from 'skills' to 'selected_skills' for clarity
- Modify application context metadata documentation to clarify in-memory state usage
- Add test case to verify configured skills are added without filtering existing skills
- Update internal skill directory handling to use renamed variable consistently

* refactor(agent): restructure agent wrapper components and session storage

- Move CcFileSessionStore to separate module for better organization
- Add SDK package version logging in base agent wrapper
- Update Claude Code agent to use new session store structure with project keys
- Refactor Claude Code agent wrapper to use proper type hints and SDK integration
- Add support for server tool use events in Claude Code message processing
- Improve error handling and resource cleanup in streaming operations
- Update Codex agent wrapper with proper type annotations and configuration
- Remove deprecated system prompt mode handling from Claude Code wrapper
- Fix session path construction for Claude Code transcript storage
- Update dependency injection and configuration handling patterns

* fix(cc_agent_wrapper): resolve Claude Code SDK integration issues

- Added dataclass import and created _BlockState for content block metadata tracking
- Implemented proper MCP server name constant and tool context ID validation
- Fixed tool_context_id injection to prevent duplicate assignment errors
- Resolved skills parameter handling in build_options method
- Enhanced job tools integration with MCP servers mapping validation
- Replaced deprecated block_ids/block_types/tool_call_names with block_states dict
- Updated message_delta to emit USAGE chunks instead of REPLY_END
- Fixed stream result handling to ensure proper REPLY_END emission
- Improved error handling for session mirror failures and rate limits
- Added proper cleanup for expected trailing errors in streams
- Refactored Codex agent wrapper initialization and configuration management
- Removed obsolete system_prompt_mode from default config
- Enhanced test coverage for new block state and error handling features
- Fixed async generator handling with aclosing context manager
- Improved chunk type mapping for Claude Code SDK events

* refactor(tests): remove demo config tests from config parser test suite

- Removed test_demo_config_registers_llm_jobs function and its assertions
- Eliminated verification of LLM demo job configurations
- Removed checks for agent wrapper component settings
- Deleted assertions for model configurations and parameters
- Cleaned up deprecated test cases related to demo config parsing

* refactor(evolve): simplify Claude Code session store path structure

- Removed redundant project key subdirectory from session link generation
- Updated CcFileSessionStore initialization to use direct session directory path
- Maintained existing session layout compatibility for backward compatibility
- Added unit tests to verify session persistence behavior with existing transcripts
- Ensured UUID-based session files remain accessible at expected locations
- Preserved existing session directory structure without additional nesting

* refactor(agent): defer optional Codex SDK imports until first use

- Moved openai-codex imports inside functions to avoid mandatory dependencies
- Added TYPE_CHECKING guard for development time type checking only
- Implemented lazy loading mechanism with _get_async_codex_class function
- Updated AsyncCodex initialization to occur on demand rather than at module level
- Maintained backward compatibility while improving import performance
- Added test case to verify package import works without optional Codex SDK
- Updated agentscope dependency to version 2.0.4.post1 in pyproject.toml

* test(embedded): add compatibility tests for in-process ReMe embedding

- Add test suite for QwenPaw-style embedded configurations
- Verify optional defaults remain preserved in embedded configs
- Ensure in-process application API stays compatible
- Test model injection and lifecycle management compatibility
- Remove obsolete hermes agent plugin tests
- Update CLI import test to cover multiple optional SDKs
- Block claude_agent_sdk and openai_codex during import testing
2026-07-20 23:52:14 +08:00

1045 lines
35 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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 reme.components.keyword_index import BM25Index
from reme.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_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_BM25Index_regex_" in path
assert path.endswith("_v1.pkl")
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 bm25.document_ids == {"d1"}
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())
def test_index_file_isolated_by_component_name():
"""Different BM25Index names must not share one persisted pickle."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
alpha = BM25Index(name="alpha")
alpha_tokenizer = RegexTokenizer(filter_stopwords=False)
alpha.tokenizer = alpha_tokenizer
alpha._owned.append(alpha_tokenizer)
await alpha.start()
beta = BM25Index(name="beta")
beta_tokenizer = RegexTokenizer(filter_stopwords=False)
beta.tokenizer = beta_tokenizer
beta._owned.append(beta_tokenizer)
await beta.start()
assert alpha.index_file != beta.index_file
await alpha.add_docs({"d1": "alpha only"})
await alpha.dump()
await alpha.close()
assert beta.n_docs == 0
assert await beta.retrieve("alpha", limit=1) == {}
await beta.close()
run(go())
def test_index_file_isolated_by_tokenizer_config():
"""Tokenizer settings that affect tokens must map to different index files."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
unfiltered = BM25Index()
unfiltered_tokenizer = RegexTokenizer(filter_stopwords=False)
unfiltered.tokenizer = unfiltered_tokenizer
unfiltered._owned.append(unfiltered_tokenizer)
await unfiltered.start()
filtered = BM25Index()
filtered_tokenizer = RegexTokenizer(filter_stopwords=True)
filtered.tokenizer = filtered_tokenizer
filtered._owned.append(filtered_tokenizer)
await filtered.start()
assert unfiltered.index_file != filtered.index_file
await unfiltered.close()
await filtered.close()
run(go())
def test_tokenizer_fingerprint_ignores_stopwords_absolute_path():
"""Same stopwords content should not fork indexes by install path."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
stopwords_a = os.path.join(tmp, "a", "stopwords")
stopwords_b = os.path.join(tmp, "b", "stopwords")
os.makedirs(os.path.dirname(stopwords_a))
os.makedirs(os.path.dirname(stopwords_b))
with open(stopwords_a, "w", encoding="utf-8") as f:
f.write("alpha\nbeta\n")
with open(stopwords_b, "w", encoding="utf-8") as f:
f.write("alpha\nbeta\n")
first = BM25Index()
first_tokenizer = RegexTokenizer(filter_stopwords=True, stopwords_path=stopwords_a)
first.tokenizer = first_tokenizer
first._owned.append(first_tokenizer)
await first.start()
second = BM25Index()
second_tokenizer = RegexTokenizer(filter_stopwords=True, stopwords_path=stopwords_b)
second.tokenizer = second_tokenizer
second._owned.append(second_tokenizer)
await second.start()
assert first._tokenizer_config()["stopwords_sha256"] == second._tokenizer_config()["stopwords_sha256"]
assert "stopwords_path" not in first._tokenizer_config()
assert first._tokenizer_fingerprint() == second._tokenizer_fingerprint()
assert first.index_file == second.index_file
with open(stopwords_b, "w", encoding="utf-8") as f:
f.write("alpha\ngamma\n")
assert first._tokenizer_fingerprint() != second._tokenizer_fingerprint()
await first.close()
await second.close()
run(go())
def test_dump_failure_is_not_silent():
"""A failed write must be observable by callers."""
async def go():
from unittest.mock import patch
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello"})
with patch("builtins.open", side_effect=OSError("disk full")):
try:
await bm25.dump()
except OSError:
pass
else:
raise AssertionError("expected dump() to raise OSError")
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()
assert not bm25.index_file.exists()
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)} 项测试通过!")