Spaces:
Runtime error
Runtime error
File size: 7,310 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 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 | """Persistent memory for explicit user corrections, with optional private HF Dataset sync.
Local-only by default (stdlib). Set EIM_CORRECTIONS_REPO or EIM_MEMORY_REPO to a
private Hugging Face Dataset repo and HF_TOKEN to a token with write access to keep
corrections across ephemeral Space restarts. Remote sync failures never break chat.
"""
from __future__ import annotations
import json
import os
import re
import shutil
import threading
import time
from pathlib import Path
_EN_MARKERS = (
"that's wrong", "that is wrong", "wrong answer", "not what i asked", "i said",
"i told you", "you repeated", "don't repeat", "do not repeat", "you made the same",
"not correct", "you forgot", "i already said", "stop doing", "instead of",
)
_AR_MARKERS = (
"غلط", "مو هذا", "مو هيج", "مو هيچ", "قلتلك", "كلتلك", "نفس الخطأ", "نفس الاخطاء",
"لا تكرر", "لا تعيد", "كررت", "نسيت", "مو اللي طلبته", "ما طلبت", "مو صحيح", "خطأ",
)
_STOP = set("the a an and or to of for in on with this that it is are was were be do did you your i me my we please fix make write code answer about from into using use".split())
def _tokens(text: str) -> set[str]:
words = re.findall(r"[a-zA-Z0-9_]+|[\u0600-\u06FF]+", (text or "").lower())
return {w for w in words if len(w) > 1 and w not in _STOP}
class CorrectionMemory:
MAX_RECORDS = 300
PUSH_DELAY = 2.0
def __init__(self, path: str | None = None, repo: str | None = None):
self.path = os.path.abspath(path or os.environ.get("EIM_CORRECTIONS_PATH", "eim_corrections.jsonl"))
self.repo = (repo if repo is not None else (
os.environ.get("EIM_CORRECTIONS_REPO") or os.environ.get("EIM_MEMORY_REPO", "")
)).strip()
self.token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACEHUB_API_TOKEN") or None
self._lock = threading.RLock()
self._timer: threading.Timer | None = None
self.sync_status = "disabled" if not self.repo else "configured"
Path(self.path).parent.mkdir(parents=True, exist_ok=True)
if self.repo:
self._pull_remote()
@staticmethod
def is_correction(text: str) -> bool:
low = (text or "").lower()
return any(x in low for x in _EN_MARKERS + _AR_MARKERS)
def add(self, text: str) -> bool:
text = (text or "").strip()
if not text or not self.is_correction(text):
return False
stored_text = text[:2000]
norm = " ".join(stored_text.split()).casefold()
with self._lock:
existing = self._load()
if any(" ".join(r.get("text", "").split()).casefold() == norm for r in existing):
return False
rows = (existing + [{"ts": int(time.time()), "text": stored_text}])[-self.MAX_RECORDS:]
temporary = f"{self.path}.{os.getpid()}.{threading.get_ident()}.tmp"
try:
with open(temporary, "w", encoding="utf-8") as f:
for row in rows:
f.write(json.dumps(row, ensure_ascii=False) + "\n")
os.replace(temporary, self.path)
finally:
try:
os.unlink(temporary)
except OSError:
pass
self._schedule_push()
return True
def _load(self) -> list[dict]:
rows = []
try:
with open(self.path, encoding="utf-8") as f:
for line in f:
try:
row = json.loads(line)
if isinstance(row, dict) and isinstance(row.get("text"), str):
rows.append(row)
except (ValueError, TypeError):
continue
except OSError:
pass
return rows[-self.MAX_RECORDS:]
def relevant(self, query: str, limit: int = 4) -> list[str]:
rows = self._load()
if not rows:
return []
q = _tokens(query)
scored = []
for i, row in enumerate(rows):
words = _tokens(row.get("text", ""))
overlap = len(q & words) / max(1, len(q | words))
# Recency is a tie-breaker, not a replacement for relevance.
score = overlap + 0.015 * (i / max(1, len(rows) - 1))
if overlap > 0 or i >= len(rows) - 3:
scored.append((score, i, row.get("text", "")))
scored.sort(reverse=True)
return [text for _, _, text in scored[:max(1, limit)] if text]
def prompt(self, query: str, limit: int = 4) -> str:
lessons = self.relevant(query, limit)
if not lessons:
return ""
bullets = "\n".join(f"- {item}" for item in lessons)
return ("Persistent user corrections from earlier turns. Treat these as constraints; do not repeat "
"rejected approaches. If a correction conflicts with the current explicit request, follow the current request.\n"
+ bullets)
def _pull_remote(self) -> None:
"""Pull corrections from a configured dataset repo; tolerate missing repo/file/offline mode."""
try:
from huggingface_hub import hf_hub_download
local = hf_hub_download(
repo_id=self.repo,
repo_type="dataset",
filename=os.path.basename(self.path),
token=self.token,
local_dir=os.path.dirname(self.path),
)
if os.path.abspath(local) != self.path and os.path.isfile(local):
shutil.copyfile(local, self.path)
self.sync_status = "pulled"
except Exception as exc: # A missing file/new repo is normal on first launch.
self.sync_status = f"pull-unavailable:{type(exc).__name__}"
def _schedule_push(self) -> None:
if not self.repo:
return
with self._lock:
if self._timer is not None:
self._timer.cancel()
self._timer = threading.Timer(self.PUSH_DELAY, self.sync_now)
self._timer.daemon = True
self._timer.start()
def sync_now(self) -> bool:
"""Push the JSONL file to the configured dataset. Returns True only after upload succeeds."""
if not self.repo:
self.sync_status = "disabled"
return False
if not self.token:
self.sync_status = "push-unavailable:missing-token"
return False
try:
from huggingface_hub import HfApi
with self._lock:
api = HfApi(token=self.token)
api.create_repo(repo_id=self.repo, repo_type="dataset", private=True, exist_ok=True)
api.upload_file(
path_or_fileobj=self.path,
path_in_repo=os.path.basename(self.path),
repo_id=self.repo,
repo_type="dataset",
commit_message="Update EIM user-correction memory",
)
self.sync_status = "pushed"
return True
except Exception as exc:
self.sync_status = f"push-unavailable:{type(exc).__name__}"
return False
|