Skip to content

Commit bfc627c

Browse files
committed
feat: add BYOK support and dynamic LLM provider selection via headers
1 parent 9ff61e6 commit bfc627c

6 files changed

Lines changed: 148 additions & 19 deletions

File tree

src/app/api/v1/chat.py

Lines changed: 25 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
import logging
22

3-
from fastapi import APIRouter, Depends, HTTPException, Request
3+
from fastapi import APIRouter, Depends, Header, HTTPException, Request
44
from langgraph.graph.state import CompiledStateGraph
55

66
from src.app.core.dependencies import get_agent_workflow
@@ -18,6 +18,9 @@ async def chat_endpoint(
1818
request: Request,
1919
payload: ChatRequest,
2020
workflow: CompiledStateGraph = Depends(get_agent_workflow),
21+
authorization: str | None = Header(default=None),
22+
x_api_key: str | None = Header(default=None),
23+
x_provider: str | None = Header(default=None),
2124
):
2225
"""
2326
Submit a question to the Vortex agentic workflow.
@@ -26,6 +29,25 @@ async def chat_endpoint(
2629
- The RAG pipeline (retrieve → grade → generate) for technical questions.
2730
- A direct LLM response for general queries.
2831
"""
32+
provider = x_provider.lower() if x_provider else None
33+
if provider and provider not in ["gemini", "anthropic", "ollama"]:
34+
raise HTTPException(
35+
status_code=400,
36+
detail=(
37+
f"Unsupported LLM provider: {x_provider}. "
38+
"Supported providers: gemini, anthropic, ollama"
39+
),
40+
)
41+
42+
api_key = None
43+
if authorization:
44+
if authorization.startswith("Bearer "):
45+
api_key = authorization[7:]
46+
else:
47+
api_key = authorization
48+
elif x_api_key:
49+
api_key = x_api_key
50+
2951
try:
3052
initial_state = {
3153
"question": payload.query,
@@ -34,6 +56,8 @@ async def chat_endpoint(
3456
"steps": ["received_query"],
3557
"route": "",
3658
"retry_count": 0,
59+
"api_key": api_key,
60+
"provider": provider,
3761
}
3862

3963
result = await workflow.ainvoke(initial_state)

src/app/core/llm.py

Lines changed: 16 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -5,17 +5,20 @@
55
from .config import settings
66

77

8-
@functools.lru_cache(maxsize=1)
9-
def get_llm() -> BaseChatModel:
8+
@functools.lru_cache(maxsize=32)
9+
def get_llm(api_key: str | None = None, provider: str | None = None) -> BaseChatModel:
1010
"""
11-
Returns the appropriate LLM instance based on configuration.
11+
Returns the appropriate LLM instance based on configuration or dynamic inputs.
1212
Supports Gemini (free), Ollama (local), and Anthropic (paid).
1313
1414
The instance is cached so that all graph nodes share a single LLM client
15-
per process, avoiding redundant re-initialization on every node call.
15+
per process/key, avoiding redundant re-initialization on every node call.
1616
"""
17-
if settings.llm_provider == "gemini":
18-
if not settings.gemini_api_key:
17+
active_provider = provider or settings.llm_provider
18+
19+
if active_provider == "gemini":
20+
active_key = api_key or settings.gemini_api_key
21+
if not active_key:
1922
raise ValueError(
2023
"GEMINI_API_KEY must be set to use Gemini. "
2124
"Get a free key at https://aistudio.google.com/"
@@ -24,11 +27,11 @@ def get_llm() -> BaseChatModel:
2427

2528
return ChatGoogleGenerativeAI(
2629
model=settings.gemini_model,
27-
google_api_key=settings.gemini_api_key,
30+
google_api_key=active_key,
2831
temperature=0.0,
2932
)
3033

31-
elif settings.llm_provider == "ollama":
34+
elif active_provider == "ollama":
3235
try:
3336
from langchain_ollama import ChatOllama
3437
except ImportError as exc:
@@ -42,8 +45,9 @@ def get_llm() -> BaseChatModel:
4245
temperature=0.0,
4346
)
4447

45-
elif settings.llm_provider == "anthropic":
46-
if not settings.anthropic_api_key:
48+
elif active_provider == "anthropic":
49+
active_key = api_key or settings.anthropic_api_key
50+
if not active_key:
4751
raise ValueError("ANTHROPIC_API_KEY must be set to use Anthropic.")
4852
try:
4953
from langchain_anthropic import ChatAnthropic
@@ -54,9 +58,9 @@ def get_llm() -> BaseChatModel:
5458
) from exc
5559
return ChatAnthropic(
5660
model=settings.anthropic_model,
57-
api_key=settings.anthropic_api_key,
61+
api_key=active_key,
5862
temperature=0.0,
5963
)
6064

6165
else:
62-
raise ValueError(f"Unsupported LLM provider: {settings.llm_provider}")
66+
raise ValueError(f"Unsupported LLM provider: {active_provider}")

src/app/graph/nodes.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -49,7 +49,7 @@ def grade_documents_node(state: AgentState) -> dict[str, Any]:
4949
steps = list(state.get("steps", []))
5050
steps.append("grade_documents")
5151

52-
llm = get_llm()
52+
llm = get_llm(api_key=state.get("api_key"), provider=state.get("provider"))
5353
structured_llm_grader = llm.with_structured_output(Grade)
5454

