5. FastAPI for GenAI¶
Intermediate · 14 min read
This page puts topics 1–4 together into the shape of a real LLM backend. Every external piece (the LLM, the vector store) sits behind a small interface, so the whole app runs and is tested here with fakes — and switching to OpenAI, Anthropic or Pinecone touches one function.
5.1 Project layout¶
app/
├── main.py # FastAPI() + routers + middleware + lifespan
├── config.py # Settings (pydantic-settings), get_settings()
├── schemas.py # Pydantic request / response models
├── deps.py # get_llm(), get_retriever(), auth
├── services/
│ ├── llm.py # LLM client wrapper (OpenAI / Anthropic / fake)
│ ├── retriever.py # vector search
│ └── agent.py # tool loop
└── routers/
├── chat.py
├── rag.py
└── agent.py
tests/
└── test_api.py # TestClient + dependency_overrides with fakes
Below everything is in one block so you can run it; in a real project, split it as above.
5.2 A swappable LLM client¶
from typing import Protocol, AsyncIterator
import asyncio
class LLM(Protocol):
async def complete(self, messages: list[dict]) -> str: ...
def stream(self, messages: list[dict]) -> AsyncIterator[str]: ...
class FakeLLM:
"""Deterministic stand-in: used in tests, demos and local development."""
async def complete(self, messages: list[dict]) -> str:
return f"[fake] {messages[-1]['content'][:60]}"
async def stream(self, messages: list[dict]):
for word in (await self.complete(messages)).split(" "):
await asyncio.sleep(0)
yield word + " "
The real one has the same two methods:
# no-run — services/llm.py
from openai import AsyncOpenAI
class OpenAILLM:
def __init__(self, model: str, api_key: str):
self.client, self.model = AsyncOpenAI(api_key=api_key, timeout=30, max_retries=2), model
async def complete(self, messages):
r = await self.client.chat.completions.create(model=self.model, messages=messages)
return r.choices[0].message.content
async def stream(self, messages):
s = await self.client.chat.completions.create(model=self.model, messages=messages, stream=True)
async for chunk in s:
if chunk.choices and chunk.choices[0].delta.content:
yield chunk.choices[0].delta.content
5.3 Schemas and dependencies¶
from typing import Annotated, Literal
from fastapi import Depends, FastAPI, HTTPException
from fastapi.testclient import TestClient
from pydantic import BaseModel, Field
class Message(BaseModel):
role: Literal["system", "user", "assistant"]
content: str = Field(min_length=1, max_length=8000)
class ChatRequest(BaseModel):
session_id: str
message: str = Field(min_length=1, max_length=4000)
class ChatResponse(BaseModel):
session_id: str
reply: str
turns: int
class Source(BaseModel):
doc_id: str
text: str
score: float
class RAGRequest(BaseModel):
question: str = Field(min_length=3)
top_k: int = Field(3, ge=1, le=10)
class RAGResponse(BaseModel):
answer: str
sources: list[Source]
_llm = FakeLLM()
def get_llm() -> LLM:
return _llm
LLMDep = Annotated[LLM, Depends(get_llm)]
app = FastAPI(title="GenAI service")
5.4 Chat with session memory¶
The server keeps the conversation per session_id (in memory here; Redis or Postgres in production) and
trims old turns so the prompt stays within the context window:
SYSTEM = {"role": "system", "content": "You are a concise assistant."}
SESSIONS: dict[str, list[dict]] = {}
MAX_HISTORY = 6 # last 6 messages = 3 turns
@app.post("/chat", response_model=ChatResponse)
async def chat(req: ChatRequest, llm: LLMDep):
history = SESSIONS.setdefault(req.session_id, [])
history.append({"role": "user", "content": req.message})
reply = await llm.complete([SYSTEM, *history[-MAX_HISTORY:]])
history.append({"role": "assistant", "content": reply})
return ChatResponse(session_id=req.session_id, reply=reply, turns=len(history) // 2)
client = TestClient(app)
for text in ["Hi, I'm Priya", "What is a vector database?"]:
print(client.post("/chat", json={"session_id": "s1", "message": text}).json())
{'session_id': 's1', 'reply': "[fake] Hi, I'm Priya", 'turns': 1}
{'session_id': 's1', 'reply': '[fake] What is a vector database?', 'turns': 2}
5.5 RAG endpoint — answer + sources¶
Return the sources with every answer, so the UI can show citations and you can debug bad answers:
KB = {
"refunds.md": "Refunds are processed within 5 working days of approval.",
"shipping.md": "Orders ship within 48 hours. Express delivery is available in metro cities.",
"account.md": "Reset your password from Settings > Security.",
}
class KeywordRetriever: # stand-in for Pinecone / pgvector / Qdrant
def search(self, query: str, top_k: int) -> list[Source]:
q = set(query.lower().split())
scored = [Source(doc_id=d, text=t, score=len(q & set(t.lower().split())) / len(q))
for d, t in KB.items()]
return sorted([s for s in scored if s.score > 0], key=lambda s: -s.score)[:top_k]
def get_retriever() -> KeywordRetriever:
return KeywordRetriever()
@app.post("/rag/query", response_model=RAGResponse)
async def rag_query(req: RAGRequest, llm: LLMDep,
retriever: Annotated[KeywordRetriever, Depends(get_retriever)]):
sources = retriever.search(req.question, req.top_k)
if not sources:
return RAGResponse(answer="I couldn't find that in the knowledge base.", sources=[])
context = "\n".join(f"[{s.doc_id}] {s.text}" for s in sources)
prompt = f"Answer only from the context.\n\nContext:\n{context}\n\nQuestion: {req.question}"
answer = await llm.complete([{"role": "user", "content": prompt}])
return RAGResponse(answer=answer, sources=sources)
r = client.post("/rag/query", json={"question": "how long are refunds processed"}).json()
print([(s["doc_id"], round(s["score"], 2)) for s in r["sources"]])
print(client.post("/rag/query", json={"question": "crypto staking rewards"}).json()["answer"])
Say 'I don't know' in code, not just in the prompt
If retrieval finds nothing relevant, don't call the LLM at all. It's cheaper, faster, and the most reliable guard against made-up answers.
5.6 Streaming endpoint¶
import json
from fastapi.responses import StreamingResponse
@app.post("/chat/stream")
async def chat_stream(req: ChatRequest, llm: LLMDep):
async def events():
try:
async for token in llm.stream([SYSTEM, {"role": "user", "content": req.message}]):
yield f"data: {json.dumps({'token': token})}\n\n"
except Exception:
yield f"data: {json.dumps({'error': 'generation failed'})}\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(events(), media_type="text/event-stream")
with client.stream("POST", "/chat/stream", json={"session_id": "s2", "message": "stream me"}) as r:
tokens = [json.loads(l[6:])["token"] for l in r.iter_lines()
if l.startswith("data: {")]
print("".join(tokens).strip())
5.7 An agent endpoint (tool loop)¶
An agent = the LLM picks a tool, your code runs it, the result goes back, repeat until a final answer. The fake "planner" below plays the LLM's role; with a real model it is one tool-calling API call per step.
TOOLS = {
"get_weather": lambda city: f"28°C and cloudy in {city}",
"search_kb": lambda query: KeywordRetriever().search(query, 1)[0].text,
}
class AgentRequest(BaseModel):
task: str
max_steps: int = Field(4, ge=1, le=8)
class Step(BaseModel):
tool: str
args: dict
result: str
class AgentResponse(BaseModel):
answer: str
steps: list[Step]
def fake_planner(task: str, steps: list[Step]) -> dict:
"""Returns {'tool': ..., 'args': ...} or {'final': ...} — the LLM's job in a real agent."""
if not steps and "weather" in task.lower():
return {"tool": "get_weather", "args": {"city": "Bengaluru"}}
if not steps:
return {"tool": "search_kb", "args": {"query": task}}
return {"final": f"Based on {steps[-1].tool}: {steps[-1].result}"}
@app.post("/agent/run", response_model=AgentResponse)
async def agent_run(req: AgentRequest):
steps: list[Step] = []
for _ in range(req.max_steps): # hard limit: agents can loop forever
decision = fake_planner(req.task, steps)
if "final" in decision:
return AgentResponse(answer=decision["final"], steps=steps)
if decision["tool"] not in TOOLS:
raise HTTPException(400, f"Unknown tool {decision['tool']}")
result = TOOLS[decision["tool"]](**decision["args"])
steps.append(Step(tool=decision["tool"], args=decision["args"], result=result))
raise HTTPException(status_code=422, detail="Agent did not finish within max_steps")
r = client.post("/agent/run", json={"task": "What's the weather like?"}).json()
print(r["answer"])
print([s["tool"] for s in r["steps"]])
print(client.post("/agent/run", json={"task": "reset password"}).json()["answer"])
Based on get_weather: 28°C and cloudy in Bengaluru
['get_weather']
Based on search_kb: Reset your password from Settings > Security.
Returning the steps makes the agent observable — the UI can show "Checking weather…", and you can
log every decision for debugging and evals.
5.8 Testing with fakes¶
Override the LLM dependency and assert on behaviour — no API key, no cost, no flaky network:
class BrokenLLM(FakeLLM):
async def complete(self, messages):
raise RuntimeError("provider down")
def test_rag_returns_sources():
r = client.post("/rag/query", json={"question": "express delivery cities"})
assert r.status_code == 200
assert r.json()["sources"][0]["doc_id"] == "shipping.md"
def test_validation():
assert client.post("/rag/query", json={"question": "hi", "top_k": 99}).status_code == 422
def test_provider_failure_is_500():
app.dependency_overrides[get_llm] = lambda: BrokenLLM()
try:
r = TestClient(app, raise_server_exceptions=False).post(
"/chat", json={"session_id": "x", "message": "hello"})
assert r.status_code == 500
finally:
app.dependency_overrides.clear()
for t in (test_rag_returns_sources, test_validation, test_provider_failure_is_500):
t()
print("passed:", t.__name__)
In a real project put these in tests/test_api.py and run pytest — see
Testing with pytest.
5.9 Deployment checklist¶
# production: several worker processes, no --reload
uvicorn app.main:app --host 0.0.0.0 --port 8000 --workers 4
# or
fastapi run app/main.py --workers 4
FROM python:3.12-slim
WORKDIR /app
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
COPY app ./app
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]
- Secrets from environment variables / secret manager — never in the image or the repo.
- Timeouts and retries on every LLM call; return
502/504with a friendly message. - Auth + rate limits per user — every request costs money.
- Limits on input size,
top_k,max_tokens, agent steps. - State outside the process (Redis/Postgres) — with 4 workers, in-memory
SESSIONSis 4 different dicts. - Logging of request ID, user, model, tokens, latency; tracing with LangSmith / Langfuse / OpenTelemetry.
- Health check endpoint (
/health) for the load balancer.
Interview questions¶
Why FastAPI for LLM services rather than Flask?
Native async (many concurrent, slow LLM calls per worker), Pydantic validation of request/response bodies, automatic OpenAPI docs, dependency injection that makes swapping LLM clients and testing with fakes easy, and first-class streaming support.
How do you stream LLM output to a browser?
Use the provider's streaming API, wrap it in an async generator, return a StreamingResponse with
media_type="text/event-stream", and send SSE data: events (tokens, then metadata, then [DONE]).
The frontend reads it with EventSource or fetch + a stream reader. Disable proxy buffering.
What breaks when you scale a FastAPI LLM app to multiple workers?
Anything held in process memory: chat history, caches, rate-limit counters, background tasks.
Move them to Redis/Postgres/a queue. Also watch for blocking sync calls inside async def, and
upstream rate limits from the LLM provider.
Practice¶
- Add a
DELETE /chat/{session_id}that clears a session. - Add a per-session limit of 20 messages that returns
429when exceeded, and a test for it.
Next: Streamlit — a chat UI for this API in 40 lines.