Umithan's picture
Create app.py
984e1e7 verified
Raw
History Blame Contribute Delete
2.06 kB
import os
import json
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import StreamingResponse
from fastapi.middleware.cors import CORSMiddleware
from huggingface_hub import hf_hub_download
from llama_cpp import Llama
app = FastAPI()
# Enable CORS for frontend integration
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Download the Gemma model from Hugging Face Hub if not cached locally
MODEL_FILE = "gemma-4-E2B-it-IQ4_NL.gguf"
if not os.path.exists(MODEL_FILE):
print("Downloading Gemma model, please wait...")
hf_hub_download(
repo_id="unsloth/gemma-4-E2B-it-GGUF",
filename=MODEL_FILE,
local_dir="."
)
# Initialize local LLM instance with 2048 context length
llm = Llama(model_path=f"./{MODEL_FILE}", n_ctx=2048, n_threads=2)
@app.post("/v1/chat/completions")
async def chat_completion(request: Request):
# Secure the endpoint using a custom API Bearer Token
api_key = request.headers.get("Authorization")
if api_key != f"Bearer {os.environ.get('MY_SECRET_KEY', 'default_pass')}":
raise HTTPException(status_code=401, detail="Unauthorized access.")
body = await request.json()
messages = body.get("messages", [])
# Format chat history into standard LLM prompt template
prompt = ""
for msg in messages:
role = msg.get("role")
content = msg.get("content")
prompt += f"<|im_start|>{role}\n{content}<|im_end|>\n"
prompt += "<|im_start|>assistant\n"
# Generate streaming response from the model
output = llm(prompt, max_tokens=512, stream=True)
def stream_generator():
for chunk in output:
token = chunk['choices'][0]['text']
data = {"choices": [{"delta": {"content": token}, "finish_reason": None}]}
yield f"data: {json.dumps(data)}\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(stream_generator(), media_type="text/event-stream")