File size: 7,203 Bytes
8fa3dd6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
from __future__ import annotations

import argparse
from array import array
import hashlib
import json
import math
from pathlib import Path
import re
import sqlite3
from typing import Iterable, List, Sequence, Tuple

WORD_RE = re.compile(r"[A-Za-z0-9_]+")


def connect(db_path: str) -> sqlite3.Connection:
    Path(db_path).parent.mkdir(parents=True, exist_ok=True)
    con = sqlite3.connect(db_path)
    con.execute("PRAGMA journal_mode=WAL")
    con.execute("PRAGMA synchronous=NORMAL")
    return con


def init_db(con: sqlite3.Connection, dim: int = 384) -> None:
    con.execute(
        """
        CREATE TABLE IF NOT EXISTS docs (
            id INTEGER PRIMARY KEY AUTOINCREMENT,
            title TEXT,
            source TEXT,
            text TEXT NOT NULL,
            meta_json TEXT DEFAULT '{}',
            vector BLOB,
            dim INTEGER NOT NULL
        )
        """
    )
    try:
        con.execute("CREATE VIRTUAL TABLE IF NOT EXISTS docs_fts USING fts5(title, text, source, content='docs', content_rowid='id')")
    except sqlite3.OperationalError as exc:
        raise SystemExit("Your SQLite build needs FTS5 enabled for this RAG store.") from exc
    con.execute("CREATE INDEX IF NOT EXISTS idx_docs_source ON docs(source)")
    con.execute("CREATE TABLE IF NOT EXISTS meta (key TEXT PRIMARY KEY, value TEXT)")
    con.execute("INSERT OR REPLACE INTO meta(key, value) VALUES('vector_dim', ?)", (str(dim),))
    con.commit()


def chunk_text(text: str, max_chars: int = 1200, overlap: int = 120) -> Iterable[str]:
    text = re.sub(r"\s+", " ", text).strip()
    if not text:
        return
    start = 0
    while start < len(text):
        end = min(len(text), start + max_chars)
        yield text[start:end]
        if end == len(text):
            break
        start = max(0, end - overlap)


def hashed_embedding(text: str, dim: int = 384) -> array:
    """Deterministic no-external-model lexical embedding.

    This is not a neural embedding. It is a bootstrap retriever that can later be replaced by an
    embedding head trained inside Ares.
    """
    vec = array("f", [0.0]) * dim
    words = [w.lower() for w in WORD_RE.findall(text)]
    features: List[str] = []
    features.extend(words)
    features.extend([" ".join(words[i : i + 2]) for i in range(max(0, len(words) - 1))])
    features.extend([" ".join(words[i : i + 3]) for i in range(max(0, len(words) - 2))])
    for feat in features:
        h = hashlib.blake2b(feat.encode("utf-8"), digest_size=8).digest()
        bucket = int.from_bytes(h[:4], "little") % dim
        sign = 1.0 if (h[4] & 1) else -1.0
        vec[bucket] += sign
    norm = math.sqrt(sum(x * x for x in vec)) or 1.0
    for i, x in enumerate(vec):
        vec[i] = x / norm
    return vec


def vector_to_blob(vec: array) -> bytes:
    return vec.tobytes()


def blob_to_vector(blob: bytes) -> array:
    vec = array("f")
    vec.frombytes(blob)
    return vec


def cosine(a: array, b: array) -> float:
    return float(sum(x * y for x, y in zip(a, b)))


def add_doc(con: sqlite3.Connection, title: str, text: str, source: str = "", meta=None, dim: int = 384) -> int:
    meta = meta or {}
    vec = hashed_embedding(title + "\n" + text, dim=dim)
    cur = con.execute(
        "INSERT INTO docs(title, source, text, meta_json, vector, dim) VALUES (?, ?, ?, ?, ?, ?)",
        (title, source, text, json.dumps(meta), vector_to_blob(vec), dim),
    )
    doc_id = int(cur.lastrowid)
    con.execute("INSERT INTO docs_fts(rowid, title, text, source) VALUES (?, ?, ?, ?)", (doc_id, title, text, source))
    return doc_id


