chopratejas commited on
Commit
ad26396
Β·
1 Parent(s): 2da7d1d

Add cloud provider support via LiteLLM backend

Browse files

Enables Headroom proxy to work with AWS Bedrock, Google Vertex AI,
Azure OpenAI, and 100+ other providers via LiteLLM.

Usage:
headroom proxy --backend bedrock --region us-west-2
headroom proxy --backend vertex_ai --region us-central1
headroom proxy --backend azure --region eastus

Features:
- Automatic format translation (Anthropic API <-> provider APIs)
- Streaming support with proper SSE event translation
- Uses existing cloud credentials (AWS, GCP, Azure)
- Shorthand: --backend bedrock (expands to litellm-bedrock)
- Stats tracking per provider

README.md CHANGED
@@ -213,6 +213,20 @@ ANTHROPIC_BASE_URL=http://localhost:8787 claude
213
  OPENAI_BASE_URL=http://localhost:8787/v1 cursor
214
  ```
215
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
216
  ### Option 2: LangChain Integration
217
 
218
  ```bash
 
213
  OPENAI_BASE_URL=http://localhost:8787/v1 cursor
214
  ```
215
 
216
+ **Using AWS Bedrock, Google Vertex, or Azure?** Route through Headroom:
217
+
218
+ ```bash
219
+ # AWS Bedrock (uses your AWS credentials)
220
+ headroom proxy --backend bedrock --region us-west-2
221
+ ANTHROPIC_BASE_URL=http://localhost:8787 claude
222
+
223
+ # Google Vertex AI
224
+ headroom proxy --backend vertex_ai --region us-central1
225
+
226
+ # Azure OpenAI
227
+ headroom proxy --backend azure --region eastus
228
+ ```
229
+
230
  ### Option 2: LangChain Integration
231
 
232
  ```bash
