feat: implement local code scout, training pipeline, and MCP tools

This commit is contained in:
emil28092005
2026-09-16 03:57:09 +03:00
parent 2e5ab98c56
commit ba24db2be5
33 changed files with 4347 additions and 27 deletions
+52
View File
@@ -0,0 +1,52 @@
import json
import pytest
from micro_scout.data import audit_splits, prepare
from micro_scout.io import write_jsonl
def test_overlap_audit_rejects_renamed_clones(tmp_path):
for split in ("train", "validation", "test"):
write_jsonl(
tmp_path / f"{split}.jsonl",
[
{
"id": split,
"repo": split,
"code_hash": split,
"query_hash": split,
"structural_hash": "same-shape",
}
],
)
with pytest.raises(ValueError, match="structural_hash"):
audit_splits(tmp_path)
def test_overlap_audit_passes_disjoint_records(tmp_path):
for split in ("train", "validation", "test"):
write_jsonl(
tmp_path / f"{split}.jsonl",
[{key: split for key in ("id", "repo", "code_hash", "query_hash", "structural_hash")}],
)
assert set(audit_splits(tmp_path).values()) == {0}
def test_dataset_failure_does_not_publish_partial_output(tmp_path):
pytest.importorskip("pyarrow")
output = tmp_path / "prepared"
with pytest.raises(FileNotFoundError):
prepare(tmp_path / "missing", output, {"train": 10, "validation": 10, "test": 10})
assert not output.exists()
assert not list(tmp_path.glob(".prepare-*"))
def test_json_writer_rejects_non_finite_metrics(tmp_path):
from micro_scout.io import atomic_json
destination = tmp_path / "metrics.json"
atomic_json(destination, {"score": 1.0})
with pytest.raises(ValueError):
atomic_json(destination, {"score": float("nan")})
assert json.loads(destination.read_text()) == {"score": 1.0}
+116
View File
@@ -0,0 +1,116 @@
import json
import numpy as np
import pytest
torch = pytest.importorskip("torch")
transformers = pytest.importorskip("transformers")
from micro_scout.encoder import Encoder # noqa: E402
from micro_scout.io import atomic_json, write_jsonl # noqa: E402
from micro_scout.train import train # noqa: E402
@pytest.fixture
def tiny_model(tmp_path):
"""Offline random BERT fixture tests plumbing, never used for reported quality."""
from transformers import BertConfig, BertModel, BertTokenizerFast
root = tmp_path / "tiny-model"
root.mkdir()
words = [
"[PAD]",
"[UNK]",
"[CLS]",
"[SEP]",
"[MASK]",
"read",
"write",
"file",
"sort",
"numbers",
"return",
"open",
"def",
"parse",
"text",
"a",
"b",
"(",
")",
":",
]
(root / "vocab.txt").write_text("\n".join(words))
tokenizer = BertTokenizerFast(vocab_file=str(root / "vocab.txt"))
tokenizer.save_pretrained(root)
torch.manual_seed(17)
model = BertModel(
BertConfig(
vocab_size=len(words),
hidden_size=16,
num_hidden_layers=1,
num_attention_heads=2,
intermediate_size=32,
max_position_embeddings=64,
)
)
model.save_pretrained(root)
return root
def test_embedding_save_reload_equivalence(tiny_model, tmp_path):
encoder = Encoder(str(tiny_model), max_length=32, query_length=16)
texts = ["read file", "sort numbers"]
vectors = encoder.encode(texts, query=True, batch_size=1)
assert vectors.shape == (2, 16)
np.testing.assert_allclose(np.linalg.norm(vectors, axis=1), 1, atol=1e-6)
saved = tmp_path / "saved"
encoder.save(saved)
loaded = Encoder(str(saved))
np.testing.assert_allclose(loaded.encode(texts, query=True), vectors, atol=1e-6)
assert encoder.fingerprint == loaded.fingerprint
def test_training_updates_weights_and_can_resume(tiny_model, tmp_path):
data = tmp_path / "data"
data.mkdir()
rows = [
{"query": "read file", "code": "def read ( ) : return open ( )", "repo": "train"},
{"query": "sort numbers", "code": "def sort ( a ) : return numbers", "repo": "train"},
{"query": "write text", "code": "def write ( text ) : return text", "repo": "train"},
{"query": "parse file", "code": "def parse ( file ) : return file", "repo": "train"},
]
write_jsonl(data / "train.jsonl", rows)
write_jsonl(data / "validation.jsonl", [{**r, "repo": "validation"} for r in rows[:2]])
atomic_json(data / "manifest.json", {"fixture": True})
config = {
"base_model": str(tiny_model),
"revision": None,
"max_length": 32,
"query_length": 16,
"batch_size": 2,
"epochs": 1,
"learning_rate": 0.001,
"weight_decay": 0.01,
"temperature": 0.05,
"warmup_ratio": 0,
"eval_every": 1,
"seed": 17,
"threads": 1,
"max_minutes": 1,
"mixed_precision": False,
}
before = Encoder(str(tiny_model), max_length=32, query_length=16)
output = tmp_path / "run"
result = train(data, output, config, "cpu")
assert result["optimizer_steps"] == 2
after = Encoder(str(output / "last"))
assert any(
not torch.equal(a, b)
for a, b in zip(before.model.parameters(), after.model.parameters(), strict=True)
)
resumed = train(data, output, config, "cpu", output / "last")
assert resumed["optimizer_steps"] == 2
state = torch.load(output / "last/training_state.pt", weights_only=True)
assert state["step"] == 2
assert json.loads((output / "result.json").read_text())["test_set_used_for_selection"] is False
+71
View File
@@ -0,0 +1,71 @@
import asyncio
import os
import sys
import pytest
pytest.importorskip("mcp")
from mcp import ClientSession, StdioServerParameters # noqa: E402
from mcp.client.stdio import stdio_client # noqa: E402
from micro_scout.index import build_index # noqa: E402
def test_real_stdio_tool_roundtrip_and_stale_read(tmp_path):
root = tmp_path / "repo"
root.mkdir()
source = root / "reader.py"
source.write_text("def read_file(path):\n return open(path).read()\n")
index = tmp_path / "index.sqlite"
build_index(root, index)
async def roundtrip():
parameters = StdioServerParameters(
command=sys.executable,
args=[
"-m",
"micro_scout",
"serve",
"--index",
str(index),
"--trace",
str(tmp_path / "trace.jsonl"),
],
env=dict(os.environ),
)
async with (
stdio_client(parameters) as (reader, writer),
ClientSession(reader, writer) as session,
):
await session.initialize()
tools = await session.list_tools()
assert {t.name for t in tools.tools} == {
"scout_search",
"scout_read",
"scout_status",
"scout_feedback",
}
result = await session.call_tool("scout_search", {"query": "read file"})
assert not result.isError
payload = result.structuredContent
hit = payload["results"][0]
assert hit["verified"] and hit["path"] == "reader.py"
read = await session.call_tool("scout_read", {"symbol_id": hit["id"]})
assert not read.isError
feedback = await session.call_tool(
"scout_feedback",
{
"request_id": payload["request_id"],
"useful_ids": [hit["id"]],
"outcome": "helpful",
},
)
assert feedback.structuredContent["weights_updated"] is False
source.write_text("# changed after indexing\n")
stale = await session.call_tool("scout_read", {"symbol_id": hit["id"]})
assert stale.isError
status = await session.call_tool("scout_status")
assert not status.isError
asyncio.run(asyncio.wait_for(roundtrip(), timeout=30))
+239
View File
@@ -0,0 +1,239 @@
import json
import subprocess
import numpy as np
import pytest
from micro_scout.index import Index, build_index
from micro_scout.lexical import BM25, reciprocal_rank_fusion
from micro_scout.metrics import ranks_from_scores, retrieval_metrics
from micro_scout.scout import Scout, StaleReferenceError
from micro_scout.symbols import build_edges, parse_source, source_paths
@pytest.fixture
def repository(tmp_path):
root = tmp_path / "repo"
root.mkdir()
(root / "files.py").write_text('''def read_file(path):
"""Read text from a file."""
return open(path).read()
def parse_file(path):
return read_file(path).splitlines()
class Writer:
def save(self, path, text):
with open(path, "w") as stream:
stream.write(text)
''')
(root / "numbers.py").write_text("def add_numbers(a, b):\n return a + b\n")
return root
def test_bm25_and_empty_corpus():
index = BM25(["read file content", "write file content", "sort numbers"])
assert np.argmax(index.score("sort numbers")) == 2
assert np.all(index.score("notpresent") == 0)
assert BM25([]).score("x").shape == (0,)
def test_stable_ranks_and_single_positive_metrics():
scores = np.array([[1, 1, 0], [3, 2, 1], [0, 1, 2]], dtype=float)
ranks = ranks_from_scores(scores)
assert ranks.tolist() == [1, 2, 1]
assert retrieval_metrics(ranks)["mrr"] == pytest.approx(5 / 6)
with pytest.raises(ValueError):
ranks_from_scores(np.array([[float("nan")]]))
def test_fusion_does_not_create_lexical_matches_for_zero_scores():
result = reciprocal_rank_fusion(np.zeros(3), np.array([0.2, 0.9, 0.1]))
assert np.argmax(result) == 1
assert result[1] == pytest.approx(0.5 / 61)
def test_ranges_decorators_nested_symbols_and_graph():
text = "@decorator\ndef outer():\n def inner():\n return 1\n return inner()\n"
symbols = parse_source("a.py", text)
assert [(s.name, s.start_line, s.end_line) for s in symbols] == [
("outer", 1, 5),
("outer.inner", 3, 4),
]
assert build_edges(symbols) == [
{"source": symbols[0].id, "target": symbols[1].id, "kind": "contains"}
]
def test_invalid_python_uses_line_chunks():
symbols = parse_source("broken.py", "def incomplete(\n unfinished", chunk_lines=1)
assert [s.kind for s in symbols] == ["chunk", "chunk"]
assert [s.start_line for s in symbols] == [1, 2]
def test_gitignore_and_symlinks_are_respected(repository, tmp_path):
subprocess.run(["git", "init", "-q", str(repository)], check=True)
(repository / ".gitignore").write_text("ignored.py\n")
(repository / "ignored.py").write_text("secret = 1")
(repository / ".hidden.py").write_text("secret = 2")
outside = tmp_path / "outside.py"
outside.write_text("secret = 3")
(repository / "linked.py").symlink_to(outside)
names = {p.name for p in source_paths(repository)}
assert names == {"files.py", "numbers.py"}
def test_search_returns_verified_references_and_budget(repository, tmp_path):
path = tmp_path / "index.sqlite"
build_index(repository, path)
scout = Scout(Index(path))
response = scout.search("read file", mode="lexical", max_chars=200, top_k=2)
assert response["results"]
assert response["returned_chars"] <= 200
assert response["returned_chars"] == sum(
len(r["content"]) for r in response["results"] + response["neighbors"]
)
for item in response["results"] + response["neighbors"]:
lines = (repository / item["path"]).read_text().splitlines()
assert item["content"] == "\n".join(lines[item["start_line"] - 1 : item["end_line"]])
assert item["verified"]
def test_changed_and_deleted_files_never_return_stale_content(repository, tmp_path):
path = tmp_path / "index.sqlite"
build_index(repository, path)
scout = Scout(Index(path))
symbol = next(s for s in scout.index.symbols if s.name == "read_file")
(repository / "files.py").write_text("# now completely different\n")
with pytest.raises(StaleReferenceError, match="Source changed"):
scout.read(symbol.id)
response = scout.search("read file", mode="lexical")
assert not response["results"]
assert response["warnings"]
(repository / "files.py").unlink()
with pytest.raises(StaleReferenceError, match="Source unavailable"):
scout.read(symbol.id)
def test_replacing_source_with_symlink_is_rejected(repository, tmp_path):
path = tmp_path / "index.sqlite"
build_index(repository, path)
scout = Scout(Index(path))
symbol = next(s for s in scout.index.symbols if s.name == "add_numbers")
outside = tmp_path / "outside.py"
outside.write_text((repository / "numbers.py").read_text())
(repository / "numbers.py").unlink()
(repository / "numbers.py").symlink_to(outside)
with pytest.raises(ValueError, match="symlink"):
scout.read(symbol.id)
class FakeEncoder:
"""Deterministic embeddings exercise index integrity, not model quality."""
fingerprint = "test-encoder-v1"
dimension = 4
def __init__(self):
self.encoded = 0
def encode(self, texts, **kwargs):
self.encoded += len(texts)
return np.tile(np.array([1, 0, 0, 0], dtype=np.float32), (len(texts), 1))
def test_dense_index_reuses_vectors_and_rejects_wrong_model(repository, tmp_path):
path = tmp_path / "index.sqlite"
encoder = FakeEncoder()
first = build_index(repository, path, encoder)
assert encoder.encoded == first["symbols"]
second = build_index(repository, path, encoder)
assert encoder.encoded == first["symbols"]
assert second["reused_embeddings"] == first["symbols"]
assert Index(path).vectors.shape == (first["symbols"], 4)
encoder.fingerprint = "different-model"
with pytest.raises(ValueError, match="fingerprints differ"):
Scout(Index(path), encoder)
def test_failed_refresh_keeps_previous_complete_index(repository, tmp_path):
path = tmp_path / "index.sqlite"
first = build_index(repository, path)
with pytest.raises(ValueError, match="exceeds"):
build_index(repository, path, max_symbols=1)
assert Index(path).metadata["snapshot"] == first["snapshot"]
@pytest.mark.parametrize(
"kwargs",
[
{"query": ""},
{"query": "x", "top_k": 0},
{"query": "x", "max_chars": 1},
{"query": "x", "mode": "bad"},
{"query": "x", "mode": "dense"},
],
)
def test_search_input_validation(repository, tmp_path, kwargs):
path = tmp_path / "index.sqlite"
build_index(repository, path)
with pytest.raises(ValueError):
Scout(Index(path)).search(**kwargs)
def test_feedback_is_logged_without_weight_update(repository, tmp_path):
path, trace = tmp_path / "index.sqlite", tmp_path / "trace.jsonl"
build_index(repository, path)
scout = Scout(Index(path), trace_path=trace)
result = scout.search("read file", mode="lexical")
answer = scout.feedback(result["request_id"], [result["results"][0]["id"]], "helpful")
assert answer == {"recorded": True, "weights_updated": False}
assert [json.loads(line)["event"] for line in trace.read_text().splitlines()] == [
"search",
"feedback",
]
def test_feedback_rejects_unknown_requests(repository, tmp_path):
path = tmp_path / "index.sqlite"
build_index(repository, path)
scout = Scout(Index(path), trace_path=tmp_path / "trace.jsonl")
with pytest.raises(ValueError, match="Unknown or expired"):
scout.feedback("a" * 32, [], "helpful")
def test_context_does_not_repeat_overlapping_source_lines(repository, tmp_path):
path = tmp_path / "index.sqlite"
build_index(repository, path)
scout = Scout(Index(path))
response = scout.search("Writer save write path text", mode="lexical")
seen = set()
for item in response["results"] + response["neighbors"]:
locations = {(item["path"], i) for i in range(item["start_line"], item["end_line"] + 1)}
assert not seen & locations
seen |= locations
def test_language_and_documentation_filters(repository, tmp_path):
(repository / "README.md").write_text("uniquedocumentationneedle")
(repository / "client.ts").write_text("function uniquetypescriptneedle() { return 1; }")
path = tmp_path / "index.sqlite"
build_index(repository, path)
scout = Scout(Index(path))
assert not scout.search("uniquedocumentationneedle", mode="lexical")["results"]
docs = scout.search("uniquedocumentationneedle", mode="lexical", include_docs=True)
assert docs["results"][0]["path"] == "README.md"
code = scout.search("uniquetypescriptneedle", mode="lexical", language="typescript")
assert code["results"][0]["path"] == "client.ts"
assert not scout.search("uniquetypescriptneedle", mode="lexical", language="python")["results"]
def test_repository_cluster_bootstrap_retains_group_correlation():
from micro_scout.evaluate import paired_mrr_interval
result = paired_mrr_interval(
np.array([1, 1, 1, 10]), np.array([2, 2, 2, 2]), ["a", "a", "a", "b"]
)
assert result["delta"] == pytest.approx(0.275)
assert result["ci95"] == pytest.approx([-0.4, 0.5])
assert result["repository_clusters"] == 2
+72
View File
@@ -0,0 +1,72 @@
import ast
from micro_scout.data import normalize_row
from micro_scout.text import code_fingerprints, lexical_tokens, strip_python_documentation
def test_strip_documentation_preserves_runtime_strings_and_unicode():
code = '''def café(value):
"""Find the secret target description."""
# also remove a comment
message = "keep this literal # content"
return message + value
'''
clean = strip_python_documentation(code)
assert "secret target" not in clean
assert "remove a comment" not in clean
assert '"keep this literal # content"' in clean
ast.parse(clean)
def test_nested_docstrings_are_removed():
code = '''class C:
"""Outer text."""
def run(self):
"""Inner text."""
return 42
'''
clean = strip_python_documentation(code)
assert "Outer text" not in clean and "Inner text" not in clean
ast.parse(clean)
def test_comment_removal_does_not_change_multiline_literal():
code = 'def x():\n text = """a\n# literal\nb"""\n return text\n'
assert "# literal" in strip_python_documentation(code)
def test_fingerprint_detects_renamed_clone():
a = code_fingerprints("def add(a, b):\n return a + b\n")
b = code_fingerprints("def sum_values(x, y):\n return x + y\n")
assert a[0] != b[0] and a[1] == b[1]
def test_tokenizer_splits_identifiers_and_keeps_exact_name():
assert lexical_tokens("parseHTTP get_user_id") == [
"parsehttp",
"parse",
"http",
"get_user_id",
"get",
"user",
"id",
]
def test_dataset_normalization_uses_no_docstring_as_code():
row = normalize_row(
{
"repo": "Example/Project",
"path": "src/files.py",
"url": "https://github.com/Example/Project/blob/abc/src/files.py#L1-L5",
"docstring": "Read every nonempty line from the given input file.",
"code": '''def read_lines(path):
"""Read every nonempty line from the given input file."""
with open(path) as stream:
return [line.strip() for line in stream if line.strip()]
''',
}
)
assert row is not None
assert row["query"] not in row["code"]
assert row["repo"] == "example/project"