import os import asyncio from contextlib import asynccontextmanager from fastapi import FastAPI, HTTPException from fastapi.responses import StreamingResponse from pydantic import BaseModel from typing import List, Optional import logging from engine import init_engine, get_engine logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # Configuration MODEL_PATH = os.getenv("MODEL_PATH", "model/model.gguf") MODEL_URL = os.getenv("MODEL_URL", "https://huggingface.co/prithivMLmods/Nanbeige4.1-3B-f32-GGUF/resolve/main/Nanbeige4.1-3B.Q8_0.gguf") class GenerateRequest(BaseModel): prompt: str max_tokens: int = 256 temperature: float = 0.7 stream: bool = True class BatchRequest(BaseModel): prompts: List[str] max_tokens: int = 256 temperature: float = 0.7 def download_model(): """Download model if not exists""" if not os.path.exists(MODEL_PATH): os.makedirs(os.path.dirname(MODEL_PATH), exist_ok=True) logger.info(f"Downloading model from {MODEL_URL}") import urllib.request urllib.request.urlretrieve(MODEL_URL, MODEL_PATH) logger.info("Model downloaded") @asynccontextmanager async def lifespan(app: FastAPI): # Startup logger.info("Starting up...") download_model() init_engine(MODEL_PATH, n_ctx=4096, n_threads=4) logger.info("Ready!") yield # Shutdown logger.info("Shutting down...") app = FastAPI(title="Nanbeige LLM API", lifespan=lifespan) @app.post("/generate") async def generate(req: GenerateRequest): """Single prompt generation with streaming""" engine = get_engine() if req.stream: async def stream_generator(): async for token in engine.generate_stream( req.prompt, max_tokens=req.max_tokens, temperature=req.temperature ): yield token return StreamingResponse( stream_generator(), media_type="text/plain" ) else: # Non-streaming: collect all tokens chunks = [] async for token in engine.generate_stream( req.prompt, max_tokens=req.max_tokens, temperature=req.temperature ): chunks.append(token) return {"text": "".join(chunks)} @app.post("/generate_batch") async def generate_batch(req: BatchRequest): """Batch generation (multiple prompts)""" engine = get_engine() results = await engine.generate_batch( req.prompts, max_tokens=req.max_tokens, temperature=req.temperature ) return {"results": results} @app.get("/health") async def health(): return {"status": "ok", "model_loaded": get_engine()._model is not None}