knowledge-base/search_kb.py

234 lines
8.6 KiB
Python

import os
import sys
import sqlite3
import requests
import numpy as np
from rich.console import Console
from rich.panel import Panel
from rich.text import Text
# --- CONFIGURATION ---
OLLAMA_URL = os.environ.get("OLLAMA_URL", "http://localhost:11434/api/embed")
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")
console = Console()
def get_query_embedding(query_text):
"""Fetches embedding for the search query 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:
console.print(f"[bold red]Error getting query embedding from Ollama:[/bold red] {e}")
console.print("[yellow]Make sure Ollama is running and the model qwen3-embedding:0.6b is pulled.[/yellow]")
sys.exit(1)
def load_table_data(cursor, table_name, source_type):
"""Loads text chunks and embeddings from a specific table."""
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):
"""Retrieves the full content of a KB article or PDF page."""
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 as e:
console.print(f"[dim yellow]Warning: could not read file {file_path}: {e}[/dim yellow]")
# 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:
# Reconstruct page content
file_name = item['source_name']
page_num = item['page_num']
table_name = "os_release_note" if item['source_type'] == 'release_note' else "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 search(query, top_k=5, target_source=None, show_full=False):
"""Performs cosine similarity search against the embedded documents."""
if not os.path.exists(DB_PATH):
console.print(f"[bold red]Database not found at {DB_PATH}.[/bold red] Please run embed_docs.py first.")
return
# 1. Embed query
console.print(f"[dim]Generating embedding for query...[/dim]")
query_emb = np.array(get_query_embedding(query), dtype=np.float32)
# 2. Connect to database and load data
conn = sqlite3.connect(DB_PATH)
cursor = conn.cursor()
all_data = []
all_embeddings = []
tables_to_load = [
("knowledge_base", "kb"),
("network_book", "book")
]
if target_source:
tables_to_load = [t for t in tables_to_load if t[1] == target_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:
# Table might not exist yet if indexing has not run for this source
continue
conn.close()
if not all_embeddings:
console.print("[bold yellow]No matching tables or data found in the database. Run embed_docs.py to index documents.[/bold yellow]")
return
# 3. Calculate cosine similarity
console.print(f"[dim]Searching {len(all_embeddings)} document chunks...[/dim]")
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)
# Avoid zero division
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)
# 4. Extract Top-K
top_indices = np.argsort(similarities)[::-1][:top_k]
# 5. Display output
console.print(Panel(
f"[bold green]Query:[/bold green] '{query}'\n[dim]Model: {EMBEDDING_MODEL} | Top {top_k} results shown below[/dim]",
border_style="green",
expand=False
))
for rank, idx in enumerate(top_indices, 1):
item = all_data[idx]
score = similarities[idx]
title_text = f"{rank}. {item['title']} (Score: [bold green]{score:.4f}[/bold green])"
meta_info = f"Type: [bold blue]{item['source_type'].upper()}[/bold blue]"
if item['page_num']:
meta_info += f" | Page: [bold yellow]{item['page_num']}[/bold yellow]"
if item['url']:
meta_info += f" | [link={item['url']}]URL[/link]"
if not show_full:
meta_info += " | [dim yellow]Add --full to view complete doc/page[/dim yellow]"
content_to_show = Text(item['text'])
else:
if item['source_type'] == 'kb':
from rich.markdown import Markdown
try:
content_to_show = Markdown(get_full_content(item))
except Exception:
content_to_show = Text(get_full_content(item))
else:
content_to_show = Text(get_full_content(item))
console.print(Panel(
content_to_show,
title=title_text,
subtitle=meta_info,
border_style="cyan",
padding=(1, 2)
))
if __name__ == "__main__":
if len(sys.argv) < 2:
console.print("[bold yellow]Usage:[/bold yellow] python3 search_kb.py \"<search query>\" [--source kb|book] [--top-k <number>] [--full]")
sys.exit(0)
query = sys.argv[1]
# Parse args
target_source = None
top_k = 5
show_full = "--full" in sys.argv
if "--source" in sys.argv:
try:
idx = sys.argv.index("--source")
target_source = sys.argv[idx + 1]
if target_source not in ["kb", "book"]:
raise ValueError
except Exception:
console.print("[bold red]Error:[/bold red] --source must be one of: kb, book")
sys.exit(1)
if "--top-k" in sys.argv:
try:
idx = sys.argv.index("--top-k")
top_k = int(sys.argv[idx + 1])
except Exception:
console.print("[bold red]Error:[/bold red] --top-k must be an integer")
sys.exit(1)
search(query, top_k=top_k, target_source=target_source, show_full=show_full)