import os import gc import csv import shutil import tempfile import logging import asyncio import traceback from concurrent.futures import ThreadPoolExecutor from contextlib import asynccontextmanager import httpx import lancedb import pandas as pd import pyarrow as pa from fastapi import FastAPI, File, UploadFile, Form, HTTPException, Depends, Security from fastapi.security.api_key import APIKeyHeader from pydantic import BaseModel from sentence_transformers import SentenceTransformer from pypdf import PdfReader from conversation import ( build_general_answer_prompt, RAG_CONVERSATION_TURNS, build_rag_answer_prompt, build_retrieval_query, build_scope_classification_prompt, format_conversation_block, ) from llm_client import ( LLM_PROVIDER, OLLAMA_MODEL, OLLAMA_NUM_PREDICT, OLLAMA_READ_TIMEOUT, OLLAMA_URL, llm_generate, llm_provider_info, ) from reranker import ( RERANK_ENABLED, RERANK_MIN_SCORE, build_rerank_query, get_reranker, load_reranker, rerank_results, reranker_info, retrieval_meta_from_results, should_skip_rerank, unload_reranker, ) from guardrails import ( blocked_response, build_guardrails_meta, check_input, decide_answer_mode, guardrails_info, is_conversational_ack, is_llm_refusal, is_scope_refusal, is_unusable_assistant_output, not_in_documents_message, chunks_likely_answer_question, out_of_scope_refusal_message, parse_scope_classification, sanitize_output, ) logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Constants # --------------------------------------------------------------------------- BUCKET_PATH = "/data" # HF Storage Bucket (FUSE mount, slow) LOCAL_DB_PATH = "/tmp/lancedb_store" # Local disk (fast, ephemeral) BUCKET_DB_PATH = os.path.join(BUCKET_PATH, "lancedb_store") TABLE_NAME = "documents" def _normalize_ollama_base_url(raw: str) -> str: """ OLLAMA_URL must be the Ollama **base** origin only (no /api/... path). Strip accidental suffixes users paste from OpenAI-compatible clients. """ u = (raw or "").strip().rstrip("/") for suffix in ( "/v1/chat/completions", "/v1/chat", "/v1", "/api/generate", "/api/tags", "/api", ): if u.lower().endswith(suffix.lower()): u = u[: -len(suffix)].rstrip("/") return u def ollama_url_config_issue(url: str) -> str | None: """ Detect common misconfigurations that yield HTML 404 from HF (not Ollama JSON). Returns a short hint for logs /health, or None if the shape looks OK. """ if not url: return "OLLAMA_URL is empty — set it in the rag-api Space secrets to your Ollama Space Direct URL." low = url.lower() if "huggingface.co" in low and "/spaces/" in low: return ( "OLLAMA_URL looks like a huggingface.co Spaces **page** URL. " "Use the **Direct URL** from the Ollama Space → Settings (ends with `.hf.space`), not huggingface.co/spaces/…" ) if "huggingface.co" in low: return "OLLAMA_URL must not point at huggingface.co — use the Ollama Space Direct `.hf.space` URL." return None # Health uses OLLAMA_URL from llm_client (normalized there). _ollama_issue = ollama_url_config_issue(OLLAMA_URL) if _ollama_issue and LLM_PROVIDER == "ollama": logger.error("OLLAMA_URL configuration: %s Current value: %s", _ollama_issue, OLLAMA_URL) # Vector retrieval — wide recall; reranker narrows to RAG_CONTEXT_CHUNKS RAG_RETRIEVE_K = max( 1, min(int(os.environ.get("RAG_RETRIEVE_K", os.environ.get("RAG_TOP_K", "24"))), 32), ) RAG_TOP_K = RAG_RETRIEVE_K # backward-compatible alias for health / docs # Chunks passed to the LLM after rerank or distance narrowing RAG_CONTEXT_CHUNKS = max( 1, min(int(os.environ.get("RAG_CONTEXT_CHUNKS", "6")), RAG_RETRIEVE_K) ) # Drop retrieved rows much worse than the best hit (distance-only path when rerank skipped). # Set RAG_RELATIVE_DISTANCE_CAP=0 to disable filtering (only cap by RAG_CONTEXT_CHUNKS). RAG_RELATIVE_DISTANCE_CAP = float(os.environ.get("RAG_RELATIVE_DISTANCE_CAP", "1.5")) # LLM generation limits — re-exported from llm_client for / health RAG_MAX_CONTEXT_CHARS = int(os.environ.get("RAG_MAX_CONTEXT_CHARS", "4500")) LLM_FALLBACK_CONTEXT_MAX_CHARS = int(os.environ.get("LLM_FALLBACK_CONTEXT_MAX_CHARS", "2800")) ALLOWED_EXTENSIONS = {".pdf", ".txt", ".csv", ".docx"} # Shared secret for API authentication. Set RAG_API_SECRET as a HF Space secret. # Leave empty only for local dev (no auth enforced when unset). RAG_API_SECRET = os.environ.get("RAG_API_SECRET", "") _api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False) def _require_secret(key: str = Security(_api_key_header)): if RAG_API_SECRET and key != RAG_API_SECRET: raise HTTPException(status_code=401, detail="Unauthorized") # --------------------------------------------------------------------------- # Global handles # --------------------------------------------------------------------------- db: lancedb.DBConnection = None embed_model: SentenceTransformer = None def sync_from_bucket(): """Copy LanceDB data from bucket to local disk on startup.""" try: if os.path.exists(BUCKET_DB_PATH) and os.listdir(BUCKET_DB_PATH): logger.info("Syncing LanceDB from bucket to local disk...") if os.path.exists(LOCAL_DB_PATH): shutil.rmtree(LOCAL_DB_PATH) shutil.copytree(BUCKET_DB_PATH, LOCAL_DB_PATH) logger.info("Sync from bucket complete") return except OSError as e: logger.warning(f"Bucket read failed ({e}), starting fresh") os.makedirs(LOCAL_DB_PATH, exist_ok=True) logger.info("No existing data in bucket, starting fresh") def sync_to_bucket(): """Incremental sync: only copy new/modified files from local to bucket.""" try: os.makedirs(BUCKET_DB_PATH, exist_ok=True) # Build set of local files (relative paths) local_files = set() for root, dirs, files in os.walk(LOCAL_DB_PATH): for fname in files: src = os.path.join(root, fname) rel = os.path.relpath(src, LOCAL_DB_PATH) local_files.add(rel) dst = os.path.join(BUCKET_DB_PATH, rel) os.makedirs(os.path.dirname(dst), exist_ok=True) # Copy only if dest missing or source is newer/different size if not os.path.exists(dst): shutil.copy2(src, dst) else: src_stat = os.stat(src) dst_stat = os.stat(dst) if (src_stat.st_size != dst_stat.st_size or src_stat.st_mtime > dst_stat.st_mtime): shutil.copy2(src, dst) # Remove bucket files that no longer exist locally (e.g. compacted away) for root, dirs, files in os.walk(BUCKET_DB_PATH): for fname in files: dst = os.path.join(root, fname) rel = os.path.relpath(dst, BUCKET_DB_PATH) if rel not in local_files: os.remove(dst) logger.info(f"Removed stale bucket file: {rel}") logger.info("Incremental sync to bucket complete") except Exception as e: logger.error(f"Failed to sync to bucket: {e}") _sync_executor = ThreadPoolExecutor(max_workers=1) _rerank_executor = ThreadPoolExecutor(max_workers=1) _sync_lock = asyncio.Lock() async def sync_to_bucket_async(): """Run sync in background thread so uploads return immediately.""" if _sync_lock.locked(): logger.info("Sync already in progress, skipping") return async with _sync_lock: loop = asyncio.get_event_loop() await loop.run_in_executor(_sync_executor, sync_to_bucket) KEEP_ALIVE_INTERVAL = 5 * 60 # 5 minutes # Public URL — HF only counts external traffic to prevent sleep SELF_PUBLIC_URL = os.environ.get( "SPACE_URL", "https://idnameraj-rag-vs.hf.space" ) async def keep_alive_loop(): """Ping own public URL every 5 min so HF counts it as external traffic.""" while True: await asyncio.sleep(KEEP_ALIVE_INTERVAL) try: async with httpx.AsyncClient(verify=False) as client: resp = await client.get(f"{SELF_PUBLIC_URL}/", timeout=15) logger.info(f"Keep-alive ping via public URL: {resp.status_code}") except Exception as e: # Fallback to localhost so the app at least stays warm internally try: async with httpx.AsyncClient(verify=False) as client: await client.get("http://localhost:7860/", timeout=10) logger.info("Keep-alive fallback (localhost) OK") except Exception: logger.warning(f"Keep-alive failed: {e}") @asynccontextmanager async def lifespan(app: FastAPI): global db, embed_model # Restore data from bucket to fast local disk sync_from_bucket() db = lancedb.connect(LOCAL_DB_PATH) migrate_documents_schema_if_needed() embed_model = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2") load_reranker() # Start keep-alive background task keep_alive_task = asyncio.create_task(keep_alive_loop()) logger.info( "Startup complete (LLM_PROVIDER=%s, rerank=%s)", LLM_PROVIDER, RERANK_ENABLED, ) if LLM_PROVIDER == "ollama" and ollama_url_config_issue(OLLAMA_URL): logger.error("Fix OLLAMA_URL on this Space — see / health field ollama_url_issue") yield # Cancel keep-alive and persist to bucket keep_alive_task.cancel() sync_to_bucket() unload_reranker() del embed_model, db gc.collect() app = FastAPI(title="RAG API", lifespan=lifespan) @app.get("/") async def health(): ollama_ok = False ollama_http: int | str = "n/a" if LLM_PROVIDER == "ollama": try: async with httpx.AsyncClient(verify=False) as client: r = await client.get(f"{OLLAMA_URL}/api/tags", timeout=10) ollama_http = r.status_code ollama_ok = r.status_code == 200 except Exception as e: ollama_http = str(e)[:120] table_count = 0 if TABLE_NAME in get_table_names(): try: table_count = db.open_table(TABLE_NAME).count_rows() except Exception: pass return { "status": "ok", **llm_provider_info(), "ollama": ( "connected" if ollama_ok else ("n/a" if LLM_PROVIDER != "ollama" else "unreachable") ), "ollama_http": ollama_http, "ollama_url_issue": ollama_url_config_issue(OLLAMA_URL) if LLM_PROVIDER == "ollama" else None, "ollama_read_timeout_s": OLLAMA_READ_TIMEOUT, "ollama_num_predict": OLLAMA_NUM_PREDICT, "rag_retrieve_k": RAG_RETRIEVE_K, "rag_top_k": RAG_TOP_K, "rag_context_chunks": RAG_CONTEXT_CHUNKS, "rag_relative_distance_cap": RAG_RELATIVE_DISTANCE_CAP, **reranker_info(), **guardrails_info(), "rag_conversation_turns": RAG_CONVERSATION_TURNS, "accepted_formats": list(ALLOWED_EXTENSIONS), "chunks_in_db": table_count, } # --------------------------------------------------------------------------- # Text extraction # --------------------------------------------------------------------------- def extract_text_from_pdf(filepath: str) -> str: reader = PdfReader(filepath) return "\n".join(page.extract_text() or "" for page in reader.pages) def extract_text_from_txt(filepath: str) -> str: with open(filepath, "r", encoding="utf-8", errors="ignore") as f: return f.read() def extract_text_from_csv(filepath: str) -> str: with open(filepath, "r", encoding="utf-8", errors="ignore") as f: reader = csv.reader(f) rows = [", ".join(row) for row in reader] return "\n".join(rows) def extract_text_from_docx(filepath: str) -> str: from docx import Document doc = Document(filepath) return "\n".join(para.text for para in doc.paragraphs) EXTRACTORS = { ".pdf": extract_text_from_pdf, ".txt": extract_text_from_txt, ".csv": extract_text_from_csv, ".docx": extract_text_from_docx, } # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def get_table_names() -> list[str]: """Get list of table names, handling both old and new LanceDB API.""" result = db.list_tables() if hasattr(result, "tables"): return result.tables return list(result) def documents_table_schema_field_names() -> set[str]: """Column names on the documents table (empty if missing).""" if db is None or TABLE_NAME not in get_table_names(): return set() try: return {f.name for f in db.open_table(TABLE_NAME).schema} except Exception as e: logger.warning("Could not read documents schema: %s", e) return set() def migrate_documents_schema_if_needed() -> None: """ Tables created before source_id was added have no such column; deletes and appends then fail. Add the column (null/empty for existing rows) or rewrite. """ global db if db is None or TABLE_NAME not in get_table_names(): return try: table = db.open_table(TABLE_NAME) names = {f.name for f in table.schema} except Exception as e: logger.warning("migrate: could not inspect table: %s", e) return if "source_id" in names: return logger.info("Migrating LanceDB table %r: adding source_id (legacy schema)", TABLE_NAME) try: table.add_columns(pa.field("source_id", pa.string())) logger.info("Migration: add_columns(source_id) complete") except Exception as e: logger.warning("add_columns failed (%s); rewriting table", e) _migrate_documents_table_rewrite() def _migrate_documents_table_rewrite() -> None: """Last-resort migration: materialize full table with source_id column.""" global db tbl = db.open_table(TABLE_NAME) df = tbl.to_pandas() df["source_id"] = "" db.drop_table(TABLE_NAME) db.create_table(TABLE_NAME, df) logger.info("Migration: full table rewrite with source_id complete") def _sql_string_literal(val: str) -> str: """Escape a value for use inside LanceDB / DataFusion SQL string literals.""" return "'" + str(val).replace("'", "''") + "'" def _lance_delete_where(predicate: str) -> None: """Run table.delete(predicate); logs and swallows errors for legacy schemas.""" if TABLE_NAME not in get_table_names(): return table = db.open_table(TABLE_NAME) try: table.delete(predicate) except Exception as e: logger.warning(f"LanceDB delete skipped or failed ({e!s}); predicate={predicate[:200]}") def delete_vectors_exact(tenant_id: str, project_id: str, source_id: str) -> None: t, p, s = _sql_string_literal(tenant_id), _sql_string_literal(project_id), _sql_string_literal(source_id) _lance_delete_where(f"tenant_id = {t} AND project_id = {p} AND source_id = {s}") def delete_vectors_prefix(tenant_id: str, project_id: str, source_id_prefix: str) -> None: """Delete rows where source_id starts with prefix (DocChat manual Q&A per-chunk keys).""" t, p = _sql_string_literal(tenant_id), _sql_string_literal(project_id) pref = _sql_string_literal(source_id_prefix) _lance_delete_where( f"tenant_id = {t} AND project_id = {p} AND starts_with(source_id, {pref})" ) def delete_vectors_project(tenant_id: str, project_id: str) -> None: t, p = _sql_string_literal(tenant_id), _sql_string_literal(project_id) _lance_delete_where(f"tenant_id = {t} AND project_id = {p}") def chunk_text(text: str, chunk_size: int = 500, overlap: int = 50) -> list[str]: chunks = [] start = 0 while start < len(text): end = start + chunk_size chunks.append(text[start:end]) start = end - overlap return [c.strip() for c in chunks if c.strip()] def _truncate_context(context: str, max_chars: int = RAG_MAX_CONTEXT_CHARS) -> str: """Keep prompt size bounded so CPU Ollama finishes within read timeout.""" context = (context or "").strip() if len(context) <= max_chars: return context cut = context[:max_chars] last_break = max(cut.rfind("\n\n"), cut.rfind(". ")) if last_break > max_chars // 2: cut = cut[: last_break + 1] return cut.rstrip() + "\n\n[... context truncated ...]" def _excerpt_answer_from_context(context: str) -> str: """When the LLM refuses but we retrieved chunks, return a direct excerpt.""" excerpt = _truncate_context(context, LLM_FALLBACK_CONTEXT_MAX_CHARS) return ( "Here is the most relevant information from your documents:\n\n" f"{excerpt}" ) SCOPE_CLASSIFIER_ENABLED = os.environ.get("RAG_SCOPE_CLASSIFIER_ENABLED", "true").lower() in ( "1", "true", "yes", ) async def classify_request_out_of_scope(question: str) -> bool: """LLM scope check — works for any document-only chatbot domain.""" if not SCOPE_CLASSIFIER_ENABLED: return False try: raw = await llm_generate( build_scope_classification_prompt(question), user_question=question, ) return parse_scope_classification(raw) == "out_of_scope" except Exception as e: logger.warning("Scope classification failed (%s) — defaulting to in-scope", e) return False async def _resolve_grounded_refusal( question: str, llm_answer: str, raw_context: str, source_count: int, retrieval_meta: dict | None = None, ) -> tuple[str, dict]: """Choose scope refusal vs document excerpt after a grounded LLM refusal.""" meta: dict = {} retrieval_meta = retrieval_meta or {} if is_scope_refusal(llm_answer): logger.info("LLM returned document-scope refusal for %s retrieved chunks", source_count) meta["scope_refusal"] = True return llm_answer.strip(), meta if await classify_request_out_of_scope(question): logger.info( "Scope classifier marked request out-of-scope (%s chunks) — declining", source_count, ) meta["out_of_scope_refusal"] = True return out_of_scope_refusal_message(), meta relevance_ok = retrieval_meta.get("relevance_gate_passed", True) chunks_relevant = chunks_likely_answer_question(question, raw_context) if not relevance_ok or not chunks_relevant: logger.info( "Retrieved chunks unrelated to question (gate=%s, overlap=%s) — not excerpting", relevance_ok, chunks_relevant, ) meta["not_in_documents"] = True return not_in_documents_message(), meta logger.info( "LLM refused despite %s relevant retrieved chunks — using context excerpt", source_count, ) meta["llm_refusal_excerpt"] = True return _excerpt_answer_from_context(raw_context), meta def _package_query_response( answer: str, sources: list, retrieval_meta: dict, *, translated: bool = False, llm_fallback: bool = False, ) -> dict: mode = retrieval_meta.get("mode", "grounded") payload = { "answer": sanitize_output(answer, mode=mode, sources=sources), "sources": sources, "translated": translated, "retrieval": retrieval_meta, "guardrails": build_guardrails_meta(blocked=False), } if llm_fallback: payload["llm_fallback"] = True return payload async def _general_llm_fallback_answer( question: str, chat_history: list, session_facts: list[str], answer_language: str | None, retrieval_meta: dict, ) -> dict: """Fallback to direct LLM answering when retrieval has no useful grounding.""" conversation_block = format_conversation_block(chat_history, session_facts) prompt = build_general_answer_prompt(question, conversation_block) answer = await llm_generate(prompt, user_question=question) final_answer = answer.strip() if not final_answer or is_unusable_assistant_output(final_answer): final_answer = ( "I'm here to help. Could you tell me a bit more about what you're looking for?" ) target_lang = (answer_language or "").strip().lower()[:12] translated = False if target_lang and target_lang != "en" and final_answer: try: final_answer = await llm_translate(final_answer, target_lang) translated = True except Exception as e: logger.warning("General fallback translation failed for lang=%s (%s)", target_lang, e) retrieval_meta = { **retrieval_meta, "mode": "general", "general_llm_fallback": True, "relevance_gate_passed": False, } return _package_query_response( final_answer, [], retrieval_meta, translated=translated, ) def _fallback_answer_from_context(context: str) -> str: """When Ollama times out, return a short excerpt instead of HTTP 502.""" excerpt = _truncate_context(context, LLM_FALLBACK_CONTEXT_MAX_CHARS) return ( "The language model took too long to respond. Here are the most relevant " f"passages from your documents:\n\n{excerpt}" ) def _fallback_answer_llm_http(context: str, status_code: int) -> str: """When the LLM gateway returns HTTP errors — provider down or misconfigured.""" excerpt = _truncate_context(context, LLM_FALLBACK_CONTEXT_MAX_CHARS) if LLM_PROVIDER == "openrouter": if status_code in (401, 403): hint = ( f"OpenRouter returned HTTP {status_code}. Check OPENROUTER_API_KEY on the rag-api Space " "(Repository secrets) and that the key is valid." ) elif status_code == 404: hint = ( "OpenRouter returned HTTP 404. Check OPENROUTER_MODEL " f"(current: {os.environ.get('OPENROUTER_MODEL', 'openrouter/free')})." ) else: hint = ( f"OpenRouter returned HTTP {status_code}. The service may be overloaded; " "check https://openrouter.ai/status or try again." ) elif status_code == 404: hint = ( "The Ollama Hugging Face Space returned HTTP 404 (HTML error page, not the Ollama API). " "Fix: open your Ollama Space on huggingface.co, wait until it shows **Running** (not Building/Paused), " "then in Space **Settings** copy the **Direct URL** (ends with `.hf.space`) and set it as " "`OLLAMA_URL` on the rag-api Space. Do not use `huggingface.co/spaces/...` as OLLAMA_URL. " "The value must be the **base** URL only (no `/api` or `/v1` path). " f"rag-api is currently calling: {OLLAMA_URL}/api/generate. " "Open that host /api/tags in a browser; it must return JSON. " "If the URL is correct but you still see 404, set OLLAMA_MODEL on rag-api to a model that exists on that Ollama Space." ) else: hint = ( f"The Ollama Space returned HTTP {status_code}. It may be overloaded or restarting; " "check the Space logs on Hugging Face." ) return f"{hint}\n\nRetrieved excerpts from your documents:\n\n{excerpt}" def _fallback_answer_ollama_http(context: str, status_code: int) -> str: return _fallback_answer_llm_http(context, status_code) # Language name lookup used in the translation prompt. _LANG_NAMES = { "ar": "Arabic", "de": "German", "es": "Spanish", "fr": "French", "hi": "Hindi", "id": "Indonesian", "it": "Italian", "ja": "Japanese", "ko": "Korean", "nl": "Dutch", "pl": "Polish", "pt": "Portuguese", "ru": "Russian", "tr": "Turkish", "vi": "Vietnamese", "zh": "Chinese", } async def llm_translate(text: str, target_lang: str) -> str: """Translate an English answer via the configured LLM provider.""" lang_name = _LANG_NAMES.get(target_lang, target_lang) prompt = ( f"Translate the following text into {lang_name}. " "Output ONLY the translated text — no explanations, no preamble, no notes. " "Preserve Markdown formatting (bold, bullet lists, tables, links). " "Keep proper nouns, brand/product names, URLs, and code snippets unchanged.\n\n" f"{text}" ) try: translated = await llm_generate(prompt) return translated.strip() if translated and translated.strip() else text except Exception as e: logger.warning("llm_translate failed for lang=%s: %s — returning English", target_lang, e) return text def _narrow_results_for_context(results: pd.DataFrame) -> pd.DataFrame: """Keep strongest matches: sort by vector distance, drop loose tail, cap chunk count.""" if results.empty: return results out = results if "_distance" in out.columns: out = out.sort_values("_distance", ascending=True) best = float(out["_distance"].iloc[0]) if RAG_RELATIVE_DISTANCE_CAP > 1.0 and best > 1e-9: limit_d = best * RAG_RELATIVE_DISTANCE_CAP filtered = out[out["_distance"] <= limit_d] if not filtered.empty: out = filtered return out.head(RAG_CONTEXT_CHUNKS) async def refine_retrieval_results( results: pd.DataFrame, question: str, chat_history: list, ) -> tuple[pd.DataFrame, dict]: """ Stage-2 retrieval: cross-encoder rerank when enabled, else distance narrowing. Runs CPU reranking off the async event loop. """ if results.empty: return results, retrieval_meta_from_results( results, path="empty", skipped_rerank=True, retrieve_k=RAG_RETRIEVE_K ) original = results.copy() if RERANK_ENABLED and get_reranker() is not None and not should_skip_rerank(results): rerank_query = build_rerank_query(question, chat_history) loop = asyncio.get_event_loop() def _run_rerank(): return rerank_results(rerank_query, original) try: reranked = await loop.run_in_executor(_rerank_executor, _run_rerank) if reranked.empty: logger.warning( "Rerank filtered all %s candidates (min_score=%s); " "falling back to vector distance order", len(original), RERANK_MIN_SCORE, ) results = _narrow_results_for_context(original) meta = retrieval_meta_from_results( results, path="rerank_fallback", skipped_rerank=True, retrieve_k=RAG_RETRIEVE_K, ) meta["rerank_empty_fallback"] = True return results, meta results = reranked.head(RAG_CONTEXT_CHUNKS) meta = retrieval_meta_from_results( results, path="rerank", skipped_rerank=False, retrieve_k=RAG_RETRIEVE_K ) return results, meta except Exception as e: logger.error("Rerank failed (%s); falling back to vector distance order", e) results = _narrow_results_for_context(original) meta = retrieval_meta_from_results( results, path="rerank_error_fallback", skipped_rerank=True, retrieve_k=RAG_RETRIEVE_K, ) meta["rerank_error"] = str(e)[:200] return results, meta path = "distance_skip" if RERANK_ENABLED else "distance" results = _narrow_results_for_context(results) meta = retrieval_meta_from_results( results, path=path, skipped_rerank=True, retrieve_k=RAG_RETRIEVE_K ) return results, meta # --------------------------------------------------------------------------- # POST /upload # --------------------------------------------------------------------------- @app.post("/upload", dependencies=[Depends(_require_secret)]) async def upload_document( file: UploadFile = File(...), tenant_id: str = Form("default"), project_id: str = Form("default"), source_id: str = Form(""), ): ext = os.path.splitext(file.filename or "")[1].lower() if ext not in ALLOWED_EXTENSIONS: raise HTTPException( status_code=400, detail=f"Unsupported format '{ext}'. Accepted: {', '.join(ALLOWED_EXTENSIONS)}", ) tmp = tempfile.NamedTemporaryFile(delete=False, suffix=ext) try: shutil.copyfileobj(file.file, tmp) tmp.close() text = EXTRACTORS[ext](tmp.name) if not text.strip(): raise HTTPException( status_code=400, detail="No extractable text found in file." ) chunks = chunk_text(text) sid = (source_id or "").strip() migrate_documents_schema_if_needed() if sid and TABLE_NAME in get_table_names(): delete_vectors_exact(tenant_id, project_id, sid) BATCH_SIZE = 16 all_records: list[dict] = [] vectors = [] for i in range(0, len(chunks), BATCH_SIZE): batch = chunks[i : i + BATCH_SIZE] vectors = embed_model.encode(batch).tolist() for chunk_text_item, vec in zip(batch, vectors): all_records.append({ "text": chunk_text_item, "vector": vec, "tenant_id": tenant_id, "project_id": project_id, "source_id": sid, }) df = pd.DataFrame(all_records) stored_count = len(all_records) if TABLE_NAME in get_table_names(): table = db.open_table(TABLE_NAME) table.add(df) else: db.create_table(TABLE_NAME, df) del df, all_records, vectors, chunks, text gc.collect() # Persist to bucket in background (non-blocking) asyncio.create_task(sync_to_bucket_async()) return { "status": "ok", "chunks_stored": stored_count, "filename": file.filename, "tenant_id": tenant_id, "project_id": project_id, "source_id": sid, } except HTTPException: raise except Exception as e: logger.error(f"Upload failed: {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"Upload failed: {e}") finally: os.unlink(tmp.name) # --------------------------------------------------------------------------- # POST /query # --------------------------------------------------------------------------- class ChatTurn(BaseModel): user: str = "" assistant: str = "" class QueryRequest(BaseModel): query: str tenant_id: str = "default" project_id: str = "default" # If true, return top matching chunks + distances only (no Ollama) — for indexing QA retrieve_only: bool = False answer_language: str | None = None # Conversational RAG — prior turns and Redis session facts from the backend chat_history: list[ChatTurn] = [] session_facts: list[str] = [] @app.post("/query", dependencies=[Depends(_require_secret)]) async def query_documents(req: QueryRequest): if not req.retrieve_only: input_check = check_input(req.query) if input_check["blocked"]: return blocked_response(input_check["reason"]) if TABLE_NAME not in get_table_names(): if req.retrieve_only: raise HTTPException(status_code=404, detail="No documents uploaded yet.") retrieval_meta = { "path": "no_table", "skipped_rerank": True, "retrieve_k": RAG_RETRIEVE_K, "context_chunks": 0, "mode": "general", "relevance_gate_passed": False, "gate_fail_reasons": ["no_table"], } return await _general_llm_fallback_answer( req.query, req.chat_history, req.session_facts, req.answer_language, retrieval_meta, ) table = db.open_table(TABLE_NAME) retrieval_query = build_retrieval_query( req.query, req.chat_history, req.session_facts ) query_vec = embed_model.encode(retrieval_query or req.query).tolist() schema_cols = documents_table_schema_field_names() has_tenant_scope = "tenant_id" in schema_cols and "project_id" in schema_cols t_lit = _sql_string_literal(str(req.tenant_id)) p_lit = _sql_string_literal(str(req.project_id)) where_clause = f"tenant_id = {t_lit} AND project_id = {p_lit}" search = table.search(query_vec) if has_tenant_scope: results = search.where(where_clause).limit(RAG_RETRIEVE_K).to_pandas() else: logger.warning( "LanceDB %r missing tenant_id/project_id — unscoped search (legacy data)", TABLE_NAME, ) results = search.limit(RAG_RETRIEVE_K).to_pandas() if results.empty: if req.retrieve_only: return {"chunks": [], "count": 0, "tenant_id": req.tenant_id, "project_id": req.project_id} retrieval_meta = { "path": "no_results", "skipped_rerank": True, "retrieve_k": RAG_RETRIEVE_K, "context_chunks": 0, "mode": "general", "relevance_gate_passed": False, "gate_fail_reasons": ["no_results"], } return await _general_llm_fallback_answer( req.query, req.chat_history, req.session_facts, req.answer_language, retrieval_meta, ) results, retrieval_meta = await refine_retrieval_results( results, req.query, req.chat_history ) mode_meta = decide_answer_mode(req.query, results) retrieval_meta = {**retrieval_meta, **mode_meta} if results.empty: if req.retrieve_only: return { "chunks": [], "count": 0, "tenant_id": req.tenant_id, "project_id": req.project_id, "retrieval": retrieval_meta, } return await _general_llm_fallback_answer( req.query, req.chat_history, req.session_facts, req.answer_language, retrieval_meta, ) if not req.retrieve_only and is_conversational_ack(req.query): return await _general_llm_fallback_answer( req.query, req.chat_history, req.session_facts, req.answer_language, retrieval_meta, ) if req.retrieve_only: chunks_out = [] has_distance = "_distance" in results.columns has_rerank = "_rerank_score" in results.columns for _, row in results.iterrows(): item = { "text": row.get("text", ""), "tenant_id": row.get("tenant_id"), "project_id": row.get("project_id"), "source_id": row.get("source_id"), } if has_distance: item["distance"] = float(row["_distance"]) if has_rerank: item["rerank_score"] = float(row["_rerank_score"]) chunks_out.append(item) return { "chunks": chunks_out, "count": len(chunks_out), "tenant_id": req.tenant_id, "project_id": req.project_id, "retrieval": retrieval_meta, } raw_context = "\n\n".join(results["text"].tolist()) context = _truncate_context(raw_context) conversation_block = format_conversation_block(req.chat_history, req.session_facts) prompt = build_rag_answer_prompt(context, req.query, conversation_block) source_texts = results["text"].tolist() try: answer = await llm_generate(prompt, user_question=req.query) except httpx.ReadTimeout: logger.error("LLM read timeout after all retries — returning context fallback") return _package_query_response( _fallback_answer_from_context(raw_context), source_texts, {**retrieval_meta, "mode": "grounded"}, llm_fallback=True, ) except httpx.HTTPStatusError as e: code = e.response.status_code if code in (404, 429, 502, 503, 504): logger.error( "LLM gateway HTTP %s — returning excerpt fallback (check %s config)", code, LLM_PROVIDER, ) return _package_query_response( _fallback_answer_llm_http(raw_context, code), source_texts, {**retrieval_meta, "mode": "grounded"}, llm_fallback=True, ) logger.error(f"LLM error: {traceback.format_exc()}") raise HTTPException(status_code=502, detail=f"LLM error: HTTP {code}") except Exception as e: logger.error(f"LLM error: {traceback.format_exc()}") if isinstance(e, (httpx.ConnectError, httpx.ConnectTimeout)): return _package_query_response( _fallback_answer_from_context(raw_context), source_texts, {**retrieval_meta, "mode": "grounded"}, llm_fallback=True, ) raise HTTPException(status_code=502, detail=f"LLM error: {e}") # Step 2: translate if a non-English language was requested final_answer = answer.strip() if not final_answer or is_unusable_assistant_output(final_answer): logger.info( "LLM output unusable despite %s retrieved chunks — resolving refusal", len(source_texts), ) if raw_context.strip(): final_answer, refusal_meta = await _resolve_grounded_refusal( req.query, final_answer or "", raw_context, len(source_texts), retrieval_meta, ) retrieval_meta = {**retrieval_meta, **refusal_meta} else: try: return await _general_llm_fallback_answer( req.query, req.chat_history, req.session_facts, req.answer_language, {**retrieval_meta, "general_llm_after_unusable_output": True}, ) except Exception as e: logger.warning("General fallback after unusable LLM output failed (%s)", e) final_answer = ( "I'm here to help. Could you tell me a bit more about what you're looking for?" ) elif raw_context.strip() and is_llm_refusal(final_answer): final_answer, refusal_meta = await _resolve_grounded_refusal( req.query, final_answer, raw_context, len(source_texts), retrieval_meta, ) retrieval_meta = {**retrieval_meta, **refusal_meta} target_lang = (req.answer_language or "").strip().lower()[:12] translated = False if target_lang and target_lang != "en" and final_answer: try: final_answer = await llm_translate(final_answer, target_lang) translated = True except Exception as e: logger.warning("Post-translation failed for lang=%s (%s) — returning English", target_lang, e) return _package_query_response( final_answer, source_texts, {**retrieval_meta, "mode": "grounded"}, translated=translated, ) # --------------------------------------------------------------------------- # POST /delete-source — remove all chunks for one logical source (DocChat doc / URL / CSV) # --------------------------------------------------------------------------- class DeleteSourceRequest(BaseModel): tenant_id: str project_id: str source_id: str @app.post("/delete-source", dependencies=[Depends(_require_secret)]) async def delete_source_vectors(req: DeleteSourceRequest): migrate_documents_schema_if_needed() delete_vectors_exact(req.tenant_id, str(req.project_id), req.source_id) asyncio.create_task(sync_to_bucket_async()) return {"status": "ok"} # --------------------------------------------------------------------------- # POST /delete-sources-by-prefix — e.g. manual Q&A keys manual_qa_{ds.id}_* # --------------------------------------------------------------------------- class DeleteSourcePrefixRequest(BaseModel): tenant_id: str project_id: str source_id_prefix: str @app.post("/delete-sources-by-prefix", dependencies=[Depends(_require_secret)]) async def delete_source_vectors_by_prefix(req: DeleteSourcePrefixRequest): migrate_documents_schema_if_needed() delete_vectors_prefix(req.tenant_id, str(req.project_id), req.source_id_prefix) asyncio.create_task(sync_to_bucket_async()) return {"status": "ok"} # --------------------------------------------------------------------------- # POST /delete-project — all vectors for a tenant + project # --------------------------------------------------------------------------- class DeleteProjectVectorsRequest(BaseModel): tenant_id: str project_id: str @app.post("/delete-project", dependencies=[Depends(_require_secret)]) async def delete_project_vectors(req: DeleteProjectVectorsRequest): migrate_documents_schema_if_needed() delete_vectors_project(req.tenant_id, str(req.project_id)) asyncio.create_task(sync_to_bucket_async()) return {"status": "ok"} # --------------------------------------------------------------------------- # GET /debug-storage # --------------------------------------------------------------------------- @app.get("/debug-storage", dependencies=[Depends(_require_secret)]) async def debug_storage(): """Diagnostic: check local vs bucket storage status.""" info = {"local_db": LOCAL_DB_PATH, "bucket_db": BUCKET_DB_PATH} # Local /tmp try: local_files = [] for root, _, files in os.walk(LOCAL_DB_PATH): for f in files: fp = os.path.join(root, f) local_files.append({"path": os.path.relpath(fp, LOCAL_DB_PATH), "size": os.path.getsize(fp)}) info["local_file_count"] = len(local_files) info["local_total_bytes"] = sum(f["size"] for f in local_files) except Exception as e: info["local_error"] = str(e) # Bucket mount try: info["bucket_mount_exists"] = os.path.exists(BUCKET_PATH) info["bucket_mount_writable"] = os.access(BUCKET_PATH, os.W_OK) info["bucket_mount_contents"] = os.listdir(BUCKET_PATH) if os.path.exists(BUCKET_PATH) else [] except OSError as e: info["bucket_mount_error"] = str(e) # Bucket DB dir try: if os.path.exists(BUCKET_DB_PATH): bucket_files = [] for root, _, files in os.walk(BUCKET_DB_PATH): for f in files: fp = os.path.join(root, f) bucket_files.append({"path": os.path.relpath(fp, BUCKET_DB_PATH), "size": os.path.getsize(fp)}) info["bucket_file_count"] = len(bucket_files) info["bucket_total_bytes"] = sum(f["size"] for f in bucket_files) else: info["bucket_db_exists"] = False except OSError as e: info["bucket_db_error"] = str(e) # Test sync try: sync_to_bucket() info["sync_result"] = "success" except Exception as e: info["sync_result"] = f"failed: {e}" return info # --------------------------------------------------------------------------- # POST /reset # --------------------------------------------------------------------------- @app.post("/reset", dependencies=[Depends(_require_secret)]) async def reset_db(): for name in get_table_names(): db.drop_table(name) await sync_to_bucket_async() return {"status": "ok", "message": "All tables cleared"}