redfernstech/cpp_llama
0
1from fastapi import FastAPI, HTTPException, Depends, Header, Request
2from pydantic import BaseModel
3import os
4import logging
5import time
6from langchain_community.llms import LlamaCpp
7from dotenv import load_dotenv
8
9# Load environment variables
10load_dotenv()
11
12# Configure logging
13logging.basicConfig(level=logging.INFO)
14
15# API keys from .env
16API_KEYS = {
17 "user1": os.getenv("API_KEY_USER1"),
18 "user2": os.getenv("API_KEY_USER2"),
19}
20
21app = FastAPI()
22
23# API Key Authentication
24def verify_api_key(request: Request, api_key: str = Header(None, alias="X-API-Key")):
25 logging.info(f"Received Headers: {request.headers}")
26 if not api_key:
27 raise HTTPException(status_code=401, detail="API key is missing")
28
29 api_key = api_key.strip()
30 if api_key not in API_KEYS.values():
31 raise HTTPException(status_code=401, detail="Invalid API key")
32
33 return api_key
34
35# OpenAI-compatible request format
36class OpenAIRequest(BaseModel):
37 model: str
38 messages: list
39 stream: bool = False
40
41# Initialize LangChain with Llama.cpp
42def get_llm():
43 model_path = "/app/Meta-Llama-3-8B-Instruct.Q4_0.gguf"
44 return LlamaCpp(model_path=model_path, n_ctx=2048)
45
46@app.post("/v1/chat/completions")
47def generate_text(request: OpenAIRequest, api_key: str = Depends(verify_api_key)):
48 try:
49 llm = get_llm()
50
51 # Extract last user message
52 user_message = next((msg["content"] for msg in reversed(request.messages) if msg["role"] == "user"), None)
53 if not user_message:
54 raise HTTPException(status_code=400, detail="User message is required")
55
56 response_text = llm.invoke(user_message)
57
58 response = {
59 "id": "chatcmpl-123",
60 "object": "chat.completion",
61 "created": int(time.time()),
62 "model": request.model,
63 "choices": [
64 {
65 "index": 0,
66 "message": {"role": "assistant", "content": response_text},
67 "finish_reason": "stop",
68 }
69 ],
70 "usage": {
71 "prompt_tokens": len(user_message.split()),
72 "completion_tokens": len(response_text.split()),
73 "total_tokens": len(user_message.split()) + len(response_text.split()),
74 }
75 }
76
77 return response
78
79 except Exception as e:
80 logging.error(f"Error generating response: {e}")
81 raise HTTPException(status_code=500, detail="Internal server error")
82 