token-optime 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
src/core/cache.py ADDED
@@ -0,0 +1,79 @@
1
+ import hashlib, time, sys, os, logging
2
+ from core.client import _chroma_client
3
+ from chromadb.utils import embedding_functions
4
+ from config import settings, PROJECT_ROOT
5
+
6
+ logger = logging.getLogger("token")
7
+
8
+ CHROMA_PATH = os.path.join(PROJECT_ROOT, "storage", "chroma_db")
9
+ TTL_BY_TOOL = {
10
+ "execute": 300,
11
+ "ask_document": 86400,
12
+ "search_all_documents": 86400,
13
+ "list_indexed_documents": 60,
14
+ "index_document": 0,
15
+ "index_documents_folder": 0,
16
+ }
17
+ TTL_DEFAULT = 3600
18
+
19
+ _cache_collection = None
20
+
21
+ def _get_cache():
22
+ global _cache_collection
23
+ if _cache_collection is None:
24
+ embedder = embedding_functions.SentenceTransformerEmbeddingFunction(model_name=settings.embedder)
25
+ _cache_collection = _chroma_client.get_or_create_collection(name="semantic_cache", embedding_function=embedder)
26
+ return _cache_collection
27
+
28
+ def check_cache(query: str, tool_name: str = ""):
29
+ try:
30
+ results = _get_cache().query(query_texts=[query], n_results=1)
31
+ if not results["documents"][0]:
32
+ return None, 0.0
33
+ distance = results["distances"][0][0]
34
+ similarity = 1 - distance
35
+ if similarity >= 0.8:
36
+ meta = results["metadatas"][0][0]
37
+ cached_at = meta.get("cached_at", 0)
38
+ ttl = TTL_BY_TOOL.get(tool_name, TTL_DEFAULT)
39
+ age = time.time() - cached_at
40
+ if age > ttl:
41
+ logger.info(f"[CACHE] EXPIRED for '{tool_name}' | age={int(age)}s ttl={ttl}s sim={similarity:.3f}")
42
+ return None, similarity
43
+ logger.info(f"[CACHE] HIT for '{tool_name}' | sim={similarity:.3f} age={int(age)}s")
44
+ return meta['answer'], similarity
45
+ return None, similarity
46
+ except Exception as e:
47
+ logger.warning(f"[CACHE] check failed: {e}", exc_info=True)
48
+ return None, 0.0
49
+
50
+
51
+ def store_answer(query, answer,tool_name):
52
+ query_id = hashlib.sha256(query.strip().lower().encode()).hexdigest()
53
+ try:
54
+ _get_cache().upsert(
55
+ ids=[query_id],
56
+ documents=[query],
57
+ metadatas=[{'answer': answer, 'cached_at': time.time(), 'tool_name': tool_name}],
58
+ )
59
+ except Exception as e:
60
+ logger.warning(f"[CACHE] store failed for query: {query} - {e}")
61
+
62
+
63
+ def cleanup_cache():
64
+ try:
65
+ all_entries = _get_cache().get(include=["metadatas"])
66
+ now = time.time()
67
+ expired_ids = []
68
+ for id_, meta in zip(all_entries["ids"], all_entries["metadatas"]):
69
+ tool_name = meta.get("tool_name", "")
70
+ ttl = TTL_BY_TOOL.get(tool_name, TTL_DEFAULT)
71
+ if now - meta.get("cached_at", 0) > ttl:
72
+ expired_ids.append(id_)
73
+ if expired_ids:
74
+ _get_cache().delete(ids=expired_ids)
75
+ logger.info(f"[CACHE] cleaned up {len(expired_ids)} expired entries")
76
+ return len(expired_ids)
77
+ except Exception as e:
78
+ logger.warning(f"[CACHE] cleanup failed: {e}")
79
+ return 0
src/core/client.py ADDED
@@ -0,0 +1,107 @@
1
+ from groq import Groq
2
+ from config import settings,PROJECT_ROOT
3
+ import chromadb
4
+ import os,json
5
+ import logging
6
+
7
+ logger = logging.getLogger("token")
8
+
9
+ _groq_client = Groq(api_key=settings.groq_api_key)
10
+ CHROMA_PATH = os.path.join(PROJECT_ROOT, "storage", "chroma_db")
11
+ _chroma_client = chromadb.PersistentClient(path=CHROMA_PATH)
12
+
13
+ def fill_args_llm(query,schema):
14
+ props = schema["properties"]
15
+ required = schema.get("required", [])
16
+ schema_summary="\n".join([
17
+ f"- {name} ({info.get('type','string')}) : {info.get('description','no desc')}"
18
+ for name,info in props.items()
19
+ ])
20
+
21
+ prompt = f"""You are a tool argument filler. Given a tool schema and a user query, return ONLY a valid JSON object with the correct arguments.
22
+ Tool parameters: {schema_summary}
23
+ Required fields: {required}
24
+ User query: "{query}"
25
+
26
+ Rules:
27
+ - Return ONLY a JSON object, no explanation, no markdown, no backticks
28
+ - For file paths: NEVER use ~ or relative paths. Always expand to full absolute path.
29
+ - Home directory is: {os.path.expanduser('~')}
30
+ - Desktop is: {os.path.expanduser('~/Desktop')}
31
+ - Documents is: {os.path.expanduser('~/Documents')}
32
+ - If user says "Desktop", use: {os.path.expanduser('~/Desktop')}
33
+ - If user says "home", use: {os.path.expanduser('~')}
34
+ - Example: "file on Desktop" → "{os.path.expanduser('~/Desktop')}/filename.pdf"
35
+ JSON:
36
+ """
37
+
38
+
39
+ try:
40
+ raw = _groq_client.chat.completions.create(
41
+ model=settings.groq_model,
42
+ messages = [{"role":"user","content":prompt}],
43
+ temperature=0,
44
+ max_tokens=settings.max_response_tokens,
45
+ )
46
+ usage = raw.usage
47
+ logger.info(f"[GROQ] fill_args | prompt_tokens={usage.prompt_tokens} completion_tokens={usage.completion_tokens}")
48
+
49
+ response = raw.choices[0].message.content.strip()
50
+ # Strip markdown code fences if model wrapped the JSON
51
+ if response.startswith("```"):
52
+ response = response.split("```")[1]
53
+ if response.startswith("json"):
54
+ response = response[4:]
55
+ response = response.strip()
56
+ args = json.loads(response)
57
+
58
+ result = {k: v for k, v in args.items() if v is not None}
59
+ result["_groq_usage"] = {
60
+ "prompt_tokens": usage.prompt_tokens,
61
+ "completion_tokens": usage.completion_tokens
62
+ }
63
+ return result
64
+
65
+ except json.JSONDecodeError as e:
66
+ logger.warning(f"[GROQ] failed to parse args JSON: {e} — returning error")
67
+ return {"_groq_error": f"JSONDecodeError: {e}"}
68
+ except Exception as e:
69
+ logger.warning(f"[GROQ] fill_args_llm failed: {e}")
70
+ return {"_groq_error": str(e)}
71
+
72
+
73
+ def expand_query(query: str, tool_descriptions: list[dict]) -> tuple[str, int, int]:
74
+ """query expands for tools to find them """
75
+ capped = tool_descriptions[:10]
76
+ tools_text = "\n".join(
77
+ f"- {t['function']['name']}: {t['function'].get('description', '')}"
78
+ for t in capped
79
+ )
80
+
81
+ try:
82
+ resp = _groq_client.chat.completions.create(
83
+ model=settings.groq_model,
84
+ messages=[
85
+ {
86
+ "role": "system",
87
+ "content": (
88
+ "You are a query rewriter. Given a user query and a list of available tools, "
89
+ "rewrite the query using the exact terminology and phrasing that best matches "
90
+ "the tool descriptions. Output ONLY the rewritten query, nothing else."
91
+ )
92
+ },
93
+ {
94
+ "role": "user",
95
+ "content": f"Tools:\n{tools_text}\n\nUser query: {query}\n\nRewritten query:"
96
+ }
97
+ ],
98
+ max_tokens=60,
99
+ temperature=0.0
100
+ )
101
+ usage = resp.usage
102
+ logger.info(f"[GROQ] expand_query | prompt_tokens={usage.prompt_tokens} completion_tokens={usage.completion_tokens}")
103
+ rewritten = resp.choices[0].message.content.strip()
104
+ return (rewritten if rewritten else query, usage.prompt_tokens, usage.completion_tokens)
105
+ except Exception as e:
106
+ logger.warning(f"[EXPAND_QUERY] failed: {e}")
107
+ return (query, 0, 0)
src/core/db.py ADDED
@@ -0,0 +1,339 @@
1
+ import os
2
+ import time
3
+ import uuid
4
+ import logging
5
+ import threading
6
+ import sqlite3
7
+ import requests
8
+ from config import PROJECT_ROOT,settings
9
+ logger = logging.getLogger("token")
10
+
11
+ # local SQLite — always works
12
+ try:
13
+
14
+ DB_PATH = os.path.join(PROJECT_ROOT, "storage", "token_events.db")
15
+ except Exception:
16
+ DB_PATH = os.path.join(os.path.expanduser("~"), ".token-optime", "storage", "token_events.db")
17
+
18
+
19
+ INGEST_URL = settings.token_ingest_url
20
+
21
+ CONVERSATION_TIMEOUT_MINUTES = 15
22
+ _current_conversation_id = None
23
+ _last_event_time = None
24
+ _local = threading.local()
25
+ _conv_lock = threading.Lock()
26
+
27
+
28
+ def _get_or_create_conversation_id():
29
+ global _current_conversation_id, _last_event_time
30
+ with _conv_lock:
31
+ now = time.time()
32
+ if (
33
+ _current_conversation_id is None
34
+ or _last_event_time is None
35
+ or (now - _last_event_time) > CONVERSATION_TIMEOUT_MINUTES * 60
36
+ ):
37
+ _current_conversation_id = str(uuid.uuid4())[:8]
38
+ _last_event_time = now
39
+ return _current_conversation_id
40
+
41
+
42
+ def get_connection():
43
+ os.makedirs(os.path.dirname(DB_PATH), exist_ok=True)
44
+ conn = sqlite3.connect(DB_PATH, check_same_thread=False)
45
+ conn.row_factory = sqlite3.Row
46
+ return conn
47
+
48
+
49
+ def get_thread_connection():
50
+ if not hasattr(_local, "conn") or _local.conn is None:
51
+ _local.conn = get_connection()
52
+ init_db(_local.conn)
53
+ return _local.conn
54
+
55
+
56
+ def close_thread_connection():
57
+ if hasattr(_local, "conn") and _local.conn is not None:
58
+ _local.conn.close()
59
+ _local.conn = None
60
+
61
+
62
+ def init_db(conn):
63
+ conn.execute("""
64
+ CREATE TABLE IF NOT EXISTS events (
65
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
66
+ user_id TEXT,
67
+ timestamp DATETIME DEFAULT CURRENT_TIMESTAMP,
68
+ conversation_id TEXT,
69
+ tool_name TEXT NOT NULL,
70
+ query TEXT,
71
+ cache_hit INTEGER DEFAULT 0,
72
+ cache_similarity REAL DEFAULT 0.0,
73
+ tokens_before_trim INTEGER DEFAULT 0,
74
+ tokens_after_trim INTEGER DEFAULT 0,
75
+ trim_saved INTEGER DEFAULT 0,
76
+ schema_tokens_full INTEGER DEFAULT 0,
77
+ schema_tokens_selected INTEGER DEFAULT 0,
78
+ schema_tokens_saved INTEGER DEFAULT 0,
79
+ groq_prompt_tokens INTEGER DEFAULT 0,
80
+ groq_completion_tokens INTEGER DEFAULT 0,
81
+ doc_id TEXT,
82
+ success INTEGER DEFAULT 1
83
+ )
84
+ """)
85
+ for col in ["user_id", "conversation_id"]:
86
+ try:
87
+ conn.execute(f"ALTER TABLE events ADD COLUMN {col} TEXT")
88
+ except Exception:
89
+ pass
90
+ conn.commit()
91
+ logger.info(f"[DB] SQLite ready at {DB_PATH}")
92
+
93
+
94
+ def _post_to_cloud(data: dict):
95
+ """Send event to EC2 API"""
96
+ try:
97
+ requests.post(INGEST_URL,json=data,timeout=3,)
98
+ except Exception as e:
99
+ logger.debug(f"[DB] Cloud ingest failed (local SQLite has it): {e}")
100
+
101
+
102
+ def insert_event(**kwargs):
103
+ conn = get_thread_connection()
104
+ kwargs.setdefault("user_id", os.getenv("TOKEN_USER_ID", "unknown"))
105
+ kwargs.setdefault("conversation_id", _get_or_create_conversation_id())
106
+
107
+ fields = [
108
+ "user_id", "conversation_id",
109
+ "tool_name", "query", "cache_hit", "cache_similarity",
110
+ "tokens_before_trim", "tokens_after_trim", "trim_saved",
111
+ "schema_tokens_full", "schema_tokens_selected", "schema_tokens_saved",
112
+ "doc_id", "success",
113
+ "groq_prompt_tokens", "groq_completion_tokens",
114
+ ]
115
+
116
+ data = {f: kwargs.get(f, None) for f in fields}
117
+
118
+ # always write to local SQLite
119
+ placeholders = ", ".join(["?" for _ in fields])
120
+ columns = ", ".join(fields)
121
+ conn.execute(
122
+ f"INSERT INTO events ({columns}) VALUES ({placeholders})",
123
+ list(data.values()),
124
+ )
125
+ conn.commit()
126
+
127
+ _post_to_cloud(data)
128
+
129
+
130
+ def _fetchall(conn, sql, params=()):
131
+ rows = conn.execute(sql, params).fetchall()
132
+ return [dict(r) for r in rows]
133
+
134
+
135
+
136
+ def _fetchone(conn, sql, params=()):
137
+ row = conn.execute(sql, params).fetchone()
138
+ return dict(row) if row else {}
139
+
140
+
141
+ def get_summary(conn):
142
+ return _fetchone(conn, """
143
+ SELECT
144
+ COUNT(*) as total_calls,
145
+ COALESCE(SUM(cache_hit), 0) as cache_hits,
146
+ COALESCE(ROUND(AVG(cache_hit) * 100, 1), 0.0) as hit_rate_pct,
147
+ COALESCE(SUM(trim_saved), 0) as total_trim_saved,
148
+ COALESCE(SUM(schema_tokens_saved), 0) as total_schema_saved,
149
+ COALESCE(SUM(trim_saved + schema_tokens_saved), 0) as total_tokens_saved
150
+ FROM events
151
+ """)
152
+
153
+
154
+
155
+ def get_recent_events(conn, limit=50):
156
+ is_pg = "psycopg2" in type(conn).__module__
157
+ if is_pg:
158
+ return _fetchall(conn, "SELECT * FROM events ORDER BY timestamp DESC LIMIT %s", (limit,))
159
+ else:
160
+ return _fetchall(conn, "SELECT * FROM events ORDER BY timestamp DESC LIMIT ?", (limit,))
161
+
162
+
163
+ def get_tool_stats(conn):
164
+ return _fetchall(conn,"""
165
+ SELECT
166
+ tool_name,
167
+ COUNT(*) AS calls,
168
+ SUM(cache_hit) AS hits,
169
+ SUM(trim_saved) AS trim_saved,
170
+ SUM(schema_tokens_saved) AS schema_saved
171
+ FROM events
172
+ GROUP BY tool_name
173
+ ORDER BY calls DESC
174
+ """)
175
+
176
+
177
+ def get_token_analysis(conn):
178
+ rows = _fetchone(conn, """
179
+ SELECT
180
+ COUNT(*) AS total_queries,
181
+ COALESCE(SUM(schema_tokens_saved), 0) AS schema_saved,
182
+ COALESCE(SUM(schema_tokens_full), 0) AS schema_full,
183
+ COALESCE(SUM(schema_tokens_selected), 0) AS schema_selected,
184
+ COALESCE(SUM(trim_saved), 0) AS trim_saved,
185
+ COALESCE(SUM(tokens_before_trim), 0) AS tokens_before_trim,
186
+ COALESCE(SUM(tokens_after_trim), 0) AS tokens_after_trim,
187
+ COALESCE(SUM(cache_hit), 0) AS cache_hits,
188
+ COUNT(CASE WHEN cache_hit = 0 THEN 1 END) as cache_misses
189
+ FROM events
190
+ """)
191
+
192
+ schema_saved = rows.get("schema_saved",0)
193
+ schema_full = rows.get("schema_full",0)
194
+ schema_selected = rows.get("schema_selected",0)
195
+ trim_saved = rows.get("trim_saved",0)
196
+ tokens_before = rows.get("tokens_before_trim",0)
197
+ tokens_after = rows.get("tokens_after_trim",0)
198
+ cache_hits = rows.get("cache_hits",0)
199
+ cache_misses = rows.get("cache_misses",0)
200
+ total_queries = rows.get("total_queries",0)
201
+
202
+ total_saved = schema_saved + trim_saved
203
+
204
+ actual_without= schema_full + tokens_before
205
+ actual_with = schema_selected + tokens_after
206
+
207
+ pct_saved = round((total_saved / actual_without * 100), 1) if actual_without > 0 else 0.0
208
+
209
+ return {
210
+ "total_queries": total_queries,
211
+ "cache_hits": cache_hits,
212
+ "cache_misses": cache_misses,
213
+ "schema_tokens_without_tom": schema_full,
214
+ "schema_tokens_with_tom": schema_selected,
215
+ "schema_saved": schema_saved,
216
+ "tokens_before_trim": tokens_before,
217
+ "tokens_after_trim": tokens_after,
218
+ "trim_saved": trim_saved,
219
+ "total_saved": total_saved,
220
+ "actual_without_tom": actual_without,
221
+ "actual_with_tom": actual_with,
222
+ "pct_saved": pct_saved,
223
+ }
224
+
225
+ def get_event_count(conn) -> int:
226
+ """Quick row count for health check."""
227
+ row = _fetchone(conn, "SELECT COUNT(*) AS cnt FROM events")
228
+ return row.get("cnt", 0)
229
+
230
+ def get_conversation_stats(conn, limit=20):
231
+ is_pg = "psycopg2" in type(conn).__module__
232
+ if is_pg:
233
+ return _fetchall(conn, """
234
+ SELECT
235
+ conversation_id,
236
+ MIN(timestamp) AS started_at,
237
+ COUNT(*) AS total_calls,
238
+ SUM(cache_hit) AS cache_hits,
239
+ SUM(trim_saved) AS trim_saved,
240
+ SUM(schema_tokens_saved) AS schema_saved,
241
+ SUM(trim_saved + schema_tokens_saved) AS total_saved
242
+ FROM events
243
+ WHERE conversation_id IS NOT NULL
244
+ GROUP BY conversation_id
245
+ ORDER BY started_at DESC
246
+ LIMIT %s
247
+ """, (limit,))
248
+ else:
249
+ return _fetchall(conn, """
250
+ SELECT
251
+ conversation_id,
252
+ MIN(timestamp) AS started_at,
253
+ COUNT(*) AS total_calls,
254
+ SUM(cache_hit) AS cache_hits,
255
+ SUM(trim_saved) AS trim_saved,
256
+ SUM(schema_tokens_saved) AS schema_saved,
257
+ SUM(trim_saved + schema_tokens_saved) AS total_saved
258
+ FROM events
259
+ WHERE conversation_id IS NOT NULL
260
+ GROUP BY conversation_id
261
+ ORDER BY started_at DESC
262
+ LIMIT ?
263
+ """, (limit,))
264
+
265
+
266
+
267
+ def get_groq_usage(conn):
268
+ return _fetchone(conn, """
269
+ SELECT
270
+ COALESCE(SUM(groq_prompt_tokens), 0) AS total_prompt_tokens,
271
+ COALESCE(SUM(groq_completion_tokens), 0) AS total_completion_tokens,
272
+ COALESCE(SUM(groq_prompt_tokens + groq_completion_tokens), 0) AS total_groq_tokens
273
+ FROM events
274
+ """)
275
+
276
+
277
+ def get_timeseries(conn, hours: int = 24):
278
+ is_pg = "psycopg2" in type(conn).__module__
279
+ if is_pg:
280
+ return _fetchall(conn, """
281
+ SELECT
282
+ to_char(date_trunc('hour', timestamp), 'YYYY-MM-DD HH24:00') AS hour,
283
+ COUNT(*) AS total_calls,
284
+ COALESCE(SUM(cache_hit), 0) AS cache_hits,
285
+ COALESCE(SUM(schema_tokens_saved), 0) AS schema_saved,
286
+ COALESCE(SUM(trim_saved), 0) AS trim_saved,
287
+ COALESCE(SUM(groq_prompt_tokens + groq_completion_tokens), 0) AS groq_tokens
288
+ FROM events
289
+ WHERE timestamp >= NOW() - INTERVAL '%s hours'
290
+ GROUP BY date_trunc('hour', timestamp)
291
+ ORDER BY date_trunc('hour', timestamp) ASC
292
+ """, (hours,))
293
+ else:
294
+ return _fetchall(conn, """
295
+ SELECT
296
+ strftime('%Y-%m-%d %H:00', timestamp) AS hour,
297
+ COUNT(*) AS total_calls,
298
+ COALESCE(SUM(cache_hit), 0) AS cache_hits,
299
+ COALESCE(SUM(schema_tokens_saved), 0) AS schema_saved,
300
+ COALESCE(SUM(trim_saved), 0) AS trim_saved,
301
+ COALESCE(SUM(groq_prompt_tokens + groq_completion_tokens), 0) AS groq_tokens
302
+ FROM events
303
+ WHERE timestamp >= datetime('now', ? || ' hours')
304
+ GROUP BY hour
305
+ ORDER BY hour ASC
306
+ """, (f"-{hours}",))
307
+
308
+
309
+ def get_period_summary(conn, period: str = "today"):
310
+ is_pg = "psycopg2" in type(conn).__module__
311
+ if is_pg:
312
+ if period == "today":
313
+ where = "DATE(timestamp) = CURRENT_DATE"
314
+ elif period == "yesterday":
315
+ where = "DATE(timestamp) = CURRENT_DATE - INTERVAL '1 day'"
316
+ elif period == "week":
317
+ where = "timestamp >= NOW() - INTERVAL '7 days'"
318
+ else:
319
+ where = "1=1"
320
+ else:
321
+ if period == "today":
322
+ where = "date(timestamp) = date('now')"
323
+ elif period == "yesterday":
324
+ where = "date(timestamp) = date('now', '-1 day')"
325
+ elif period == "week":
326
+ where = "timestamp >= datetime('now', '-7 days')"
327
+ else:
328
+ where = "1=1"
329
+
330
+ return _fetchone(conn, f"""
331
+ SELECT
332
+ COUNT(*) AS total_calls,
333
+ COALESCE(SUM(cache_hit), 0) AS cache_hits,
334
+ COALESCE(SUM(schema_tokens_saved + trim_saved), 0) AS total_saved,
335
+ COALESCE(SUM(groq_prompt_tokens), 0) AS groq_prompt,
336
+ COALESCE(SUM(groq_completion_tokens), 0) AS groq_completion
337
+ FROM events
338
+ WHERE {where}
339
+ """)
@@ -0,0 +1,159 @@
1
+ from config import settings, PROJECT_ROOT
2
+ import os, logging,re
3
+
4
+ logger = logging.getLogger("token")
5
+ CHROMA_PATH = os.path.join(PROJECT_ROOT, "storage", "chroma_db")
6
+
7
+ _embeddings = None
8
+ _splitter = None
9
+ _vectorstore = None
10
+ _reranker = None
11
+
12
+ def _get_vectorstore():
13
+ global _embeddings, _splitter, _vectorstore, _reranker
14
+
15
+ if _vectorstore is None:
16
+
17
+ from langchain_text_splitters import RecursiveCharacterTextSplitter
18
+ from langchain_huggingface import HuggingFaceEmbeddings
19
+ from langchain_chroma import Chroma
20
+ from langchain_community.cross_encoders import HuggingFaceCrossEncoder
21
+ from langchain_classic.retrievers.document_compressors.cross_encoder_rerank import CrossEncoderReranker
22
+
23
+
24
+ _embeddings = HuggingFaceEmbeddings(model_name=settings.embedder)
25
+ _splitter = RecursiveCharacterTextSplitter(chunk_size=1200, chunk_overlap=200)
26
+ _vectorstore = Chroma(collection_name="document", embedding_function=_embeddings, persist_directory=CHROMA_PATH)
27
+ _reranker_model = HuggingFaceCrossEncoder(model_name="cross-encoder/ms-marco-MiniLM-L-6-v2")
28
+ _reranker = CrossEncoderReranker(model=_reranker_model, top_n=3)
29
+
30
+ return _vectorstore
31
+
32
+
33
+ def list_index_doc():
34
+ """List the documents that are already indexed. this is called before index doc"""
35
+ try:
36
+ result = _get_vectorstore().get(include=['metadatas'])
37
+ seen = {}
38
+ for meta in result['metadatas']:
39
+ doc_id = meta.get("doc_id")
40
+ source = meta.get("source", "unknown")
41
+ if doc_id and doc_id not in seen:
42
+ seen[doc_id] = source
43
+ return seen
44
+ except Exception as e:
45
+ logger.error(f"[LIST_DOC] empty : {e}")
46
+ return {}
47
+
48
+
49
+ def index_doc(file_path, doc_id):
50
+ """index the doc uploaded by the user"""
51
+ from langchain_community.document_loaders import PyMuPDFLoader
52
+ from langchain_core.documents import Document
53
+
54
+ file_path = os.path.expanduser(file_path)
55
+ if not file_path or not file_path.strip():
56
+ raise ValueError("file_path cannot be empty")
57
+ if not os.path.exists(file_path):
58
+ raise FileNotFoundError(f"File not found: {file_path}")
59
+ if os.path.getsize(file_path) == 0:
60
+ raise ValueError(f"File is empty: {file_path}")
61
+ if not doc_id or not doc_id.strip():
62
+ raise ValueError("doc_id cannot be empty")
63
+
64
+ # Validate file extension — only PDFs are supported
65
+ if not file_path.lower().endswith(".pdf"):
66
+ raise ValueError(f"Only PDF files are supported. Got: {os.path.splitext(file_path)[1] or '(no extension)'}")
67
+
68
+
69
+ doc_id = re.sub(r"[^a-zA-Z0-9_\-]", "_", doc_id.strip())
70
+ if not doc_id:
71
+ raise ValueError("doc_id became empty after sanitization — use alphanumeric characters")
72
+
73
+ vs = _get_vectorstore()
74
+ loader = PyMuPDFLoader(file_path)
75
+ pages = loader.load()
76
+ full_chunks = []
77
+ chunk_index = 0
78
+
79
+ for page in pages:
80
+ page_number = page.metadata.get("page", 0) + 1
81
+ chunks = _splitter.split_text(page.page_content)
82
+ for chunk in chunks:
83
+ full_chunks.append(
84
+ Document(page_content=chunk,
85
+ metadata={"doc_id": doc_id, "chunk_index": chunk_index,
86
+ "page": page_number, "source": file_path}))
87
+ chunk_index += 1
88
+ ids = [f"{doc_id}_{i}" for i in range(len(full_chunks))]
89
+ try:
90
+ existing = vs.get(where={"doc_id": doc_id})
91
+ if existing["ids"]:
92
+ vs.delete(ids=existing["ids"])
93
+ logger.info(f"[INDEX] removed {len(existing['ids'])} old chunks - {doc_id}")
94
+ except Exception as e:
95
+ logger.warning(f"[INDEX] could not clean old chunks for {doc_id}: {e}")
96
+ vs.add_documents(full_chunks, ids=ids)
97
+ return len(full_chunks)
98
+
99
+
100
+
101
+ def search_doc(query, doc_id, top_k=3, rerank=10):
102
+ from langchain_classic.retrievers import ContextualCompressionRetriever
103
+
104
+ vs = _get_vectorstore()
105
+ search_kwargs = {'k': rerank}
106
+ if doc_id:
107
+ search_kwargs['filter'] = {'doc_id': doc_id}
108
+ base_retriever = vs.as_retriever(search_kwargs=search_kwargs)
109
+ compress_retriever = ContextualCompressionRetriever(
110
+ base_compressor=_reranker,
111
+ base_retriever=base_retriever
112
+ )
113
+ results = compress_retriever.invoke(query)
114
+ return [
115
+ {"text": doc.page_content, "chunk_index": doc.metadata.get("chunk_index"), "page": doc.metadata.get("page")}
116
+ for doc in results[:top_k]
117
+ ]
118
+
119
+
120
+
121
+ def search_all_doc(query, top_k=3, rerank=10):
122
+ from langchain_classic.retrievers import ContextualCompressionRetriever
123
+
124
+
125
+ vs = _get_vectorstore()
126
+ base_retriever = vs.as_retriever(search_kwargs={'k': rerank})
127
+ compress_retriever = ContextualCompressionRetriever(
128
+ base_compressor=_reranker,
129
+ base_retriever=base_retriever
130
+ )
131
+ results = compress_retriever.invoke(query)
132
+ return [
133
+ {'text': doc.page_content, 'doc_id': doc.metadata.get("doc_id"),
134
+ 'page': doc.metadata.get("page"), 'source': doc.metadata.get("source")}
135
+ for doc in results[:top_k]
136
+ ]
137
+
138
+
139
+
140
+ def index_folder(folder_path):
141
+ results = {}
142
+ if not os.path.isdir(folder_path):
143
+ raise ValueError(f"Not a directory: {folder_path}")
144
+ pdf_files = [f for f in os.listdir(folder_path) if f.lower().endswith(".pdf")]
145
+ if not pdf_files:
146
+ return {}
147
+ for filename in pdf_files:
148
+ file_path = os.path.join(folder_path, filename)
149
+ doc_id = os.path.splitext(filename)[0]
150
+ try:
151
+ chunks = index_doc(file_path, doc_id)
152
+ results[doc_id] = chunks
153
+ logger.info(f"Indexed {doc_id} | {chunks} chunks")
154
+ except Exception as e:
155
+ logger.error(f"Failed to index {filename}: {e}")
156
+ results[doc_id] = 0
157
+ return results
158
+
159
+