headroom/backends/__init__.py ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Headroom Backends - API translation layers for different LLM providers.
2
+
3
+ Backends handle the translation between the proxy's canonical format
4
+ (Anthropic Messages API) and provider-specific APIs.
5
+
6
+ Uses LiteLLM for broad provider support:
7
+ - bedrock: AWS Bedrock (Claude, Cohere, Mistral, etc.)
8
+ - vertex_ai: Google Vertex AI (Claude, Gemini, etc.)
9
+ - azure: Azure OpenAI (GPT-4, etc.)
10
+ - And 100+ more providers...
11
+
12
+ Usage:
13
+ headroom proxy --backend litellm-bedrock --region us-west-2
14
+ """
15
+
16
+ from .base import Backend, BackendResponse, StreamEvent
17
+ from .litellm import LiteLLMBackend
18
+
19
+ __all__ = ["Backend", "BackendResponse", "StreamEvent", "LiteLLMBackend"]
headroom/backends/base.py ADDED
@@ -0,0 +1,122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Base backend interface for Headroom.
2
+
3
+ Backends translate between the canonical Anthropic Messages API format
4
+ (used by the proxy) and provider-specific APIs.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from abc import ABC, abstractmethod
10
+ from collections.abc import AsyncIterator
11
+ from dataclasses import dataclass, field
12
+ from typing import Any
13
+
14
+
15
+ @dataclass
16
+ class BackendResponse:
17
+ """Standardized response from a backend."""
18
+
19
+ # Response body (Anthropic Messages API format)
20
+ body: dict[str, Any]
21
+
22
+ # HTTP status code
23
+ status_code: int = 200
24
+
25
+ # Response headers to forward
26
+ headers: dict[str, str] = field(default_factory=dict)
27
+
28
+ # Error message if any
29
+ error: str | None = None
30
+
31
+
32
+ @dataclass
33
+ class StreamEvent:
34
+ """A single event from a streaming response."""
35
+
36
+ # Event type (message_start, content_block_delta, etc.)
37
+ event_type: str
38
+
39
+ # Event data (Anthropic SSE format)
40
+ data: dict[str, Any]
41
+
42
+ # Raw SSE line to forward (if available)
43
+ raw_sse: str | None = None
44
+
45
+
46
+ class Backend(ABC):
47
+ """Abstract base class for LLM API backends.
48
+
49
+ Backends are responsible for:
50
+ - Translating requests from Anthropic format to provider format
51
+ - Making API calls to the provider
52
+ - Translating responses back to Anthropic format
53
+ - Handling streaming
54
+ """
55
+
56
+ @property
57
+ @abstractmethod
58
+ def name(self) -> str:
59
+ """Backend name (e.g., 'anthropic', 'bedrock')."""
60
+ ...
61
+
62
+ @abstractmethod
63
+ async def send_message(
64
+ self,
65
+ body: dict[str, Any],
66
+ headers: dict[str, str],
67
+ ) -> BackendResponse:
68
+ """Send a non-streaming message request.
69
+
70
+ Args:
71
+ body: Request body in Anthropic Messages API format.
72
+ headers: Request headers (may include API keys, etc.).
73
+
74
+ Returns:
75
+ BackendResponse with body in Anthropic Messages API format.
76
+ """
77
+ ...
78
+
79
+ @abstractmethod
80
+ def stream_message(
81
+ self,
82
+ body: dict[str, Any],
83
+ headers: dict[str, str],
84
+ ) -> AsyncIterator[StreamEvent]:
85
+ """Stream a message request.
86
+
87
+ Args:
88
+ body: Request body in Anthropic Messages API format.
89
+ headers: Request headers.
90
+
91
+ Yields:
92
+ StreamEvent objects in Anthropic SSE format.
93
+ """
94
+ ...
95
+
96
+ @abstractmethod
97
+ def map_model_id(self, anthropic_model: str) -> str:
98
+ """Map Anthropic model ID to provider model ID.
99
+
100
+ Args:
101
+ anthropic_model: Model ID in Anthropic format (e.g., 'claude-3-opus-20240229').
102
+
103
+ Returns:
104
+ Model ID in provider format.
105
+ """
106
+ ...
107
+
108
+ @abstractmethod
109
+ def supports_model(self, model: str) -> bool:
110
+ """Check if this backend supports a model.
111
+
112
+ Args:
113
+ model: Model ID (can be Anthropic or provider format).
114
+
115
+ Returns:
116
+ True if model is supported.
117
+ """
118
+ ...
119
+
120
+ async def close(self) -> None: # noqa: B027
121
+ """Clean up resources (e.g., close HTTP clients)."""
122
+ pass
headroom/backends/litellm.py ADDED
@@ -0,0 +1,421 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LiteLLM-based backend for Headroom.
2
+
3
+ Uses LiteLLM to support 100+ providers with minimal code:
4
+ - AWS Bedrock: model="bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0"
5
+ - Azure OpenAI: model="azure/gpt-4"
6
+ - Google Vertex: model="vertex_ai/claude-3-5-sonnet"
7
+ - And many more...
8
+
9
+ LiteLLM handles all the auth and format translation internally.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import logging
15
+ import uuid
16
+ from collections.abc import AsyncIterator
17
+ from typing import Any
18
+
19
+ from .base import Backend, BackendResponse, StreamEvent
20
+
21
+ logger = logging.getLogger(__name__)
22
+
23
+ try:
24
+ import litellm
25
+ from litellm import acompletion
26
+
27
+ LITELLM_AVAILABLE = True
28
+ except ImportError:
29
+ LITELLM_AVAILABLE = False
30
+ litellm = None # type: ignore
31
+ acompletion = None # type: ignore
32
+
33
+
34
+ # Model mapping: Anthropic model IDs -> LiteLLM model strings
35
+ BEDROCK_MODEL_MAP = {
36
+ # Claude 4.5
37
+ "claude-opus-4-5-20251101": "bedrock/anthropic.claude-opus-4-5-20251101-v1:0",
38
+ "claude-sonnet-4-5-20250929": "bedrock/anthropic.claude-sonnet-4-5-20250929-v1:0",
39
+ # Claude 4
40
+ "claude-opus-4-20250514": "bedrock/anthropic.claude-opus-4-20250514-v1:0",
41
+ "claude-sonnet-4-20250514": "bedrock/anthropic.claude-sonnet-4-20250514-v1:0",
42
+ # Claude 3.7
43
+ "claude-3-7-sonnet-20250219": "bedrock/anthropic.claude-3-7-sonnet-20250219-v1:0",
44
+ # Claude 3.5
45
+ "claude-3-5-sonnet-20241022": "bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0",
46
+ "claude-3-5-sonnet-20240620": "bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
47
+ "claude-3-5-haiku-20241022": "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
48
+ # Claude 3
49
+ "claude-3-opus-20240229": "bedrock/anthropic.claude-3-opus-20240229-v1:0",
50
+ "claude-3-sonnet-20240229": "bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
51
+ "claude-3-haiku-20240307": "bedrock/anthropic.claude-3-haiku-20240307-v1:0",
52
+ }
53
+
54
+ VERTEX_MODEL_MAP = {
55
+ "claude-3-5-sonnet-20241022": "vertex_ai/claude-3-5-sonnet-v2@20241022",
56
+ "claude-3-5-sonnet-20240620": "vertex_ai/claude-3-5-sonnet@20240620",
57
+ "claude-3-opus-20240229": "vertex_ai/claude-3-opus@20240229",
58
+ "claude-3-sonnet-20240229": "vertex_ai/claude-3-sonnet@20240229",
59
+ "claude-3-haiku-20240307": "vertex_ai/claude-3-haiku@20240307",
60
+ }
61
+
62
+
63
+ class LiteLLMBackend(Backend):
64
+ """Backend using LiteLLM for multi-provider support.
65
+
66
+ Supports any provider LiteLLM supports:
67
+ - bedrock: AWS Bedrock (uses AWS credentials)
68
+ - vertex_ai: Google Vertex AI (uses GCP credentials)
69
+ - azure: Azure OpenAI (uses Azure credentials)
70
+ - And 100+ more...
71
+ """
72
+
73
+ def __init__(
74
+ self,
75
+ provider: str = "bedrock",
76
+ region: str | None = None,
77
+ **kwargs: Any,
78
+ ):
79
+ """Initialize LiteLLM backend.
80
+
81
+ Args:
82
+ provider: LiteLLM provider prefix (bedrock, vertex_ai, azure, etc.)
83
+ region: Cloud region (provider-specific)
84
+ **kwargs: Additional provider-specific config
85
+ """
86
+ if not LITELLM_AVAILABLE:
87
+ raise ImportError(
88
+ "litellm is required for LiteLLMBackend. Install with: pip install litellm"
89
+ )
90
+
91
+ self.provider = provider
92
+ self.region = region
93
+ self.kwargs = kwargs
94
+
95
+ # Select model map based on provider
96
+ if provider == "bedrock":
97
+ self._model_map = BEDROCK_MODEL_MAP
98
+ # Set AWS region for litellm
99
+ if region:
100
+ litellm.set_verbose = False # Reduce noise
101
+ elif provider == "vertex_ai":
102
+ self._model_map = VERTEX_MODEL_MAP
103
+ else:
104
+ self._model_map = {}
105
+
106
+ logger.info(f"LiteLLM backend initialized (provider={provider})")
107
+
108
+ @property
109
+ def name(self) -> str:
110
+ return f"litellm-{self.provider}"
111
+
112
+ def map_model_id(self, anthropic_model: str) -> str:
113
+ """Map Anthropic model ID to LiteLLM model string."""
114
+ # Check direct mapping
115
+ if anthropic_model in self._model_map:
116
+ return self._model_map[anthropic_model]
117
+
118
+ # If already has provider prefix, use as-is
119
+ if "/" in anthropic_model:
120
+ return anthropic_model
121
+
122
+ # Fallback: construct provider/model format
123
+ return f"{self.provider}/{anthropic_model}"
124
+
125
+ def supports_model(self, model: str) -> bool:
126
+ """Check if model is supported."""
127
+ return "claude" in model.lower() or model in self._model_map
128
+
129
+ def _convert_messages_for_litellm(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
130
+ """Convert Anthropic message format to LiteLLM/OpenAI format.
131
+
132
+ LiteLLM expects OpenAI-style messages but handles most Anthropic
133
+ content blocks automatically.
134
+ """
135
+ converted = []
136
+ for msg in messages:
137
+ role = msg.get("role", "user")
138
+ content = msg.get("content", "")
139
+
140
+ # Handle string content directly
141
+ if isinstance(content, str):
142
+ converted.append({"role": role, "content": content})
143
+ continue
144
+
145
+ # Handle content blocks (Anthropic style)
146
+ if isinstance(content, list):
147
+ # Check if it's simple text blocks only
148
+ text_parts = []
149
+ has_complex_content = False
150
+
151
+ for block in content:
152
+ if isinstance(block, dict):
153
+ if block.get("type") == "text":
154
+ text_parts.append(block.get("text", ""))
155
+ elif block.get("type") in ("tool_use", "tool_result", "image"):
156
+ has_complex_content = True
157
+ break
158
+
159
+ if not has_complex_content and text_parts:
160
+ # Simple text - join into single string
161
+ converted.append({"role": role, "content": "\n".join(text_parts)})
162
+ else:
163
+ # Complex content - pass through (LiteLLM handles it)
164
+ converted.append({"role": role, "content": content})
165
+
166
+ return converted
167
+
168
+ def _to_anthropic_response(
169
+ self,
170
+ litellm_response: Any,
171
+ original_model: str,
172
+ ) -> dict[str, Any]:
173
+ """Convert LiteLLM/OpenAI response to Anthropic format."""
174
+ msg_id = f"msg_{uuid.uuid4().hex[:24]}"
175
+
176
+ # Extract content from OpenAI format
177
+ choice = litellm_response.choices[0]
178
+ message = choice.message
179
+
180
+ # Build Anthropic content blocks
181
+ content = []
182
+ if message.content:
183
+ content.append({"type": "text", "text": message.content})
184
+
185
+ # Handle tool calls if present
186
+ if hasattr(message, "tool_calls") and message.tool_calls:
187
+ for tc in message.tool_calls:
188
+ content.append(
189
+ {
190
+ "type": "tool_use",
191
+ "id": tc.id,
192
+ "name": tc.function.name,
193
+ "input": tc.function.arguments,
194
+ }
195
+ )
196
+
197
+ # Map stop reason
198
+ stop_reason_map = {
199
+ "stop": "end_turn",
200
+ "length": "max_tokens",
201
+ "tool_calls": "tool_use",
202
+ "content_filter": "end_turn",
203
+ }
204
+ stop_reason = stop_reason_map.get(choice.finish_reason, "end_turn")
205
+
206
+ # Build usage
207
+ usage = {
208
+ "input_tokens": getattr(litellm_response.usage, "prompt_tokens", 0),
209
+ "output_tokens": getattr(litellm_response.usage, "completion_tokens", 0),
210
+ }
211
+
212
+ return {
213
+ "id": msg_id,
214
+ "type": "message",
215
+ "role": "assistant",
216
+ "content": content,
217
+ "model": original_model,
218
+ "stop_reason": stop_reason,
219
+ "stop_sequence": None,
220
+ "usage": usage,
221
+ }
222
+
223
+ async def send_message(
224
+ self,
225
+ body: dict[str, Any],
226
+ headers: dict[str, str],
227
+ ) -> BackendResponse:
228
+ """Send message via LiteLLM."""
229
+ original_model = body.get("model", "claude-3-5-sonnet-20241022")
230
+ litellm_model = self.map_model_id(original_model)
231
+
232
+ try:
233
+ # Convert messages
234
+ messages = self._convert_messages_for_litellm(body.get("messages", []))
235
+
236
+ # Build kwargs for litellm
237
+ kwargs: dict[str, Any] = {
238
+ "model": litellm_model,
239
+ "messages": messages,
240
+ }
241
+
242
+ # Optional parameters
243
+ if "max_tokens" in body:
244
+ kwargs["max_tokens"] = body["max_tokens"]
245
+ if "temperature" in body:
246
+ kwargs["temperature"] = body["temperature"]
247
+ if "top_p" in body:
248
+ kwargs["top_p"] = body["top_p"]
249
+ if "stop_sequences" in body:
250
+ kwargs["stop"] = body["stop_sequences"]
251
+
252
+ # System prompt (Anthropic puts it in body, OpenAI in messages)
253
+ if "system" in body:
254
+ system = body["system"]
255
+ if isinstance(system, str):
256
+ kwargs["messages"].insert(0, {"role": "system", "content": system})
257
+ elif isinstance(system, list):
258
+ # Anthropic list format
259
+ system_text = " ".join(
260
+ s.get("text", "") if isinstance(s, dict) else str(s) for s in system
261
+ )
262
+ kwargs["messages"].insert(0, {"role": "system", "content": system_text})
263
+
264
+ # AWS region for Bedrock
265
+ if self.provider == "bedrock" and self.region:
266
+ kwargs["aws_region_name"] = self.region
267
+
268
+ logger.debug(f"LiteLLM request: model={litellm_model}")
269
+
270
+ # Make the call
271
+ response = await acompletion(**kwargs)
272
+
273
+ # Convert to Anthropic format
274
+ anthropic_response = self._to_anthropic_response(response, original_model)
275
+
276
+ return BackendResponse(
277
+ body=anthropic_response,
278
+ status_code=200,
279
+ headers={"content-type": "application/json"},
280
+ )
281
+
282
+ except Exception as e:
283
+ logger.error(f"LiteLLM error: {e}")
284
+
285
+ # Map to Anthropic error format
286
+ error_type = "api_error"
287
+ status_code = 500
288
+
289
+ error_str = str(e).lower()
290
+ if "authentication" in error_str or "credentials" in error_str:
291
+ error_type = "authentication_error"
292
+ status_code = 401
293
+ elif "rate" in error_str or "limit" in error_str:
294
+ error_type = "rate_limit_error"
295
+ status_code = 429
296
+ elif "not found" in error_str:
297
+ error_type = "not_found_error"
298
+ status_code = 404
299
+
300
+ return BackendResponse(
301
+ body={
302
+ "type": "error",
303
+ "error": {"type": error_type, "message": str(e)},
304
+ },
305
+ status_code=status_code,
306
+ error=str(e),
307
+ )
308
+
309
+ async def stream_message(
310
+ self,
311
+ body: dict[str, Any],
312
+ headers: dict[str, str],
313
+ ) -> AsyncIterator[StreamEvent]:
314
+ """Stream message via LiteLLM."""
315
+ original_model = body.get("model", "claude-3-5-sonnet-20241022")
316
+ litellm_model = self.map_model_id(original_model)
317
+
318
+ try:
319
+ messages = self._convert_messages_for_litellm(body.get("messages", []))
320
+
321
+ kwargs: dict[str, Any] = {
322
+ "model": litellm_model,
323
+ "messages": messages,
324
+ "stream": True,
325
+ }
326
+
327
+ if "max_tokens" in body:
328
+ kwargs["max_tokens"] = body["max_tokens"]
329
+ if "temperature" in body:
330
+ kwargs["temperature"] = body["temperature"]
331
+ if "system" in body:
332
+ system = body["system"]
333
+ if isinstance(system, str):
334
+ kwargs["messages"].insert(0, {"role": "system", "content": system})
335
+
336
+ if self.provider == "bedrock" and self.region:
337
+ kwargs["aws_region_name"] = self.region
338
+
339
+ msg_id = f"msg_{uuid.uuid4().hex[:24]}"
340
+
341
+ # Emit message_start
342
+ yield StreamEvent(
343
+ event_type="message_start",
344
+ data={
345
+ "type": "message_start",
346
+ "message": {
347
+ "id": msg_id,
348
+ "type": "message",
349
+ "role": "assistant",
350
+ "content": [],
351
+ "model": original_model,
352
+ "stop_reason": None,
353
+ "stop_sequence": None,
354
+ "usage": {"input_tokens": 0, "output_tokens": 0},
355
+ },
356
+ },
357
+ )
358
+
359
+ # Emit content_block_start
360
+ yield StreamEvent(
361
+ event_type="content_block_start",
362
+ data={
363
+ "type": "content_block_start",
364
+ "index": 0,
365
+ "content_block": {"type": "text", "text": ""},
366
+ },
367
+ )
368
+
369
+ # Stream content
370
+ response = await acompletion(**kwargs)
371
+ output_tokens = 0
372
+
373
+ async for chunk in response:
374
+ if hasattr(chunk, "choices") and chunk.choices:
375
+ delta = chunk.choices[0].delta
376
+ if hasattr(delta, "content") and delta.content:
377
+ yield StreamEvent(
378
+ event_type="content_block_delta",
379
+ data={
380
+ "type": "content_block_delta",
381
+ "index": 0,
382
+ "delta": {"type": "text_delta", "text": delta.content},
383
+ },
384
+ )
385
+ output_tokens += 1 # Rough estimate
386
+
387
+ # Emit content_block_stop
388
+ yield StreamEvent(
389
+ event_type="content_block_stop",
390
+ data={"type": "content_block_stop", "index": 0},
391
+ )
392
+
393
+ # Emit message_delta with stop reason
394
+ yield StreamEvent(
395
+ event_type="message_delta",
396
+ data={
397
+ "type": "message_delta",
398
+ "delta": {"stop_reason": "end_turn", "stop_sequence": None},
399
+ "usage": {"output_tokens": output_tokens},
400
+ },
401
+ )
402
+
403
+ # Emit message_stop
404
+ yield StreamEvent(
405
+ event_type="message_stop",
406
+ data={"type": "message_stop"},
407
+ )
408
+
409
+ except Exception as e:
410
+ logger.error(f"LiteLLM streaming error: {e}")
411
+ yield StreamEvent(
412
+ event_type="error",
413
+ data={
414
+ "type": "error",
415
+ "error": {"type": "api_error", "message": str(e)},
416
+ },
417
+ )
418
+
419
+ async def close(self) -> None: # noqa: B027
420
+ """Clean up (no-op for LiteLLM)."""
421
+ pass
headroom/cli.py CHANGED
@@ -78,12 +78,26 @@ def cmd_proxy(args: argparse.Namespace) -> int:
78
  memory_inject_tools=not args.no_memory_tools,
79
  memory_inject_context=not args.no_memory_context,
80
  memory_top_k=args.memory_top_k,
 
 
 
 
81
  )
82
 
83
  memory_status = "DISABLED"
84
  if config.memory_enabled:
85
  memory_status = f"ENABLED ({config.memory_backend})"
86
 
 
 
 
 
 
 
 
 
 
 
87
  print(f"""
