#!/usr/bin/env python3 """ Momo 1.0 — local test harness. Runs a battery of 16 real-world messages across all 6 actions, checks: 1. output parses as the Momo JSON decision 2. predicted action matches expectation 3. decision passes semantic validation (momo_core.validate_decision) 4. private names NEVER leak into web_query (anonymization) 5. latency stats (avg / p95) Model resolution order: --adapter path > ./adapter > MOMO_HF_REPO env > base model only (warning) Usage: python test.py # full local test (needs torch+transformers) python test.py --mock # logic-only test: parser/validator, no model python test.py --adapter ./adapter """ from __future__ import annotations import argparse import os import statistics import sys import time from pathlib import Path from typing import Any, Dict, List, Optional sys.path.insert(0, str(Path(__file__).resolve().parent)) from momo_core import ( # noqa: E402 BASE_MODEL_ID, SYSTEM_PROMPT, parse_decision, validate_decision, ) class Case: def __init__(self, text: str, expect: List[str], anonymize: Optional[List[str]] = None, note: str = ""): self.text = text self.expect = expect # acceptable actions (first = preferred) self.anonymize = anonymize or [] # names that must NOT appear in web_query self.note = note CASES: List[Case] = [ Case("Remember that my favorite food is sushi.", ["STORE_MEMORY"]), Case("Actually, I don't like sushi anymore — ramen is my favorite now.", ["UPDATE_MEMORY"]), Case("Forget my old phone number.", ["DELETE_MEMORY"]), Case("What's my favorite food?", ["MEMORY_ONLY"]), Case("What's the weather in Tokyo today?", ["WEB_ONLY"], note="city kept (location query)"), Case("My girlfriend Sarah loves hiking — gift ideas for her birthday?", ["HYBRID"], anonymize=["Sarah"], note="Sarah must be stripped from web_query"), Case("I just adopted a Beagle named Coco, remember!", ["STORE_MEMORY"]), Case("Hi Momo!", ["MEMORY_ONLY"], note="chit-chat path, no retrieval"), Case("Is the Pixel 8 worth buying?", ["WEB_ONLY"]), Case("My boss Priya said keto is bad — is that true?", ["WEB_ONLY"], anonymize=["Priya"]), Case("Note that I'm vegan now — find me high-protein recipes.", ["HYBRID"]), Case("When is my birthday?", ["MEMORY_ONLY"]), Case("Erase everything about my breakup with Maya.", ["DELETE_MEMORY"], note="mq should reference Maya"), Case("Who won the IPL final?", ["WEB_ONLY"]), Case("I moved from Pune to Bangalore last week.", ["UPDATE_MEMORY"]), Case("What should I get Emma for her graduation?", ["HYBRID"], anonymize=["Emma"]), ] PRIVATE_NAME_LEAKS = ("Sarah", "Priya", "Maya", "Emma") def build_generation_prompt(tok, user_text: str) -> str: msgs = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_text}, ] return tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True) def load_model(device: str, adapter: str, base: str): import torch from transformers import AutoModelForCausalLM, AutoTokenizer dtype = torch.float16 if (device == "cuda" and torch.cuda.is_available()) else torch.float32 if device == "cuda" and torch.cuda.is_available(): device = "cuda" else: device = "cpu" tok = AutoTokenizer.from_pretrained(base, trust_remote_code=True) if tok.pad_token is None: tok.pad_token = tok.eos_token model = AutoModelForCausalLM.from_pretrained( base, torch_dtype=dtype, trust_remote_code=True ).to(device).eval() adapter_resolved = adapter or "./adapter" if Path("./adapter").exists() else adapter if not adapter_resolved and os.environ.get("MOMO_HF_REPO"): adapter_resolved = os.environ["MOMO_HF_REPO"] if adapter_resolved: from peft import PeftModel try: model = PeftModel.from_pretrained(model, adapter_resolved) model = model.merge_and_unload() print(f"[momo] loaded adapter: {adapter_resolved}") except Exception as e: # noqa: BLE001 print(f"[momo] WARNING: adapter load failed ({e}); using BASE model only") else: print("[momo] WARNING: no adapter found — testing BASE model (untrained Momo)") return tok, model, device def run_model(tok, model, device: str, user_text: str, max_new_tokens: int = 128) -> str: import torch prompt = build_generation_prompt(tok, user_text) inputs = tok(prompt, return_tensors="pt").to(device) with torch.inference_mode(): out = model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=False, temperature=None, top_p=None, top_k=None, pad_token_id=tok.pad_token_id or tok.eos_token_id, eos_token_id=tok.eos_token_id, ) gen = out[0][inputs["input_ids"].shape[1]:] return tok.decode(gen, skip_special_tokens=True) def evaluate_case(case: Case, raw: str) -> Dict[str, Any]: dec = parse_decision(raw) res: Dict[str, Any] = {"case": case.text, "ok": False, "expected": case.expect, "action": None, "web_query": "", "issues": []} if dec is None: res["issues"].append("unparseable output") return res res["action"] = dec["action"] res["web_query"] = dec["web_query"] if dec["action"] not in case.expect: res["issues"].append(f"action={dec['action']} expected={case.expect}") problems = validate_decision(dec) res["issues"].extend(problems) for nm in case.anonymize: if nm.lower() in dec["web_query"].lower(): res["issues"].append(f"PII LEAK: '{nm}' in web_query") res["ok"] = not res["issues"] return res def print_report(results: List[Dict[str, Any]], latencies: List[float]) -> None: print("\n" + "=" * 100) print("MOMO 1.0 — TEST REPORT") print("=" * 100) n_pass = 0 for r in results: n_pass += bool(r["ok"]) flag = "PASS" if r["ok"] else "FAIL" print(f"\n[{flag}] {r['case']}") print(f" action={r['action']} expected={r['expected']}") if r["web_query"]: print(f" web_query=\"{r['web_query']}\"") for iss in r["issues"]: print(f" ! {iss}") n = len(results) print("\n" + "-" * 100) print(f"TOTAL: {n_pass}/{n} passed") if latencies: lat_sorted = sorted(latencies) p95 = lat_sorted[max(0, int(round(0.95 * len(lat_sorted))) - 1)] print(f"latency: avg={statistics.mean(latencies)*1000:.0f}ms " f"median={statistics.median(latencies)*1000:.0f}ms p95={p95*1000:.0f}ms") print("=" * 100) def mock_mode() -> int: """No-model logic test: parser + validator round-trips on all cases' targets.""" print("[momo] MOCK MODE — testing parser/validator logic only") from momo_core import build_decision checks = [ (build_decision("STORE_MEMORY", memory_text="User's favorite food is sushi."), "STORE_MEMORY", True), (build_decision("UPDATE_MEMORY", need_memory_search=True, memory_query="user's favorite food", memory_text="User's favorite food is ramen."), "UPDATE_MEMORY", True), (build_decision("DELETE_MEMORY", need_memory_search=True, memory_query="user's phone number"), "DELETE_MEMORY", True), (build_decision("MEMORY_ONLY", need_memory_search=True, memory_query="user's birthday"), "MEMORY_ONLY", True), (build_decision("WEB_ONLY", need_web=True, web_query="weather Tokyo"), "WEB_ONLY", True), (build_decision("HYBRID", need_memory_search=True, need_web=True, memory_query="Sarah gift preferences", web_query="birthday gift ideas for girlfriend"), "HYBRID", True), (build_decision("MEMORY_ONLY"), "MEMORY_ONLY", True), # chit-chat (build_decision("WEB_ONLY", need_web=True, web_query=""), "WEB_ONLY", False), # bad (build_decision("STORE_MEMORY"), "STORE_MEMORY", False), # bad: no memory_text ] fails = 0 for dec, action, should_pass in checks: raw = __import__("json").dumps(dec) parsed = parse_decision(raw) problems = validate_decision(parsed) if parsed else ["unparseable"] passed = (parsed is not None and not problems) verdict = "PASS" if passed == should_pass else "FAIL" fails += verdict == "FAIL" print(f" [{verdict}] {action:<14} valid={passed} (expected {should_pass}) {problems or ''}") print(f"\nMOCK: {len(checks) - fails}/{len(checks)} logic checks passed") return 1 if fails else 0 def main() -> int: ap = argparse.ArgumentParser(description="Test Momo 1.0 locally") ap.add_argument("--adapter", type=str, default="", help="adapter path or HF repo id") ap.add_argument("--base", type=str, default=BASE_MODEL_ID) ap.add_argument("--device", type=str, default="cuda" if os.environ.get("CUDA_VISIBLE_DEVICES") else "cpu") ap.add_argument("--mock", action="store_true", help="logic-only test, no model download") ap.add_argument("--rounds", type=int, default=1, help="repeat battery for latency stats") args = ap.parse_args() if args.mock: return mock_mode() try: tok, model, device = load_model(args.device, args.adapter, args.base) except Exception as e: # noqa: BLE001 print(f"[momo] model load failed: {e}\n-> falling back to MOCK mode") return mock_mode() all_results: List[Dict[str, Any]] = [] latencies: List[float] = [] for rnd in range(args.rounds): for case in CASES: t0 = time.perf_counter() raw = run_model(tok, model, device, case.text) dt = time.perf_counter() - t0 res = evaluate_case(case, raw) res["latency_s"] = dt all_results.append(res) latencies.append(dt) # report on round 1 (dedup), but latency over all rounds print_report(all_results[: len(CASES)], latencies) failed = sum(1 for r in all_results[: len(CASES)] if not r["ok"]) return 1 if failed else 0 if __name__ == "__main__": sys.exit(main())