Spaces:
Runtime error
Runtime error
Download api.py from Nikpatil/chatbot_theme_identifier: direct link, hf CLI and curl.
- Browser
- Download file 15.5 kB
-
https://huggingface.co/spaces/Nikpatil/chatbot_theme_identifier/resolve/main/api.py
- Command line
-
hf download hf://spaces/Nikpatil/chatbot_theme_identifier/api.py
-
curl -L -o api.py https://huggingface.co/spaces/Nikpatil/chatbot_theme_identifier/resolve/main/api.py
15.5 kB
| import os | |
| from fastapi import FastAPI, File, UploadFile, HTTPException, Form, BackgroundTasks, Query, APIRouter | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from fastapi.responses import JSONResponse | |
| from pydantic import BaseModel | |
| from typing import List, Optional, Dict, Any | |
| import tempfile | |
| import uuid | |
| import json | |
| from datetime import datetime | |
| import shutil | |
| from pathlib import Path | |
| import asyncio | |
| # Import Services | |
| from core.extractor import process_document | |
| from services.ingestion import VectorStoreManager | |
| from services.retrieval import VectorStoreRetriever | |
| from services.theme import ThemeSynthesizer | |
| # Intialize the FastAPI app | |
| app = FastAPI( | |
| title="Document Research & Theme Identification API", | |
| description="API for processing documents, querying content, and identifying themes", | |
| version="1.0.0" | |
| ) | |
| # Add CORS middleware | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], # Allows all origins | |
| allow_credentials=True, | |
| allow_methods=["*"], | |
| allow_headers=["*"] | |
| ) | |
| # Global variables | |
| DATA_DIR = Path("./data") | |
| CHROMA_DIR = DATA_DIR / "chroma_db" | |
| UPLOAD_DIR = DATA_DIR / "uploads" | |
| # Create necessary directories | |
| DATA_DIR.mkdir(exist_ok=True) | |
| CHROMA_DIR.mkdir(exist_ok=True) | |
| UPLOAD_DIR.mkdir(exist_ok=True) | |
| # Intialize services | |
| vector_store_manager = VectorStoreManager(persist_directory=str(CHROMA_DIR)) | |
| vector_store_retriever = VectorStoreRetriever(persist_directory=str(CHROMA_DIR)) | |
| # Initialize ThemeSynthesizer with Groq API key | |
| groq_api_key = os.getenv("GROQ_API_KEY") | |
| if groq_api_key: | |
| theme_synthesizer = ThemeSynthesizer(api_key=groq_api_key) | |
| else: | |
| theme_synthesizer = None | |
| print("Warning: GROQ_API_KEY not found. Theme synthesis will not be available.") | |
| # Process queue for background tasks | |
| processing_queue = {} | |
| # Pydantic models for request/response | |
| class DocumentResponse(BaseModel): | |
| id: str | |
| filename: str | |
| summary: Optional[str] = None | |
| added_st: str | |
| status: str | |
| metadata: Optional[Dict[str, Any]] = None | |
| class QueryRequest(BaseModel): | |
| query: str | |
| provider: Optional[str] = "groq" | |
| model: Optional[str] = "llama3-8b-8192" | |
| search_depth: int = 5 | |
| include_themes: bool = True | |
| theme_threshold: int = 2 | |
| class ThemeAnalysis(BaseModel): | |
| themes: List[str] | |
| status: str | |
| agreements: Optional[List[str]] = None | |
| contradictions: Optional[List[str]] = None | |
| insights: Optional[List[str]] = None | |
| analysis: Optional[str] = None | |
| class QueryResponse(BaseModel): | |
| results: List[Dict[str, Any]] | |
| theme_analysis: Optional[ThemeAnalysis] = None | |
| query: str | |
| total_results: int | |
| class LLMProvider(BaseModel): | |
| id: str | |
| display_name: str | |
| models: List[Dict[str, str]] | |
| class LLMProvidersResponse(BaseModel): | |
| providers: Dict[str, LLMProvider] | |
| # Router for document operations | |
| document_router = APIRouter(prefix="/api/documents", tags=["Documents"]) | |
| async def upload_document( | |
| background_tasks: BackgroundTasks, | |
| file: UploadFile = File(...) | |
| ): | |
| """Upload a document for processing""" | |
| doc_id = str(uuid.uuid4()) | |
| # Save file to disk | |
| file_path = UPLOAD_DIR / f"{doc_id}_{file.filename}" | |
| try: | |
| # Create a temporary file | |
| with open(file_path, "wb") as buffer: | |
| shutil.copyfileobj(file.file, buffer) | |
| # Add to processing queue | |
| processing_queue[doc_id] = { | |
| "id": doc_id, | |
| "filename": file.filename, | |
| "added_at": datetime.now().isoformat(), | |
| "status": "processing" | |
| } | |
| # Start background processing | |
| background_tasks.add_task( | |
| process_document_task, | |
| str(file_path), | |
| file.filename, | |
| doc_id | |
| ) | |
| return { | |
| "id": doc_id, | |
| "filename": file.filename, | |
| "status": "processing", | |
| "message": "Document upload successful, processing started" | |
| } | |
| except Exception as e: | |
| if file_path.exists(): | |
| file_path.unlink() # Clean up the file if it was created | |
| raise HTTPException(status_code=500, detail=f"Error uploading document: {str(e)}") | |
| async def get_documents(): | |
| """Get all processed documents""" | |
| try: | |
| # Get documents from vector store | |
| vector_docs = vector_store_retriever.search_by_metadata({}, limit=100) | |
| # Also check processing queue for documents still processing | |
| all_docs = [] | |
| # Add documents from the vector store | |
| for doc in vector_docs: | |
| doc_id = doc.get("metadata", {}).get("doc_id") | |
| if not doc_id: | |
| continue | |
| # Try to load document info from disk | |
| doc_info_path = DATA_DIR / f"{doc_id}.json" | |
| if doc_info_path.exists(): | |
| with open(doc_info_path, "r") as f: | |
| doc_info = json.load(f) | |
| all_docs.append(DocumentResponse(**doc_info)) | |
| else: | |
| # Create from vector store data | |
| all_docs.append(DocumentResponse( | |
| id=doc_id, | |
| filename=doc.get("metadata", {}).get("filename", "Unknown"), | |
| added_at=doc.get("metadata", {}).get("timestamp", datetime.now().isoformat()), | |
| status="completed", | |
| metadata=doc.get("metadata") | |
| )) | |
| # Add documents still in processing queue | |
| for doc_id, doc_info in processing_queue.items(): | |
| # Skip if already added from vector store | |
| if any(d.id == doc_id for d in all_docs): | |
| continue | |
| all_docs.append(DocumentResponse( | |
| id=doc_id, | |
| filename=doc_info["filename"], | |
| added_at=doc_info["added_at"], | |
| status=doc_info["status"], | |
| summary=doc_info.get("summary") | |
| )) | |
| return all_docs | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Error retrieving documents: {str(e)}") | |
| async def get_document(doc_id: str): | |
| """Get details for a specific document""" | |
| try: | |
| # Check if document info exists on disk | |
| doc_info_path = DATA_DIR / f"{doc_id}.json" | |
| if doc_info_path.exists(): | |
| with open(doc_info_path, "r") as f: | |
| doc_info = json.load(f) | |
| return DocumentResponse(**doc_info) | |
| # Check processing queue | |
| if doc_id in processing_queue: | |
| return DocumentResponse( | |
| id=doc_id, | |
| filename=processing_queue[doc_id]["filename"], | |
| added_at=processing_queue[doc_id]["added_at"], | |
| status=processing_queue[doc_id]["status"], | |
| summary=processing_queue[doc_id].get("summary") | |
| ) | |
| raise HTTPException(status_code=404, detail=f"Document {doc_id} not found") | |
| except HTTPException: | |
| raise | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Error retrieving document: {str(e)}") | |
| # Router for document extraction | |
| extraction_router = APIRouter(prefix="/api", tags=["Extraction"]) | |
| async def upload_and_extract(file: UploadFile = File(...)): | |
| """Extract text from a document without storing it""" | |
| try: | |
| temp_dir = tempfile.mkdtemp() | |
| try: | |
| file_ext = os.path.splitext(file.filename)[1] | |
| temp_file_path = os.path.join(temp_dir, f"{uuid.uuid4()}{file_ext}") | |
| with open(temp_file_path, "wb") as buffer: | |
| buffer.write(await file.read()) | |
| # Extract text | |
| result = process_document(file, temp_file_path) | |
| return JSONResponse(content=result) | |
| finally: | |
| # Always clean up temp directory | |
| shutil.rmtree(temp_dir, ignore_errors=True) | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Error extracting text: {str(e)}") | |
| # Router for querying | |
| query_router = APIRouter(prefix="/api", tags=["Query"]) | |
| async def query_documents(request: QueryRequest): | |
| """Query documents and identify themes""" | |
| try: | |
| # Search vector store | |
| search_results = vector_store_retriever.search( | |
| query=request.query, | |
| k=request.search_depth | |
| ) | |
| if not search_results: | |
| return QueryResponse( | |
| results=[], | |
| theme_analysis=None, | |
| query=request.query, | |
| total_results=0 | |
| ) | |
| # Process each search result with the LLM | |
| document_responses = [] | |
| for result in search_results: | |
| doc_id = result.get("metadata", {}).get("doc_id") | |
| doc_name = result.get("metadata", {}).get("filename", "Unknown") | |
| text_chunk = result.get("text", "") | |
| if not doc_id or not text_chunk: | |
| continue | |
| # Get response from LLM | |
| if theme_synthesizer: | |
| response = theme_synthesizer.process_query( | |
| query=request.query, | |
| context=text_chunk, | |
| provider=request.provider, | |
| model=request.model | |
| ) | |
| else: | |
| response = "LLM processing not available. Theme synthesizer is not initialized." | |
| document_responses.append({ | |
| "doc_id": doc_id, | |
| "doc_name": doc_name, | |
| "response": response, | |
| "chunk_id": result.get("id", "unknown_chunk") | |
| }) | |
| # Identify themes if requested and if we have enough documents | |
| theme_analysis = None | |
| if request.include_themes and len(document_responses) >= request.theme_threshold and theme_synthesizer: | |
| theme_analysis_result = theme_synthesizer.synthesize_themes( | |
| query=request.query, | |
| document_responses=document_responses, | |
| provider=request.provider, | |
| model=request.model | |
| ) | |
| theme_analysis = ThemeAnalysis( | |
| themes=theme_analysis_result.get("themes", []), | |
| status="completed", | |
| agreements=theme_analysis_result.get("agreements", []), | |
| contradictions=theme_analysis_result.get("contradictions", []), | |
| insights=theme_analysis_result.get("insights", []), | |
| analysis=theme_analysis_result.get("analysis", "") | |
| ) | |
| elif request.include_themes: | |
| theme_analysis = ThemeAnalysis( | |
| themes=[], | |
| status="skipped", | |
| analysis="Skipped theme analysis: not enough documents or theme synthesizer not available" | |
| ) | |
| return QueryResponse( | |
| results=document_responses, | |
| theme_analysis=theme_analysis, | |
| query=request.query, | |
| total_results=len(document_responses) | |
| ) | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Error processing query: {str(e)}") | |
| # Provider API | |
| provider_router = APIRouter(prefix="/api", tags = ["LLM"]) | |
| async def get_available_providers(): | |
| """Get available LLM providers and models""" | |
| # Get API keys from environment | |
| groq_api_key = os.getenv("GROQ_API_KEY", "") | |
| providers = {} | |
| if groq_api_key: | |
| providers["groq"] = LLMProvider( | |
| id="groq", | |
| display_name="Groq", | |
| models=[ | |
| {"id": "llama3-70b-8192", "name": "Llama 3 70B", "description": "Largest Llama 3 model"}, | |
| {"id": "llama3-8b-8192", "name": "Llama 3 8B", "description": "Smaller, faster Llama 3 model"}, | |
| {"id": "mixtral-8x7b-32768", "name": "Mixtral 8x7B", "description": "Mixtral, competitor to Llama"} | |
| ] | |
| ) | |
| return {"providers": providers} | |
| # Health check | |
| async def root(): | |
| return {"status": "ok", "message": "Document Research & Theme Identification API is running"} | |
| # Background task for document processing | |
| async def process_document_task(file_path: str, filename: str, doc_id: str): | |
| try: | |
| # Extract text from the document | |
| # Extract text from the document | |
| with open(file_path, "rb") as f: | |
| file_content = f.read() | |
| # Create a temporary UploadFile-like object | |
| mock_file = type('MockFile', (), {'filename': filename}) | |
| result = process_document(mock_file, file_path) | |
| if not result.get("success", False): | |
| raise Exception(f"Failed to extract text: {result.get('error', 'Unknown error')}") | |
| text = result.get("text", "") | |
| # Create metadata | |
| metadata = { | |
| "filename": filename, | |
| "doc_id": doc_id, | |
| "timestamp": datetime.now().isoformat() | |
| } | |
| # Get a summary if there's enough content | |
| summary = "No meaningful content to summarize." | |
| if len(text) > 100 and theme_synthesizer: | |
| try: | |
| summary = theme_synthesizer.analyze_single_document(text) | |
| except Exception as e: | |
| print(f"Error generating summary: {e}") | |
| summary = f"Error generating summary: {e}" | |
| # Store document info | |
| document_info = { | |
| "id": doc_id, | |
| "filename": filename, | |
| "summary": summary, | |
| "added_at": datetime.now().isoformat(), | |
| "status": "completed", | |
| "metadata": metadata | |
| } | |
| # Save document info to disk | |
| doc_info_path = DATA_DIR / f"{doc_id}.json" | |
| with open(doc_info_path, "w") as f: | |
| json.dump(document_info, f) | |
| # Save the full text separately (could be large) | |
| doc_text_path = DATA_DIR / f"{doc_id}.txt" | |
| with open(doc_text_path, "w") as f: | |
| f.write(text) | |
| # Add to vector store | |
| vector_store_manager.add_document(text, metadata, doc_id) | |
| # Update processing status | |
| processing_queue[doc_id]["status"] = "completed" | |
| processing_queue[doc_id]["summary"] = summary | |
| except Exception as e: | |
| # Update processing status with error | |
| processing_queue[doc_id]["status"] = "failed" | |
| processing_queue[doc_id]["error"] = str(e) | |
| print(f"Error processing document {doc_id}: {e}") | |
| finally: | |
| # Clean up the uploaded file | |
| try: | |
| Path(file_path).unlink(missing_ok=True) | |
| except Exception as e: | |
| print(f"Error removing temporary file {file_path}: {e}") | |
| # Include routers | |
| app.include_router(document_router) | |
| app.include_router(extraction_router) | |
| app.include_router(query_router) | |
| app.include_router(provider_router) | |
| if __name__ == "__main__": | |
| import uvicorn | |
| uvicorn.run("api:app", host="0.0.0.0", port=8000, reload=True) | |