88
  ╔═══════════════════════════════════════════════════════════════════════╗
89
  β•‘ HEADROOM PROXY β•‘
@@ -93,12 +107,14 @@ def cmd_proxy(args: argparse.Namespace) -> int:
93
  Starting proxy server...
94
 
95
  URL: http://{config.host}:{config.port}
 
96
  Optimization: {"ENABLED" if config.optimize else "DISABLED"}
97
  Caching: {"ENABLED" if config.cache_enabled else "DISABLED"}
98
  Rate Limit: {"ENABLED" if config.rate_limit_enabled else "DISABLED"}
99
  Memory: {memory_status}
100
 
101
  Usage with Claude Code:
 
102
  ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude
103
 
104
  Usage with OpenAI-compatible clients:
@@ -544,6 +560,29 @@ Documentation: https://github.com/headroom-sdk/headroom
544
  default=10,
545
  help="Number of memories to inject as context (default: 10)",
546
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
547
  proxy_parser.set_defaults(func=cmd_proxy)
548
 
549
  # Memory eval command
 
78
  memory_inject_tools=not args.no_memory_tools,
79
  memory_inject_context=not args.no_memory_context,
80
  memory_top_k=args.memory_top_k,
81
+ # Backend (Anthropic direct, Bedrock, or LiteLLM)
82
+ backend=args.backend,
83
+ bedrock_region=args.bedrock_region or args.region,
84
+ bedrock_profile=args.bedrock_profile,
85
  )