def ingest_paths(con: sqlite3.Connection, paths: Sequence[str], dim: int = 384) -> int:
    count = 0
    files: List[Path] = []
    for item in paths:
        p = Path(item)
        if p.is_dir():
            files.extend([x for x in sorted(p.rglob("*")) if x.is_file() and x.suffix.lower() in {".txt", ".md"}])
        elif p.is_file():
            files.append(p)
        else:
            raise FileNotFoundError(item)
    for p in files:
        text = p.read_text(encoding="utf-8", errors="ignore")
        for i, chunk in enumerate(chunk_text(text)):
            add_doc(con, title=f"{p.name} chunk {i}", text=chunk, source=str(p), meta={"chunk": i}, dim=dim)
            count += 1
    con.commit()
    return count


def search(con: sqlite3.Connection, query: str, limit: int = 5, dim: int = 384) -> List[Tuple[int, str, str, float]]:
    qvec = hashed_embedding(query, dim=dim)
    # FTS candidate pool, then vector rerank. If FTS syntax rejects query, use a broad lexical fallback.
    try:
        rows = con.execute(
            """
            SELECT d.id, d.title, d.text, d.vector, bm25(docs_fts) AS bm25
            FROM docs_fts
            JOIN docs d ON docs_fts.rowid = d.id
            WHERE docs_fts MATCH ?
            ORDER BY bm25
            LIMIT ?
            """,
            (query, max(25, limit * 5)),
        ).fetchall()
    except sqlite3.OperationalError:
        terms = [t.lower() for t in WORD_RE.findall(query)]
        like = "%" + "%".join(terms[:4]) + "%" if terms else "%"
        rows = con.execute(
            "SELECT id, title, text, vector, 0.0 FROM docs WHERE lower(text) LIKE ? LIMIT ?",
            (like, max(25, limit * 5)),
        ).fetchall()
    if not rows:
        rows = con.execute(
            "SELECT id, title, text, vector, 0.0 FROM docs LIMIT ?",
            (max(25, limit * 5),),
        ).fetchall()

    scored = []
    for doc_id, title, text, blob, bm25_score in rows:
        score = cosine(qvec, blob_to_vector(blob)) - 0.01 * float(bm25_score)
        scored.append((int(doc_id), str(title), str(text), float(score)))
    scored.sort(key=lambda x: x[3], reverse=True)
    return scored[:limit]


def main() -> None:
    parser = argparse.ArgumentParser(description="Ares SQLite/FTS5 RAG bootstrap store.")
    sub = parser.add_subparsers(dest="cmd", required=True)

    p_init = sub.add_parser("init")
    p_init.add_argument("--db", required=True)
    p_init.add_argument("--dim", type=int, default=384)

    p_ingest = sub.add_parser("ingest")
    p_ingest.add_argument("--db", required=True)
    p_ingest.add_argument("--input", nargs="+", required=True)
    p_ingest.add_argument("--dim", type=int, default=384)

    p_search = sub.add_parser("search")
    p_search.add_argument("--db", required=True)
    p_search.add_argument("--query", required=True)
    p_search.add_argument("--limit", type=int, default=5)
    p_search.add_argument("--dim", type=int, default=384)

    args = parser.parse_args()
    con = connect(args.db)
    if args.cmd == "init":
        init_db(con, dim=args.dim)
        print(f"Initialized {args.db}")
    elif args.cmd == "ingest":
        init_db(con, dim=args.dim)
        n = ingest_paths(con, args.input, dim=args.dim)
        print(f"Ingested {n} chunks into {args.db}")
    elif args.cmd == "search":
        results = search(con, args.query, limit=args.limit, dim=args.dim)
        for doc_id, title, text, score in results:
            print(json.dumps({"id": doc_id, "title": title, "score": score, "text": text[:500]}, ensure_ascii=False))


if __name__ == "__main__":
    main()