307 lines
10 KiB
Python
307 lines
10 KiB
Python
import os
|
|
import sqlite3
|
|
import numpy as np
|
|
import requests
|
|
from fastapi import FastAPI, HTTPException, Query
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
from pydantic import BaseModel
|
|
from typing import List, Optional
|
|
|
|
# --- CONFIGURATION ---
|
|
OLLAMA_URL = os.environ.get("OLLAMA_URL", "http://localhost:11434/api/embed")
|
|
OLLAMA_GENERATE_URL = os.environ.get("OLLAMA_GENERATE_URL", "http://localhost:11434/api/generate")
|
|
EMBEDDING_MODEL = os.environ.get("EMBEDDING_MODEL", "qwen3-embedding:0.6b")
|
|
DB_PATH = os.environ.get("DB_PATH", "/root/work/knowledge-base/knowledge_base.db")
|
|
|
|
app = FastAPI(
|
|
title="Juniper Technical Document RAG Engine",
|
|
description="REST API for querying the technical documentation index and generating RAG answers.",
|
|
version="1.0.0"
|
|
)
|
|
|
|
# Enable CORS for frontend integration
|
|
app.add_middleware(
|
|
CORSMiddleware,
|
|
allow_origins=["*"],
|
|
allow_credentials=True,
|
|
allow_methods=["*"],
|
|
allow_headers=["*"],
|
|
)
|
|
|
|
# --- MODELS ---
|
|
class SearchRequest(BaseModel):
|
|
query: str
|
|
top_k: int = 5
|
|
source: Optional[str] = None # "kb" or "book"
|
|
|
|
class SearchResultItem(BaseModel):
|
|
id: int
|
|
source_type: str
|
|
source_name: str
|
|
title: str
|
|
url: Optional[str] = None
|
|
page_num: Optional[int] = None
|
|
text: str
|
|
score: float
|
|
|
|
class GenerateRequest(BaseModel):
|
|
query: str
|
|
top_k: int = 5
|
|
source: Optional[str] = None
|
|
llm_model: str = "qwen3-embedding:0.6b" # Default fallback, can be overridden by user to e.g. "qwen2.5"
|
|
|
|
class GenerateResponse(BaseModel):
|
|
query: str
|
|
answer: str
|
|
sources: List[SearchResultItem]
|
|
|
|
# --- UTILITIES ---
|
|
def get_query_embedding(query_text: str) -> List[float]:
|
|
"""Fetches query embedding from Ollama."""
|
|
try:
|
|
response = requests.post(OLLAMA_URL, json={
|
|
"model": EMBEDDING_MODEL,
|
|
"input": query_text
|
|
}, timeout=30)
|
|
response.raise_for_status()
|
|
data = response.json()
|
|
return data.get("embeddings", [])[0]
|
|
except Exception as e:
|
|
raise HTTPException(
|
|
status_code=503,
|
|
detail=f"Error getting query embedding from Ollama: {str(e)}. Please ensure Ollama is running and has the embedding model '{EMBEDDING_MODEL}' pulled."
|
|
)
|
|
|
|
def load_table_data(cursor, table_name: str, source_type: str):
|
|
"""Loads chunks and embeddings from database."""
|
|
if table_name == "knowledge_base":
|
|
cursor.execute("SELECT id, kb_id, salesforce_id, title, url, prefix_content, embedding FROM knowledge_base")
|
|
rows = cursor.fetchall()
|
|
data = []
|
|
embeddings = []
|
|
for row in rows:
|
|
doc_id, kb_id, salesforce_id, title, url, prefix_content, emb_bytes = row
|
|
emb = np.frombuffer(emb_bytes, dtype=np.float32)
|
|
if len(emb) == 1024:
|
|
embeddings.append(emb)
|
|
data.append({
|
|
"id": doc_id,
|
|
"source_type": source_type,
|
|
"source_name": salesforce_id,
|
|
"title": f"{kb_id} : {title}" if not title.startswith(kb_id) else title,
|
|
"url": url,
|
|
"page_num": None,
|
|
"text": prefix_content
|
|
})
|
|
return data, embeddings
|
|
else:
|
|
cursor.execute(f"SELECT id, file_name, page_num, prefix_content, embedding FROM {table_name}")
|
|
rows = cursor.fetchall()
|
|
data = []
|
|
embeddings = []
|
|
for row in rows:
|
|
doc_id, file_name, page_num, prefix_content, emb_bytes = row
|
|
emb = np.frombuffer(emb_bytes, dtype=np.float32)
|
|
if len(emb) == 1024:
|
|
embeddings.append(emb)
|
|
data.append({
|
|
"id": doc_id,
|
|
"source_type": source_type,
|
|
"source_name": file_name,
|
|
"title": file_name,
|
|
"url": None,
|
|
"page_num": page_num,
|
|
"text": prefix_content
|
|
})
|
|
return data, embeddings
|
|
|
|
def get_full_content(item: dict) -> str:
|
|
"""Retrieves full content of the document page or KB article."""
|
|
if item['source_type'] == 'kb':
|
|
kb_id = item['source_name']
|
|
file_path = f"/root/work/knowledge-base/juniper-kb/kb_markdown/{kb_id}.md"
|
|
if os.path.exists(file_path):
|
|
try:
|
|
with open(file_path, 'r', encoding='utf-8') as f:
|
|
return f.read()
|
|
except Exception:
|
|
pass
|
|
|
|
# Fallback to combining chunks from SQLite
|
|
conn = sqlite3.connect(DB_PATH)
|
|
c = conn.cursor()
|
|
c.execute("SELECT text_content FROM knowledge_base WHERE kb_id = ? ORDER BY chunk_index", (kb_id,))
|
|
chunks = [row[0] for row in c.fetchall()]
|
|
conn.close()
|
|
return "\n\n".join(chunks)
|
|
else:
|
|
file_name = item['source_name']
|
|
page_num = item['page_num']
|
|
table_name = "network_book"
|
|
|
|
conn = sqlite3.connect(DB_PATH)
|
|
c = conn.cursor()
|
|
c.execute(f"SELECT text_content FROM {table_name} WHERE file_name = ? AND page_num = ? ORDER BY chunk_index", (file_name, page_num))
|
|
chunks = [row[0] for row in c.fetchall()]
|
|
conn.close()
|
|
return "\n\n".join(chunks)
|
|
|
|
def run_retrieval(query: str, top_k: int, source: Optional[str]) -> List[SearchResultItem]:
|
|
"""Helper function to perform vector similarity search."""
|
|
if not os.path.exists(DB_PATH):
|
|
raise HTTPException(status_code=500, detail=f"Database file not found at {DB_PATH}. Please run indexing first.")
|
|
|
|
query_emb = np.array(get_query_embedding(query), dtype=np.float32)
|
|
|
|
conn = sqlite3.connect(DB_PATH)
|
|
cursor = conn.cursor()
|
|
|
|
all_data = []
|
|
all_embeddings = []
|
|
|
|
tables_to_load = [
|
|
("knowledge_base", "kb"),
|
|
("network_book", "book")
|
|
]
|
|
|
|
if source:
|
|
tables_to_load = [t for t in tables_to_load if t[1] == source]
|
|
|
|
for table_name, source_type in tables_to_load:
|
|
try:
|
|
data, embs = load_table_data(cursor, table_name, source_type)
|
|
all_data.extend(data)
|
|
all_embeddings.extend(embs)
|
|
except sqlite3.OperationalError:
|
|
continue
|
|
|
|
conn.close()
|
|
|
|
if not all_embeddings:
|
|
return []
|
|
|
|
embs_matrix = np.array(all_embeddings, dtype=np.float32)
|
|
q_norm = np.linalg.norm(query_emb)
|
|
m_norms = np.linalg.norm(embs_matrix, axis=1)
|
|
|
|
m_norms[m_norms == 0] = 1e-10
|
|
if q_norm == 0:
|
|
q_norm = 1e-10
|
|
|
|
similarities = np.dot(embs_matrix, query_emb) / (m_norms * q_norm)
|
|
top_indices = np.argsort(similarities)[::-1][:top_k]
|
|
|
|
results = []
|
|
for idx in top_indices:
|
|
item = all_data[idx]
|
|
score = float(similarities[idx])
|
|
results.append(SearchResultItem(
|
|
id=item["id"],
|
|
source_type=item["source_type"],
|
|
source_name=item["source_name"],
|
|
title=item["title"],
|
|
url=item["url"],
|
|
page_num=item["page_num"],
|
|
text=item["text"],
|
|
score=score
|
|
))
|
|
return results
|
|
|
|
# --- ENDPOINTS ---
|
|
@app.get("/")
|
|
@app.get("/health")
|
|
def health_check():
|
|
"""Provides index statistics and service health status."""
|
|
db_exists = os.path.exists(DB_PATH)
|
|
kb_count = 0
|
|
book_count = 0
|
|
|
|
if db_exists:
|
|
try:
|
|
conn = sqlite3.connect(DB_PATH)
|
|
c = conn.cursor()
|
|
c.execute("SELECT COUNT(*) FROM knowledge_base")
|
|
kb_count = c.fetchone()[0]
|
|
c.execute("SELECT COUNT(*) FROM network_book")
|
|
book_count = c.fetchone()[0]
|
|
conn.close()
|
|
except sqlite3.OperationalError:
|
|
pass
|
|
|
|
return {
|
|
"status": "healthy",
|
|
"database_connected": db_exists,
|
|
"database_path": DB_PATH,
|
|
"stats": {
|
|
"juniper_kb_chunks": kb_count,
|
|
"network_book_chunks": book_count,
|
|
"total_chunks": kb_count + book_count
|
|
},
|
|
"model_config": {
|
|
"embedding_model": EMBEDDING_MODEL,
|
|
"ollama_url": OLLAMA_URL
|
|
}
|
|
}
|
|
|
|
@app.post("/search", response_model=List[SearchResultItem])
|
|
def search_documents(req: SearchRequest):
|
|
"""Searches indexed documents for matching chunks using similarity search."""
|
|
return run_retrieval(req.query, req.top_k, req.source)
|
|
|
|
@app.post("/generate", response_model=GenerateResponse)
|
|
def generate_rag_answer(req: GenerateRequest):
|
|
"""Retrieves context chunks and prompts Ollama LLM to synthesize a detailed answer."""
|
|
# 1. Retrieve most relevant context
|
|
sources = run_retrieval(req.query, req.top_k, req.source)
|
|
|
|
if not sources:
|
|
return GenerateResponse(
|
|
query=req.query,
|
|
answer="No relevant context documents found in the database. Please verify your query or index.",
|
|
sources=[]
|
|
)
|
|
|
|
# 2. Build context block
|
|
context_blocks = []
|
|
for idx, s in enumerate(sources, 1):
|
|
context_blocks.append(f"[{idx}] Source: {s.title} ({s.source_type})\nContent: {s.text}")
|
|
|
|
context_str = "\n\n".join(context_blocks)
|
|
|
|
# 3. Build detailed prompt for the LLM
|
|
prompt = (
|
|
"You are an expert network engineering AI assistant. Answer the user's question based strictly on the provided technical documents.\n"
|
|
"If you do not know the answer or if it's not present in the context, state that you don't know based on the documents.\n\n"
|
|
"=== TECHNICAL DOCUMENTS CONTEXT ===\n"
|
|
f"{context_str}\n"
|
|
"====================================\n\n"
|
|
f"Question: {req.query}\n\n"
|
|
"Answer (be structured, clear, and reference your sources like [1], [2] when appropriate):"
|
|
)
|
|
|
|
# 4. Request generation from Ollama
|
|
try:
|
|
response = requests.post(OLLAMA_GENERATE_URL, json={
|
|
"model": req.llm_model,
|
|
"prompt": prompt,
|
|
"stream": False
|
|
}, timeout=90)
|
|
response.raise_for_status()
|
|
answer = response.json().get("response", "")
|
|
except Exception as e:
|
|
answer = (
|
|
f"[RAG Context Retrieval Successful, but LLM Generation failed]\n"
|
|
f"Error communicating with Ollama generation model '{req.llm_model}': {str(e)}.\n"
|
|
f"Please verify Ollama has the model pulled or specify a running model name in your request."
|
|
)
|
|
|
|
return GenerateResponse(
|
|
query=req.query,
|
|
answer=answer,
|
|
sources=sources
|
|
)
|
|
|
|
if __name__ == "__main__":
|
|
import uvicorn
|
|
uvicorn.run("app:app", host="0.0.0.0", port=8000, reload=True)
|