Skip to content

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())
Output
{'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"])
Output
[('refunds.md', 0.6)]
I couldn't find that in the knowledge base.

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())
Output
[fake] stream me

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"])
Output
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__)
Output
passed: test_rag_returns_sources
passed: test_validation
passed: test_provider_failure_is_500

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/504 with 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 SESSIONS is 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 429 when exceeded, and a test for it.

Next: Streamlit — a chat UI for this API in 40 lines.