syntaxhacker
use preciz SDK from pip package instead of local path
af4937c
Raw
History Blame Contribute Delete
5.31 kB
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
import os
import logging
import sys
import json
from dotenv import load_dotenv
from .config import DATASET_CONFIGS, load_prompt_template
load_dotenv()
from summarizer.llm_client import _call_llm, _parse_tool_calls, _strip_tool_tags, TOOL_HANDLERS, execute_tool
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
handlers=[logging.StreamHandler(sys.stdout)]
)
logger = logging.getLogger(__name__)
app = FastAPI(title="RAG Pipeline API", description="Multi-dataset RAG API", version="1.0.0")
MODEL_NAME = os.getenv("MODEL_NAME", "openrouter/owl-alpha")
MAX_ROUNDS = 6
pipelines = {}
logger.info(f"Starting RAG Pipeline API — model: {MODEL_NAME}")
logger.info(f"Available datasets: {list(DATASET_CONFIGS.keys())}")
def rag_qa(question: str, dataset: str = "developer-portfolio") -> str:
try:
if not pipelines:
return "RAG Pipeline is running but datasets are still loading. Please try again in a moment."
if dataset not in pipelines:
return f"Dataset '{dataset}' not available. Available datasets: {list(pipelines.keys())}"
return pipelines[dataset].answer_question(question)
except Exception as e:
return f"Error accessing RAG pipeline: {str(e)}"
def handle_rag_qa_tool(tool_input: str, user_id: str | None = None) -> str:
try:
args = json.loads(tool_input)
return rag_qa(args.get("question", ""), args.get("dataset", "developer-portfolio"))
except json.JSONDecodeError:
parts = tool_input.split(":", 1)
if len(parts) == 2:
return rag_qa(parts[1].strip(), parts[0].strip())
return rag_qa(tool_input.strip())
TOOL_HANDLERS["rag_qa"] = handle_rag_qa_tool
class Question(BaseModel):
text: str
dataset: str = "developer-portfolio"
class ChatMessage(BaseModel):
role: str
content: str
class ChatRequest(BaseModel):
messages: list[ChatMessage]
dataset: str = "developer-portfolio"
@app.post("/chat")
async def chat_with_ai(request: ChatRequest):
messages = [{"role": m.role, "content": m.content} for m in request.messages]
if request.dataset == "developer-portfolio":
system = {"role": "system", "content": load_prompt_template("system-instruction.txt")}
else:
system = {"role": "system", "content": load_prompt_template("generic-system-instruction.txt")}
messages.insert(0, system)
for _ in range(MAX_ROUNDS):
content = _call_llm(messages, model=MODEL_NAME, max_tokens=4000)
tool_calls = _parse_tool_calls(content)
if not tool_calls:
clean = _strip_tool_tags(content)
return {"response": clean if clean else content, "tool_calls": None}
clean_content = _strip_tool_tags(content)
messages.append({"role": "assistant", "content": clean_content or "Let me check that..."})
results = []
for name, inp in tool_calls:
result = execute_tool(name, inp)
results.append(result)
for result in results:
messages.append({"role": "user", "content": f"RAG result:\n{result}\n\nNow answer based on this."})
content = _call_llm(messages, model=MODEL_NAME, max_tokens=4000)
clean = _strip_tool_tags(content)
return {"response": clean if clean else content, "tool_calls": None}
@app.get("/datasets")
async def list_datasets():
return {"datasets": list(pipelines.keys())}
@app.get("/questions")
async def list_questions(dataset: str = "developer-portfolio"):
if dataset not in pipelines:
raise HTTPException(status_code=400, detail=f"Dataset '{dataset}' not available. Available datasets: {list(pipelines.keys())}")
selected_pipeline = pipelines[dataset]
questions = [doc.meta['question'] for doc in selected_pipeline.documents if 'question' in doc.meta]
return {"dataset": dataset, "questions": questions}
async def load_datasets_background():
global pipelines
from .pipeline import RAGPipeline
dataset_name = "developer-portfolio"
try:
logger.info(f"Loading dataset: {dataset_name}")
pipeline = RAGPipeline.from_preset(preset_name=dataset_name)
pipelines[dataset_name] = pipeline
logger.info(f"Successfully loaded {dataset_name}")
except Exception as e:
logger.error(f"Failed to load {dataset_name}: {e}")
logger.info(f"Background loading complete — {len(pipelines)} datasets loaded")
@app.on_event("startup")
async def startup_event():
logger.info("FastAPI application startup complete")
import asyncio
asyncio.create_task(load_datasets_background())
@app.on_event("shutdown")
async def shutdown_event():
logger.info("FastAPI application shutting down")
@app.get("/")
async def root():
return {"status": "ok", "message": "RAG Pipeline API", "version": "1.0.0", "datasets": list(pipelines.keys())}
@app.get("/health")
async def health_check():
loading_status = "complete" if "developer-portfolio" in pipelines else "loading"
return {
"status": "healthy",
"datasets_loaded": len(pipelines),
"total_datasets": 1,
"loading_status": loading_status,
"port": os.getenv("PORT", "8000"),
}