86
 
87
  memory_status = "DISABLED"
88
  if config.memory_enabled:
89
  memory_status = f"ENABLED ({config.memory_backend})"
90
 
91
+ region = args.bedrock_region or args.region
92
+ backend_status = "Anthropic (direct API)"
93
+ if config.backend != "anthropic":
94
+ # Normalize: "bedrock" -> "litellm-bedrock"
95
+ backend = config.backend
96
+ if not backend.startswith("litellm-"):
97
+ backend = f"litellm-{backend}"
98
+ provider = backend.replace("litellm-", "")
99
+ backend_status = f"{provider.upper()} via LiteLLM (region={region})"
100
+
101
  print(f"""
102
  ╔═══════════════════════════════════════════════════════════════════════╗
103
  β•‘ HEADROOM PROXY β•‘
 
107
  Starting proxy server...
108
 
109
  URL: http://{config.host}:{config.port}
110
+ Backend: {backend_status}
111
  Optimization: {"ENABLED" if config.optimize else "DISABLED"}
112
  Caching: {"ENABLED" if config.cache_enabled else "DISABLED"}
113
  Rate Limit: {"ENABLED" if config.rate_limit_enabled else "DISABLED"}
114
  Memory: {memory_status}
115
 
116
  Usage with Claude Code:
117
+ {"unset CLAUDE_CODE_USE_BEDROCK # Important!" if config.backend == "bedrock" else ""}
118
  ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude
119
 
120
  Usage with OpenAI-compatible clients:
 
560
  default=10,
561
  help="Number of memories to inject as context (default: 10)",
562
  )
563
+ # Backend configuration
564
+ proxy_parser.add_argument(
565
+ "--backend",
566
+ default="anthropic",
567
+ help=(
568
+ "API backend: 'anthropic' (direct), 'bedrock' (AWS), "
569
+ "or 'litellm-<provider>' (e.g., litellm-bedrock, litellm-vertex)"
570
+ ),
571
+ )
572
+ proxy_parser.add_argument(
573
+ "--region",
574
+ default="us-west-2",
575
+ help="Cloud region for Bedrock/Vertex/etc (default: us-west-2)",
576
+ )
577
+ proxy_parser.add_argument(
578
+ "--bedrock-region",
579
+ default=None,
580
+ help="(deprecated, use --region) AWS region for Bedrock",
581
+ )
582
+ proxy_parser.add_argument(
583
+ "--bedrock-profile",
584
+ help="AWS profile name for Bedrock (default: use default credentials)",
585
+ )
586
  proxy_parser.set_defaults(func=cmd_proxy)
587
 
588
  # Memory eval command
headroom/proxy/server.py CHANGED
@@ -53,6 +53,8 @@ except ImportError:
53
  # Add parent to path for imports
54
  sys.path.insert(0, str(Path(__file__).parent.parent.parent))
55
 
 
 
56
  from headroom.cache.compression_feedback import get_compression_feedback
57
  from headroom.cache.compression_store import get_compression_store
58
  from headroom.ccr import (
@@ -206,6 +208,12 @@ class ProxyConfig:
206
  openai_api_url: str | None = None # Custom OpenAI API URL override
207
  gemini_api_url: str | None = None # Custom Gemini API URL override
208
 
 
 
 
 
 
 
209
  # Optimization
210
  optimize: bool = True
211
  image_optimize: bool = True # Compress images using trained ML router
@@ -981,6 +989,29 @@ class HeadroomProxy:
981
  # HTTP client
982
  self.http_client: httpx.AsyncClient | None = None
983
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
984
  # Request counter for IDs
985
  self._request_counter = 0
986
  self._request_counter_lock = asyncio.Lock()
@@ -1583,7 +1614,70 @@ class HeadroomProxy:
1583
  if tools is not None:
1584
  body["tools"] = tools
1585
 
1586
- # Forward request
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1587
  url = f"{self.ANTHROPIC_API_URL}/v1/messages"
1588
 
1589
  try:
@@ -3506,6 +3600,101 @@ class HeadroomProxy:
3506
  media_type="text/event-stream",
3507
  )
3508
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3509
  async def handle_openai_chat(
3510
  self,
3511
  request: Request,
@@ -5948,6 +6137,17 @@ def run_server(
5948
  pool_info = f"max={config.max_connections}, keepalive={config.max_keepalive_connections}"
5949
  http2_status = "ENABLED" if config.http2 else "DISABLED"
5950
 
 
 
 
 
 
 
 
 
 
 
 
5951
  print(f"""
5952
  ╔══════════════════════════════════════════════════════════════════════╗
5953
  β•‘ HEADROOM PROXY SERVER β•‘
@@ -5955,6 +6155,7 @@ def run_server(
5955
  β•‘ Version: 1.0.0 β•‘
5956
  β•‘ Listening: http://{config.host}:{config.port:<5} β•‘
5957
  β•‘ Workers: {workers:<3} Concurrency Limit: {limit_concurrency:<5} β•‘
 
5958
  ╠══════════════════════════════════════════════════════════════════════╣
5959
  β•‘ FEATURES: β•‘
5960
  β•‘ Optimization: {"ENABLED " if config.optimize else "DISABLED"} β•‘
@@ -6045,6 +6246,23 @@ if __name__ == "__main__":
6045
  "--openai-api-url", help=f"Custom OpenAI API URL (default: {HeadroomProxy.OPENAI_API_URL})"
6046
  )
6047
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6048
  # Connection pool (scalability)
6049
  parser.add_argument(
6050
  "--max-connections",
@@ -6170,6 +6388,10 @@ if __name__ == "__main__":
6170
  host=_get_env_str("HEADROOM_HOST", args.host),
6171
  port=_get_env_int("HEADROOM_PORT", args.port),
6172
  openai_api_url=_get_env_str("OPENAI_TARGET_API_URL", args.openai_api_url),
 
 
 
 
6173
  optimize=optimize,
6174
  min_tokens_to_crush=_get_env_int("HEADROOM_MIN_TOKENS", args.min_tokens),
6175
  max_items_after_crush=_get_env_int("HEADROOM_MAX_ITEMS", args.max_items),
 
53
  # Add parent to path for imports
54
  sys.path.insert(0, str(Path(__file__).parent.parent.parent))
55
 
56
+ from headroom.backends import LiteLLMBackend
57
+ from headroom.backends.base import Backend
58
  from headroom.cache.compression_feedback import get_compression_feedback
59
  from headroom.cache.compression_store import get_compression_store
60
  from headroom.ccr import (
 
208
  openai_api_url: str | None = None # Custom OpenAI API URL override
209
  gemini_api_url: str | None = None # Custom Gemini API URL override
210
 
211
+ # Backend: "anthropic" (direct API), "bedrock" (AWS Bedrock), or "litellm-*" (via LiteLLM)
212
+ # LiteLLM backends: "litellm-bedrock", "litellm-vertex", "litellm-azure", etc.
213
+ backend: str = "anthropic"
214
+ bedrock_region: str = "us-west-2" # AWS region for Bedrock/LiteLLM
215
+ bedrock_profile: str | None = None # AWS profile (optional)
216
+
217
  # Optimization
218
  optimize: bool = True
219
  image_optimize: bool = True # Compress images using trained ML router
 
989
  # HTTP client
990
  self.http_client: httpx.AsyncClient | None = None
991
 
992
+ # Backend for Anthropic API (direct or via LiteLLM)
993
+ # Supports: "anthropic" (direct), "bedrock", "vertex", or "litellm-<provider>"
994
+ self.anthropic_backend: Backend | None = None
995
+ if config.backend != "anthropic":
996
+ # Normalize backend name: "bedrock" -> "litellm-bedrock"
997
+ backend = config.backend
998
+ if not backend.startswith("litellm-"):
999
+ backend = f"litellm-{backend}"
1000
+ provider = backend.replace("litellm-", "")
1001
+
1002
+ try:
1003
+ self.anthropic_backend = LiteLLMBackend(
1004
+ provider=provider,
1005
+ region=config.bedrock_region,
1006
+ )
1007
+ logger.info(
1008
+ f"LiteLLM backend enabled (provider={provider}, region={config.bedrock_region})"
1009
+ )
1010
+ except ImportError as e:
1011
+ logger.warning(f"LiteLLM backend not available: {e}")
1012
+ except Exception as e:
1013
+ logger.error(f"Failed to initialize LiteLLM backend: {e}")
1014
+
1015
  # Request counter for IDs
1016
  self._request_counter = 0
1017
  self._request_counter_lock = asyncio.Lock()
 
1614
  if tools is not None:
1615
  body["tools"] = tools
1616
 
1617
+ # Forward request - use Bedrock backend if configured, otherwise direct API
1618
+ if self.anthropic_backend is not None:
1619
+ # Route through Bedrock backend
1620
+ try:
1621
+ if stream:
1622
+ return await self._stream_response_bedrock(
1623
+ body,
1624
+ headers,
1625
+ "anthropic",
1626
+ model,
1627
+ request_id,
1628
+ original_tokens,
1629
+ optimized_tokens,
1630
+ tokens_saved,
1631
+ transforms_applied,
1632
+ tags,
1633
+ optimization_latency,
1634
+ )
1635
+ else:
1636
+ backend_response = await self.anthropic_backend.send_message(body, headers)
1637
+
1638
+ if backend_response.error:
1639
+ return JSONResponse(
1640
+ status_code=backend_response.status_code,
1641
+ content=backend_response.body,
1642
+ )
1643
+
1644
+ # Track metrics
1645
+ total_latency = (time.time() - start_time) * 1000
1646
+ usage = backend_response.body.get("usage", {})
1647
+ output_tokens = usage.get("output_tokens", 0)
1648
+
1649
+ await self.metrics.record_request(
1650
+ provider="bedrock",
1651
+ model=model,
1652
+ input_tokens=optimized_tokens,
1653
+ output_tokens=output_tokens,
1654
+ tokens_saved=tokens_saved,
1655
+ latency_ms=total_latency,
1656
+ cached=False,
1657
+ )
1658
+
1659
+ if self.cost_tracker:
1660
+ cost_usd = self.cost_tracker.estimate_cost(
1661
+ model, optimized_tokens, output_tokens
1662
+ )
1663
+ if cost_usd:
1664
+ self.cost_tracker.record_cost(cost_usd)
1665
+
1666
+ return JSONResponse(
1667
+ status_code=backend_response.status_code,
1668
+ content=backend_response.body,
1669
+ )
1670
+ except Exception as e:
1671
+ logger.error(f"[{request_id}] Bedrock backend error: {e}")
1672
+ return JSONResponse(
1673
+ status_code=500,
1674
+ content={
1675
+ "type": "error",
1676
+ "error": {"type": "api_error", "message": str(e)},
1677
+ },
1678
+ )
1679
+
1680
+ # Direct Anthropic API
1681
  url = f"{self.ANTHROPIC_API_URL}/v1/messages"
1682
 
1683
  try:
 
3600
  media_type="text/event-stream",
3601
  )
3602
 
3603
+ async def _stream_response_bedrock(
3604
+ self,
3605
+ body: dict,
3606
+ headers: dict,
3607
+ provider: str,
3608
+ model: str,
3609
+ request_id: str,
3610
+ original_tokens: int,
3611
+ optimized_tokens: int,
3612
+ tokens_saved: int,
3613
+ transforms_applied: list[str],
3614
+ tags: dict[str, str],
3615
+ optimization_latency: float,
3616
+ ) -> StreamingResponse:
3617
+ """Stream response from Bedrock backend with metrics tracking.
3618
+
3619
+ Translates Bedrock streaming events to Anthropic SSE format.
3620
+ """
3621
+ start_time = time.time()
3622
+
3623
+ # Mutable state for the generator
3624
+ stream_state: dict[str, Any] = {
3625
+ "input_tokens": 0,
3626
+ "output_tokens": 0,
3627
+ }
3628
+
3629
+ async def generate():
3630
+ try:
3631
+ assert self.anthropic_backend is not None
3632
+
3633
+ async for event in self.anthropic_backend.stream_message(body, headers):
3634
+ # Format as SSE
3635
+ if event.raw_sse:
3636
+ yield event.raw_sse.encode()
3637
+ else:
3638
+ sse_line = f"event: {event.event_type}\ndata: {json.dumps(event.data)}\n\n"
3639
+ yield sse_line.encode()
3640
+
3641
+ # Track usage from message_start event
3642
+ if event.event_type == "message_start":
3643
+ msg = event.data.get("message", {})
3644
+ usage = msg.get("usage", {})
3645
+ if "input_tokens" in usage:
3646
+ stream_state["input_tokens"] = usage["input_tokens"]
3647
+
3648
+ # Track output tokens from message_delta
3649
+ if event.event_type == "message_delta":
3650
+ usage = event.data.get("usage", {})
3651
+ if "output_tokens" in usage:
3652
+ stream_state["output_tokens"] = usage["output_tokens"]
3653
+
3654
+ # Handle errors
3655
+ if event.event_type == "error":
3656
+ logger.error(f"[{request_id}] Bedrock stream error: {event.data}")
3657
+
3658
+ except Exception as e:
3659
+ logger.error(f"[{request_id}] Bedrock streaming error: {e}")
3660
+ error_event = {
3661
+ "type": "error",
3662
+ "error": {"type": "api_error", "message": str(e)},
3663
+ }
3664
+ yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
3665
+
3666
+ finally:
3667
+ # Record metrics
3668
+ total_latency = (time.time() - start_time) * 1000
3669
+ output_tokens = stream_state["output_tokens"]
3670
+
3671
+ await self.metrics.record_request(
3672
+ provider="bedrock",
3673
+ model=model,
3674
+ input_tokens=optimized_tokens,
3675
+ output_tokens=output_tokens,
3676
+ tokens_saved=tokens_saved,
3677
+ latency_ms=total_latency,
3678
+ cached=False,
3679
+ )
3680
+
3681
+ if self.cost_tracker:
3682
+ cost_usd = self.cost_tracker.estimate_cost(
3683
+ model, optimized_tokens, output_tokens
3684
+ )
3685
+ if cost_usd:
3686
+ self.cost_tracker.record_cost(cost_usd)
3687
+
3688
+ if tokens_saved > 0:
3689
+ logger.info(
3690
+ f"[{request_id}] Bedrock {model}: saved {tokens_saved:,} tokens (streaming)"
3691
+ )
3692
+
3693
+ return StreamingResponse(
3694
+ generate(),
3695
+ media_type="text/event-stream",
3696
+ )
3697
+
3698
  async def handle_openai_chat(
3699
  self,
3700
  request: Request,
 
6137
  pool_info = f"max={config.max_connections}, keepalive={config.max_keepalive_connections}"
6138
  http2_status = "ENABLED" if config.http2 else "DISABLED"
6139
 
6140
+ # Backend status
6141
+ if config.backend == "anthropic":
6142
+ backend_status = "ANTHROPIC (direct API)"
6143
+ else:
6144
+ # Normalize: "bedrock" -> "litellm-bedrock"
6145
+ backend = config.backend
6146
+ if not backend.startswith("litellm-"):
6147
+ backend = f"litellm-{backend}"
6148
+ provider = backend.replace("litellm-", "")
6149
+ backend_status = f"{provider.upper()} via LiteLLM (region={config.bedrock_region})"
6150
+
6151
  print(f"""
6152
  ╔══════════════════════════════════════════════════════════════════════╗
6153
  β•‘ HEADROOM PROXY SERVER β•‘
 
6155
  β•‘ Version: 1.0.0 β•‘
6156
  β•‘ Listening: http://{config.host}:{config.port:<5} β•‘
6157
  β•‘ Workers: {workers:<3} Concurrency Limit: {limit_concurrency:<5} β•‘
6158
+ β•‘ Backend: {backend_status:<59}β•‘
6159
  ╠══════════════════════════════════════════════════════════════════════╣
6160
  β•‘ FEATURES: β•‘
6161
  β•‘ Optimization: {"ENABLED " if config.optimize else "DISABLED"} β•‘
 
6246
  "--openai-api-url", help=f"Custom OpenAI API URL (default: {HeadroomProxy.OPENAI_API_URL})"
6247
  )
6248
 
6249
+ # Backend (anthropic direct or bedrock)
6250
+ parser.add_argument(
6251
+ "--backend",
6252
+ choices=["anthropic", "bedrock"],
6253
+ default="anthropic",
6254
+ help="Backend for Anthropic API: 'anthropic' (direct) or 'bedrock' (AWS Bedrock)",
6255
+ )
6256
+ parser.add_argument(
6257
+ "--bedrock-region",
6258
+ default="us-west-2",
6259
+ help="AWS region for Bedrock backend (default: us-west-2)",
6260
+ )
6261
+ parser.add_argument(
6262
+ "--bedrock-profile",
6263
+ help="AWS profile for Bedrock backend (default: use default credentials)",
6264
+ )
6265
+
6266
  # Connection pool (scalability)
6267
  parser.add_argument(
6268
  "--max-connections",
 
6388
  host=_get_env_str("HEADROOM_HOST", args.host),
6389
  port=_get_env_int("HEADROOM_PORT", args.port),
6390
  openai_api_url=_get_env_str("OPENAI_TARGET_API_URL", args.openai_api_url),
6391
+ # Backend settings
6392
+ backend=_get_env_str("HEADROOM_BACKEND", args.backend), # type: ignore[arg-type]
6393
+ bedrock_region=_get_env_str("HEADROOM_BEDROCK_REGION", args.bedrock_region),
6394
+ bedrock_profile=args.bedrock_profile or os.environ.get("AWS_PROFILE"),
6395
  optimize=optimize,
6396
  min_tokens_to_crush=_get_env_int("HEADROOM_MIN_TOKENS", args.min_tokens),
6397
  max_items_after_crush=_get_env_int("HEADROOM_MAX_ITEMS", args.max_items),