Spaces:
Sleeping
Sleeping
| """ | |
| LLM Engine - Handles API calls to OpenAI and Hugging Face | |
| """ | |
| from typing import Dict, Optional, Tuple | |
| import json | |
| import time | |
| from config import ( | |
| OPENAI_API_KEY, | |
| LLM_PROVIDER, | |
| OPENAI_MODEL, | |
| HF_MODEL, | |
| TEMPERATURE, | |
| TOP_P, | |
| ) | |
| from core.utils import log_event | |
| class LLMEngine: | |
| """Base class for LLM interactions.""" | |
| def __init__(self, provider: str = LLM_PROVIDER): | |
| """ | |
| Initialize LLM Engine. | |
| Args: | |
| provider: "openai" or "huggingface" | |
| """ | |
| self.provider = provider.lower() | |
| self.model = OPENAI_MODEL if self.provider == "openai" else HF_MODEL | |
| if self.provider == "openai": | |
| self.engine = OpenAIEngine() | |
| elif self.provider == "huggingface": | |
| self.engine = HuggingFaceEngine() | |
| else: | |
| raise ValueError(f"Unknown provider: {provider}") | |
| def generate(self, prompt: str, max_tokens: int = 1000) -> Tuple[bool, str]: | |
| """ | |
| Delegate generation to the selected engine. | |
| """ | |
| return self.engine.generate(prompt, max_tokens) | |
| def stream_generate(self, prompt: str, max_tokens: int = 1000): | |
| """ | |
| Stream generation for real-time responses. | |
| """ | |
| yield from self.engine.stream_generate(prompt, max_tokens) | |
| class OpenAIEngine: | |
| """OpenAI API integration.""" | |
| def __init__(self): | |
| """Initialize OpenAI engine.""" | |
| try: | |
| import openai | |
| if not OPENAI_API_KEY: | |
| raise ValueError( | |
| "OPENAI_API_KEY is missing. Set it in your .env file." | |
| ) | |
| openai.api_key = OPENAI_API_KEY | |
| self.client = openai.OpenAI(api_key=OPENAI_API_KEY) | |
| self.model = OPENAI_MODEL | |
| except ImportError: | |
| raise ImportError("openai package not installed. Install with: pip install openai") | |
| def generate(self, prompt: str, max_tokens: int = 1000) -> Tuple[bool, str]: | |
| """ | |
| Generate text using OpenAI API. | |
| Args: | |
| prompt: Input prompt | |
| max_tokens: Max tokens in response | |
| Returns: | |
| Tuple of (success, response_text) | |
| """ | |
| try: | |
| log_event("API_CALL", f"OpenAI - Model: {self.model}, Tokens: {max_tokens}") | |
| response = self.client.chat.completions.create( | |
| model=self.model, | |
| messages=[ | |
| { | |
| "role": "system", | |
| "content": "You are a helpful AI study assistant. Provide clear, accurate, and educational responses." | |
| }, | |
| { | |
| "role": "user", | |
| "content": prompt | |
| } | |
| ], | |
| max_tokens=max_tokens, | |
| temperature=TEMPERATURE, | |
| top_p=TOP_P, | |
| ) | |
| result = response.choices[0].message.content | |
| log_event("API_SUCCESS", "OpenAI response received") | |
| return True, result | |
| except Exception as e: | |
| error_msg = str(e) | |
| log_event("API_ERROR", f"OpenAI: {error_msg}") | |
| if "insufficient_quota" in error_msg or "Error code: 429" in error_msg: | |
| return ( | |
| False, | |
| "OpenAI API Error: Insufficient quota (429). " | |
| "Add billing/credits in your OpenAI account, or switch provider by setting " | |
| "LLM_PROVIDER=huggingface in .env to use the local model." | |
| ) | |
| return False, f"OpenAI API Error: {error_msg}" | |
| def stream_generate(self, prompt: str, max_tokens: int = 1000): | |
| """ | |
| Stream responses from OpenAI API. | |
| Args: | |
| prompt: Input prompt | |
| max_tokens: Max tokens | |
| Yields: | |
| Response chunks | |
| """ | |
| try: | |
| response = self.client.chat.completions.create( | |
| model=self.model, | |
| messages=[ | |
| {"role": "system", "content": "You are a helpful AI study assistant."}, | |
| {"role": "user", "content": prompt} | |
| ], | |
| max_tokens=max_tokens, | |
| temperature=TEMPERATURE, | |
| stream=True, | |
| ) | |
| for chunk in response: | |
| if chunk.choices[0].delta.content: | |
| yield chunk.choices[0].delta.content | |
| except Exception as e: | |
| yield f"Error: {str(e)}" | |
| class HuggingFaceEngine: | |
| """Local Hugging Face model (no API).""" | |
| _shared_tokenizer = None | |
| _shared_model_obj = None | |
| _shared_model_name = None | |
| _shared_device = None | |
| def __init__(self): | |
| """Initialize local Hugging Face pipeline.""" | |
| try: | |
| import torch | |
| from transformers import AutoModelForSeq2SeqLM, AutoTokenizer | |
| self.model = "google/flan-t5-base" if HF_MODEL == "local" else HF_MODEL | |
| if ( | |
| HuggingFaceEngine._shared_model_obj is None | |
| or HuggingFaceEngine._shared_tokenizer is None | |
| or HuggingFaceEngine._shared_model_name != self.model | |
| ): | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| tokenizer = AutoTokenizer.from_pretrained(self.model) | |
| model_obj = AutoModelForSeq2SeqLM.from_pretrained(self.model).to(device) | |
| HuggingFaceEngine._shared_tokenizer = tokenizer | |
| HuggingFaceEngine._shared_model_obj = model_obj | |
| HuggingFaceEngine._shared_model_name = self.model | |
| HuggingFaceEngine._shared_device = device | |
| self.device = HuggingFaceEngine._shared_device | |
| self.tokenizer = HuggingFaceEngine._shared_tokenizer | |
| self.model_obj = HuggingFaceEngine._shared_model_obj | |
| except ImportError: | |
| raise ImportError( | |
| "transformers/torch/sentencepiece not installed. " | |
| "Install with: pip install transformers torch sentencepiece" | |
| ) | |
| def generate(self, prompt: str, max_tokens: int = 1000) -> Tuple[bool, str]: | |
| try: | |
| log_event("API_CALL", f"Local HF - {self.model}") | |
| inputs = self.tokenizer( | |
| prompt, | |
| return_tensors="pt", | |
| truncation=True, | |
| max_length=1024, | |
| ).to(self.device) | |
| generation_max_length = 200 | |
| generation_min_length = 80 | |
| import torch | |
| with torch.no_grad(): | |
| outputs = self.model_obj.generate( | |
| **inputs, | |
| max_length=generation_max_length, | |
| min_length=generation_min_length, | |
| do_sample=True, | |
| temperature=0.6, | |
| top_p=0.9, | |
| repetition_penalty=1.3, | |
| no_repeat_ngram_size=3, | |
| ) | |
| decoded_text = self.tokenizer.decode(outputs[0], skip_special_tokens=True) | |
| response = [{"generated_text": decoded_text}] | |
| result = response[0]["generated_text"] | |
| if not isinstance(result, str): | |
| result = str(result) | |
| result = result.strip() | |
| if not result: | |
| log_event("API_ERROR", "LocalHF: Empty response generated") | |
| return False, "Local Model Error: Empty response generated" | |
| log_event("API_SUCCESS", "Local model response received") | |
| return True, result | |
| except Exception as e: | |
| error_msg = str(e) | |
| log_event("API_ERROR", f"LocalHF: {error_msg}") | |
| print("FULL ERROR:", error_msg) | |
| return False, f"Local Model Error: {error_msg}" | |
| def stream_generate(self, prompt: str, max_tokens: int = 1000): | |
| """Yield single complete response for compatibility with streaming API.""" | |
| success, response = self.generate(prompt, max_tokens) | |
| yield response | |
| class SimpleCache: | |
| """Simple in-memory cache for responses.""" | |
| def __init__(self, max_size: int = 100): | |
| """Initialize cache.""" | |
| self.cache = {} | |
| self.max_size = max_size | |
| def get(self, key: str) -> Optional[str]: | |
| """Get cached response.""" | |
| return self.cache.get(key) | |
| def set(self, key: str, value: str) -> None: | |
| """Cache a response.""" | |
| if len(self.cache) >= self.max_size: | |
| # Remove oldest entry (FIFO) | |
| self.cache.pop(next(iter(self.cache))) | |
| self.cache[key] = value | |
| def clear(self) -> None: | |
| """Clear cache.""" | |
| self.cache.clear() | |
| # Global cache instance | |
| _cache = SimpleCache() | |
| def get_or_generate(prompt: str, engine: LLMEngine, max_tokens: int = 1000) -> Tuple[bool, str]: | |
| """ | |
| Get cached response or generate new one. | |
| Args: | |
| prompt: Input prompt | |
| engine: LLM engine to use | |
| max_tokens: Max tokens | |
| Returns: | |
| Tuple of (success, response) | |
| """ | |
| # Create cache key from prompt | |
| cache_key = hash(prompt) % ((2 ** 63) - 1) | |
| # Check cache | |
| cached = _cache.get(str(cache_key)) | |
| if cached: | |
| log_event("CACHE_HIT", "Using cached response") | |
| return True, cached | |
| # Generate if not cached | |
| success, response = engine.generate(prompt, max_tokens) | |
| if success: | |
| _cache.set(str(cache_key), response) | |
| return success, response | |