feat: implement local code scout, training pipeline, and MCP tools
This commit is contained in:
@@ -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}
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user