5555
system = (
@@ -91,7 +91,7 @@ def generate_node(state: AgentState) -> dict[str, Any]:
9191
steps = list(state.get("steps", []))
9292
steps.append("generate_answer")
9393

94-
llm = get_llm()
94+
llm = get_llm(api_key=state.get("api_key"), provider=state.get("provider"))
9595
prompt = ChatPromptTemplate.from_template(
9696
"You are an expert IT infrastructure support "
9797
"assistant for the Cortex and Sentinel platforms. "
@@ -126,7 +126,7 @@ def direct_response_node(state: AgentState) -> dict[str, Any]:
126126
steps = list(state.get("steps", []))
127127
steps.append("direct_response")
128128

129-
llm = get_llm()
129+
llm = get_llm(api_key=state.get("api_key"), provider=state.get("provider"))
130130
prompt = ChatPromptTemplate.from_template(
131131
"You are a helpful IT support assistant. Answer the following question "
132132
"directly and concisely.\n\n"
@@ -157,7 +157,7 @@ def rewrite_query_node(state: AgentState) -> dict[str, Any]:
157157
steps = list(state.get("steps", []))
158158
steps.append("rewrite_query")
159159

160-
llm = get_llm()
160+
llm = get_llm(api_key=state.get("api_key"), provider=state.get("provider"))
161161
system = (
162162
"You are a question re-writer that converts an input question to a better "
163163
"version optimized for vector store retrieval. Reason about the underlying "

src/app/graph/router.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -38,7 +38,7 @@ def router_node(state: AgentState) -> dict[str, Any]:
3838
steps = list(state.get("steps", []))
3939
steps.append("router")
4040

41-
llm = get_llm()
41+
llm = get_llm(api_key=state.get("api_key"), provider=state.get("provider"))
4242
structured_llm = llm.with_structured_output(RouteDecision)
4343

4444
system = (

src/app/graph/state.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
from typing import TypedDict
1+
from typing import NotRequired, TypedDict
22

33
from langchain_core.documents import Document
44

@@ -14,6 +14,8 @@ class AgentState(TypedDict):
1414
steps: Ordered list of node names executed (for observability).
1515
route: Router decision — "retrieve" or "direct" (for API response).
1616
retry_count: Number of query-rewrite retries performed (loop protection).
17+
api_key: Optional dynamic API key for the LLM provider.
18+
provider: Optional dynamic LLM provider (gemini, anthropic, ollama).
1719
"""
1820

1921
question: str
@@ -22,3 +24,5 @@ class AgentState(TypedDict):
2224
steps: list[str]
2325
route: str
2426
retry_count: int
27+
api_key: NotRequired[str | None]
28+
provider: NotRequired[str | None]

tests/test_auth.py

Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,97 @@
1+
from unittest.mock import AsyncMock
2+
3+
from fastapi.testclient import TestClient
4+
5+
from src.app.core.dependencies import get_agent_workflow
6+
from src.app.main import app
7+
8+
client = TestClient(app)
9+
10+
11+
def test_chat_byok_headers():
12+
"""
13+
Test that Authorization and X-Provider headers are correctly
14+
extracted and passed to the workflow state.
15+
"""
16+
mock_workflow = AsyncMock()
17+
mock_workflow.ainvoke.return_value = {
18+
"question": "test question",
19+
"generation": "mock response",
20+
"documents": [],
21+
"steps": ["received_query"],
22+
"route": "direct",
23+
"retry_count": 0,
24+
}
25+
26+
app.dependency_overrides[get_agent_workflow] = lambda: mock_workflow
27+
28+
# 1. Test Bearer token in Authorization header and X-Provider header
29+
headers = {
30+
"Authorization": "Bearer my-custom-gemini-key",
31+
"X-Provider": "gemini",
32+
}
33+
response = client.post(
34+
"/api/v1/chat",
35+
json={"query": "test question"},
36+
headers=headers,
37+
)
38+
39+
assert response.status_code == 200
40+
assert mock_workflow.ainvoke.call_count == 1
41+
call_state = mock_workflow.ainvoke.call_args[0][0]
42+
assert call_state["api_key"] == "my-custom-gemini-key"
43+
assert call_state["provider"] == "gemini"
44+
45+
# 2. Test X-API-Key header and different casing for X-Provider
46+
mock_workflow.reset_mock()
47+
headers = {
48+
"X-API-Key": "my-custom-anthropic-key",
49+
"X-Provider": "Anthropic",
50+
}
51+
response = client.post(
52+
"/api/v1/chat",
53+
json={"query": "test question"},
54+
headers=headers,
55+
)
56+
57+
assert response.status_code == 200
58+
assert mock_workflow.ainvoke.call_count == 1
59+
call_state = mock_workflow.ainvoke.call_args[0][0]
60+
assert call_state["api_key"] == "my-custom-anthropic-key"
61+
assert call_state["provider"] == "anthropic"
62+
63+
# 3. Test when no headers are provided (should pass None)
64+
mock_workflow.reset_mock()
65+
response = client.post(
66+
"/api/v1/chat",
67+
json={"query": "test question"},
68+
)
69+
70+
assert response.status_code == 200
71+
assert mock_workflow.ainvoke.call_count == 1
72+
call_state = mock_workflow.ainvoke.call_args[0][0]
73+
assert call_state["api_key"] is None
74+
assert call_state["provider"] is None
75+
76+
app.dependency_overrides.clear()
77+
78+
79+
def test_chat_invalid_provider():
80+
"""Test that an unsupported provider returns a 400 Bad Request error."""
81+
mock_workflow = AsyncMock()
82+
app.dependency_overrides[get_agent_workflow] = lambda: mock_workflow
83+
84+
headers = {
85+
"X-Provider": "unsupported-llm-brand",
86+
}
87+
response = client.post(
88+
"/api/v1/chat",
89+
json={"query": "test question"},
90+
headers=headers,
91+
)
92+
93+
assert response.status_code == 400
94+
assert "Unsupported LLM provider" in response.json()["detail"]
95+
assert mock_workflow.ainvoke.call_count == 0
96+
97+
app.dependency_overrides.clear()

0 commit comments

Comments
 (0)