Spaces:
Running on Zero
Running on Zero
File size: 4,488 Bytes
1e214ed | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 | """Offline regression tests for persistent user-correction memory."""
import os
import tempfile
import sys
import types
import shutil
from correction_memory import CorrectionMemory
def main():
checks = 0
def check(label, ok):
nonlocal checks
checks += 1
print(("PASS " if ok else "FAIL ") + label)
if not ok: raise SystemExit(1)
with tempfile.TemporaryDirectory() as d:
path = os.path.join(d, "corrections.jsonl")
m = CorrectionMemory(path)
check("ordinary request is not stored", not m.add("Write a Python function to sort a list"))
check("Arabic correction is detected and stored", m.add("قلتلك لا تكرر نفس الخطأ، استخدم حل عام"))
check("duplicate correction is not stored twice", not m.add("قلتلك لا تكرر نفس الخطأ، استخدم حل عام"))
check("correction survives a new memory instance", "لا تكرر" in CorrectionMemory(path).prompt("حل عام لا تكرر الخطأ"))
check("English rejection is detected", m.add("That is wrong; do not repeat the same approach"))
check("relevant lessons are surfaced", len(m.relevant("wrong approach, do not repeat")) >= 1)
with open(path, "a", encoding="utf-8") as f: f.write("{corrupt json\n")
check("corrupt line does not destroy prior memory", len(CorrectionMemory(path).relevant("wrong approach")) >= 1)
with tempfile.TemporaryDirectory() as d:
path = os.path.join(d, "bounded.jsonl")
with open(path, "w", encoding="utf-8") as f:
for i in range(305):
f.write(__import__("json").dumps({"ts": i, "text": f"غلط سجل {i} لا تكرر"}, ensure_ascii=False) + "\n")
bounded = CorrectionMemory(path)
check("memory reads only the newest bounded number of records", len(bounded._load()) == 300)
check("adding a lesson physically keeps the JSONL file bounded", bounded.add("غلط تصحيح جديد لا تكرر") and sum(1 for _ in open(path, encoding="utf-8")) == 300)
# Exercise the optional Hugging Face Dataset adapter with a fake client only;
# this checks our adapter calls, not Internet connectivity or real HF permissions.
with tempfile.TemporaryDirectory() as d:
local = os.path.join(d, "eim_corrections.jsonl")
remote_file = os.path.join(d, "remote.jsonl")
with open(remote_file, "w", encoding="utf-8") as f:
f.write('{"ts":1,"text":"لا تكرر الخطأ، استخدم حل عام"}\n')
calls = {"pull": 0, "create": 0, "upload": 0}
fake = types.ModuleType("huggingface_hub")
def fake_download(**kwargs):
calls["pull"] += 1
return remote_file
class FakeApi:
def __init__(self, token=None):
assert token == "test-token"
def create_repo(self, **kwargs):
calls["create"] += 1
def upload_file(self, **kwargs):
calls["upload"] += 1
shutil.copyfile(kwargs["path_or_fileobj"], remote_file)
fake.hf_hub_download = fake_download
fake.HfApi = FakeApi
prior_module = sys.modules.get("huggingface_hub")
old_repo, old_token = os.environ.get("EIM_CORRECTIONS_REPO"), os.environ.get("HF_TOKEN")
try:
sys.modules["huggingface_hub"] = fake
os.environ["EIM_CORRECTIONS_REPO"] = "user/private-corrections"
os.environ["HF_TOKEN"] = "test-token"
synced = CorrectionMemory(local)
check("configured private-dataset memory pulls remote file", calls["pull"] == 1 and "لا تكرر" in synced.prompt("حل عام"))
check("configured private-dataset memory pushes correction through adapter", synced.add("That is wrong; do not repeat this approach") and synced.sync_now() and calls["create"] == 1 and calls["upload"] == 1)
finally:
if prior_module is None:
sys.modules.pop("huggingface_hub", None)
else:
sys.modules["huggingface_hub"] = prior_module
if old_repo is None: os.environ.pop("EIM_CORRECTIONS_REPO", None)
else: os.environ["EIM_CORRECTIONS_REPO"] = old_repo
if old_token is None: os.environ.pop("HF_TOKEN", None)
else: os.environ["HF_TOKEN"] = old_token
print(f"All {checks} correction-memory checks passed.")
if __name__ == "__main__": main()
|