chopratejas commited on
Commit
24e92f8
·
1 Parent(s): 0c64bc4

Add LiteLLM backend routing for OpenAI endpoint and Magika content detection

Browse files

- Route /v1/chat/completions through configured LiteLLM backend (Bedrock, Azure,
Databricks, etc.) instead of hardcoding to OpenAI API
- Add send_openai_message() method to Backend base class and LiteLLMBackend
- Add Databricks provider to PROVIDER_REGISTRY
- Add /serving-endpoints/{model}/invocations endpoint for Databricks CLI compatibility
- Integrate Magika ML-based content detection in ContentRouter for improved accuracy
- Fall back to regex-based detection when Magika is unavailable

headroom/backends/base.py CHANGED
@@ -117,6 +117,29 @@ class Backend(ABC):
117
  """
118
  ...
119
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
120
  async def close(self) -> None: # noqa: B027
121
  """Clean up resources (e.g., close HTTP clients)."""
122
  pass
 
117
  """
118
  ...
119
 
120
+ async def send_openai_message(
121
+ self,
122
+ body: dict[str, Any],
123
+ headers: dict[str, str],
124
+ ) -> BackendResponse:
125
+ """Send an OpenAI-format message request.
126
+
127
+ Unlike send_message(), this takes OpenAI-format input and returns
128
+ OpenAI-format output (no Anthropic conversion). Optional - only
129
+ implemented by backends that support OpenAI-compatible APIs.
130
+
131
+ Args:
132
+ body: Request body in OpenAI chat completion format.
133
+ headers: Request headers.
134
+
135
+ Returns:
136
+ BackendResponse with body in OpenAI chat completion format.
137
+
138
+ Raises:
139
+ NotImplementedError: If backend doesn't support OpenAI format.
140
+ """
141
+ raise NotImplementedError(f"{self.name} backend does not support OpenAI format")
142
+
143
  async def close(self) -> None: # noqa: B027
144
  """Clean up resources (e.g., close HTTP clients)."""
145
  pass
headroom/backends/litellm.py CHANGED
@@ -198,6 +198,15 @@ PROVIDER_REGISTRY: dict[str, ProviderConfig] = {
198
  uses_region=True,
199
  env_vars=["AZURE_API_KEY", "AZURE_API_BASE"],
200
  ),
 
 
 
 
 
 
 
 
 
201
  }
202
 
203
 
@@ -599,3 +608,134 @@ class LiteLLMBackend(Backend):
599
  async def close(self) -> None: # noqa: B027
600
  """Clean up (no-op for LiteLLM)."""
601
  pass
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
198
  uses_region=True,
199
  env_vars=["AZURE_API_KEY", "AZURE_API_BASE"],
200
  ),
201
+ "databricks": ProviderConfig(
202
+ name="databricks",
203
+ display_name="Databricks",
204
+ model_map={}, # Pass through - Databricks uses custom model names
205
+ pass_through=True,
206
+ uses_region=False,
207
+ env_vars=["DATABRICKS_API_KEY", "DATABRICKS_API_BASE"],
208
+ model_format_hint="databricks-meta-llama-3-1-70b-instruct, databricks-dbrx-instruct, etc.",
209
+ ),
210
  }
211
 
212
 
 
608
  async def close(self) -> None: # noqa: B027
609
  """Clean up (no-op for LiteLLM)."""
610
  pass
611
+
612
+ async def send_openai_message(
613
+ self,
614
+ body: dict[str, Any],
615
+ headers: dict[str, str],
616
+ ) -> BackendResponse:
617
+ """Send OpenAI-format message via LiteLLM.
618
+
619
+ Unlike send_message(), this takes OpenAI-format input and returns
620
+ OpenAI-format output (no Anthropic conversion).
621
+
622
+ Args:
623
+ body: OpenAI chat completion request body
624
+ headers: Request headers (ignored, auth from env vars)
625
+
626
+ Returns:
627
+ BackendResponse with OpenAI-format body
628
+ """
629
+ original_model = body.get("model", "gpt-4")
630
+ litellm_model = self.map_model_id(original_model)
631
+
632
+ try:
633
+ # Build kwargs - messages already in OpenAI format
634
+ kwargs: dict[str, Any] = {
635
+ "model": litellm_model,
636
+ "messages": body.get("messages", []),
637
+ }
638
+
639
+ # Pass through OpenAI parameters
640
+ for param in [
641
+ "max_tokens",
642
+ "temperature",
643
+ "top_p",
644
+ "stop",
645
+ "tools",
646
+ "tool_choice",
647
+ "response_format",
648
+ "seed",
649
+ "n",
650
+ ]:
651
+ if param in body:
652
+ kwargs[param] = body[param]
653
+
654
+ # Provider-specific config
655
+ if self.provider == "bedrock" and self.region:
656
+ kwargs["aws_region_name"] = self.region
657
+ elif self.provider == "databricks":
658
+ # Databricks uses env vars for auth
659
+ pass
660
+
661
+ logger.debug(f"LiteLLM OpenAI request: model={litellm_model}")
662
+
663
+ # Make the call
664
+ response = await acompletion(**kwargs)
665
+
666
+ # Convert ModelResponse to dict (OpenAI format)
667
+ response_dict = {
668
+ "id": response.id,
669
+ "object": "chat.completion",
670
+ "created": response.created,
671
+ "model": original_model,
672
+ "choices": [
673
+ {
674
+ "index": c.index,
675
+ "message": {
676
+ "role": c.message.role,
677
+ "content": c.message.content,
678
+ **(
679
+ {
680
+ "tool_calls": [
681
+ {
682
+ "id": tc.id,
683
+ "type": "function",
684
+ "function": {
685
+ "name": tc.function.name,
686
+ "arguments": tc.function.arguments,
687
+ },
688
+ }
689
+ for tc in c.message.tool_calls
690
+ ]
691
+ }
692
+ if c.message.tool_calls
693
+ else {}
694
+ ),
695
+ },
696
+ "finish_reason": c.finish_reason,
697
+ }
698
+ for c in response.choices
699
+ ],
700
+ "usage": {
701
+ "prompt_tokens": response.usage.prompt_tokens,
702
+ "completion_tokens": response.usage.completion_tokens,
703
+ "total_tokens": response.usage.total_tokens,
704
+ },
705
+ }
706
+
707
+ return BackendResponse(
708
+ body=response_dict,
709
+ status_code=200,
710
+ headers={"content-type": "application/json"},
711
+ )
712
+
713
+ except Exception as e:
714
+ logger.error(f"LiteLLM OpenAI error: {e}")
715
+
716
+ # Map to OpenAI error format
717
+ error_type = "api_error"
718
+ status_code = 500
719
+
720
+ error_str = str(e).lower()
721
+ if "authentication" in error_str or "credentials" in error_str:
722
+ error_type = "invalid_api_key"
723
+ status_code = 401
724
+ elif "rate" in error_str or "limit" in error_str:
725
+ error_type = "rate_limit_exceeded"
726
+ status_code = 429
727
+ elif "not found" in error_str:
728
+ error_type = "model_not_found"
729
+ status_code = 404
730
+
731
+ return BackendResponse(
732
+ body={
733
+ "error": {
734
+ "message": str(e),
735
+ "type": error_type,
736
+ "code": error_type,
737
+ }
738
+ },
739
+ status_code=status_code,
740
+ error=str(e),
741
+ )
headroom/proxy/server.py CHANGED
@@ -4015,6 +4015,60 @@ class HeadroomProxy:
4015
  body["messages"] = optimized_messages
4016
  if tools is not None:
4017
  body["tools"] = tools
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4018
  url = f"{self.OPENAI_API_URL}/v1/chat/completions"
4019
 
4020
  try:
@@ -4196,6 +4250,59 @@ class HeadroomProxy:
4196
  headers=response_headers,
4197
  )
4198
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4199
  # =========================================================================
4200
  # OpenAI Batch API with Compression
4201
  # =========================================================================
@@ -6186,6 +6293,21 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI:
6186
  """Gemini countTokens API with compression applied."""
6187
  return await proxy.handle_gemini_count_tokens(request, model)
6188
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6189
  # =========================================================================
6190
  # Passthrough Endpoints (no compression needed)
6191
  # =========================================================================
 
4015
  body["messages"] = optimized_messages
4016
  if tools is not None:
4017
  body["tools"] = tools
4018
+
4019
+ # Route through LiteLLM backend if configured (Databricks, Bedrock, etc.)
4020
+ if self.anthropic_backend is not None:
4021
+ try:
4022
+ # Use the backend's OpenAI-format method
4023
+ backend_response = await self.anthropic_backend.send_openai_message(body, headers)
4024
+
4025
+ if backend_response.error:
4026
+ return JSONResponse(
4027
+ status_code=backend_response.status_code,
4028
+ content=backend_response.body,
4029
+ )
4030
+
4031
+ # Track metrics
4032
+ total_latency = (time.time() - start_time) * 1000
4033
+ usage = backend_response.body.get("usage", {})
4034
+ output_tokens = usage.get("completion_tokens", 0)
4035
+ total_input_tokens = usage.get("prompt_tokens", optimized_tokens)
4036
+
4037
+ await self.metrics.record_request(
4038
+ provider=self.anthropic_backend.name,
4039
+ model=model,
4040
+ input_tokens=total_input_tokens,
4041
+ output_tokens=output_tokens,
4042
+ tokens_saved=tokens_saved,
4043
+ latency_ms=total_latency,
4044
+ cached=False,
4045
+ overhead_ms=optimization_latency,
4046
+ )
4047
+
4048
+ if tokens_saved > 0:
4049
+ logger.info(
4050
+ f"[{request_id}] {model}: {original_tokens:,} → {optimized_tokens:,} "
4051
+ f"(saved {tokens_saved:,} tokens) via {self.anthropic_backend.name}"
4052
+ )
4053
+
4054
+ return JSONResponse(
4055
+ status_code=backend_response.status_code,
4056
+ content=backend_response.body,
4057
+ )
4058
+ except Exception as e:
4059
+ logger.error(f"[{request_id}] Backend error: {e}")
4060
+ return JSONResponse(
4061
+ status_code=500,
4062
+ content={
4063
+ "error": {
4064
+ "message": str(e),
4065
+ "type": "api_error",
4066
+ "code": "backend_error",
4067
+ }
4068
+ },
4069
+ )
4070
+
4071
+ # Direct OpenAI API (no backend configured)
4072
  url = f"{self.OPENAI_API_URL}/v1/chat/completions"
4073
 
4074
  try:
 
4250
  headers=response_headers,
4251
  )
4252
 
4253
+ # =========================================================================
4254
+ # Databricks Native API
4255
+ # =========================================================================
4256
+
4257
+ async def handle_databricks_invocations(
4258
+ self,
4259
+ request: Request,
4260
+ model: str,
4261
+ ) -> Response | StreamingResponse:
4262
+ """Handle Databricks native /serving-endpoints/{model}/invocations endpoint.
4263
+
4264
+ This enables using the Databricks CLI directly with Headroom:
4265
+ databricks serving-endpoints query <model> --profile HEADROOM --json '{"messages": [...]}'
4266
+
4267
+ The request/response format is identical to OpenAI chat completions,
4268
+ so we inject the model from the path and delegate to handle_openai_chat.
4269
+ """
4270
+ request_id = await self._next_request_id()
4271
+
4272
+ try:
4273
+ body = await request.json()
4274
+ except Exception as e:
4275
+ logger.error(f"[{request_id}] Failed to parse Databricks request body: {e}")
4276
+ return JSONResponse(
4277
+ status_code=400,
4278
+ content={
4279
+ "error": {"message": f"Invalid JSON: {e}", "type": "invalid_request_error"}
4280
+ },
4281
+ )
4282
+
4283
+ # Inject model from path into body (Databricks CLI passes model in URL, not body)
4284
+ body["model"] = model
4285
+
4286
+ logger.info(f"[{request_id}] Databricks invocation: model={model}")
4287
+
4288
+ # Create a new request with the modified body
4289
+ # We reuse the OpenAI chat handler since the format is identical
4290
+ from starlette.requests import Request as StarletteRequest
4291
+
4292
+ # Build new scope with the body already parsed
4293
+ scope = dict(request.scope)
4294
+
4295
+ # Create a simple receive function that returns our modified body
4296
+ body_bytes = json.dumps(body).encode()
4297
+
4298
+ async def receive():
4299
+ return {"type": "http.request", "body": body_bytes}
4300
+
4301
+ modified_request = StarletteRequest(scope, receive)
4302
+
4303
+ # Delegate to the OpenAI chat handler (same format)
4304
+ return await self.handle_openai_chat(modified_request)
4305
+
4306
  # =========================================================================
4307
  # OpenAI Batch API with Compression
4308
  # =========================================================================
 
6293
  """Gemini countTokens API with compression applied."""
6294
  return await proxy.handle_gemini_count_tokens(request, model)
6295
 
6296
+ # =========================================================================
6297
+ # Databricks Native Endpoints
6298
+ # =========================================================================
6299
+
6300
+ @app.post("/serving-endpoints/{model}/invocations")
6301
+ async def databricks_invocations(request: Request, model: str):
6302
+ """Databricks native serving endpoint - compatible with Databricks CLI.
6303
+
6304
+ This allows using the Databricks CLI directly with Headroom proxy:
6305
+ databricks serving-endpoints query <model> --profile HEADROOM --json '{"messages": [...]}'
6306
+
6307
+ The request format is identical to OpenAI chat completions.
6308
+ """
6309
+ return await proxy.handle_databricks_invocations(request, model)
6310
+
6311
  # =========================================================================
6312
  # Passthrough Endpoints (no compression needed)
6313
  # =========================================================================
headroom/transforms/content_router.py CHANGED
@@ -45,10 +45,45 @@ from typing import Any
45
  from ..config import DEFAULT_EXCLUDE_TOOLS, TransformResult
46
  from ..tokenizer import Tokenizer
47
  from .base import Transform
48
- from .content_detector import ContentType, detect_content_type
49
 
50
  logger = logging.getLogger(__name__)
51
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52
 
53
  def _create_content_signature(
54
  content_type: str,
@@ -595,7 +630,7 @@ class ContentRouter(Transform):
595
  return CompressionStrategy.MIXED
596
 
597
  # 2. Detect content type from content itself
598
- detection = detect_content_type(content)
599
  return self._strategy_from_detection(detection)
600
 
601
  def _strategy_from_detection(self, detection: Any) -> CompressionStrategy:
@@ -1183,7 +1218,7 @@ class ContentRouter(Transform):
1183
  continue
1184
 
1185
  # Detect content type for protection decisions
1186
- detection = detect_content_type(content)
1187
  is_code = detection.content_type == ContentType.SOURCE_CODE
1188
 
1189
  # Protection 2: Don't compress recent CODE
 
45
  from ..config import DEFAULT_EXCLUDE_TOOLS, TransformResult
46
  from ..tokenizer import Tokenizer
47
  from .base import Transform
48
+ from .content_detector import ContentType, DetectionResult, detect_content_type
49
 
50
  logger = logging.getLogger(__name__)
51
 
52
+ # Use Magika-based detector if available, fallback to regex
53
+ _magika_detector = None
54
+ _USE_MAGIKA = False
55
+ try:
56
+ from ..compression.detector import get_detector
57
+
58
+ _magika_detector = get_detector(prefer_magika=True)
59
+ _USE_MAGIKA = True
60
+ logger.info("ContentRouter: Using Magika ML-based content detection")
61
+ except ImportError:
62
+ logger.debug("Magika not available, using regex-based detection")
63
+
64
+
65
+ def _detect_content(content: str) -> DetectionResult:
66
+ """Detect content type using Magika if available, else regex fallback."""
67
+ if _USE_MAGIKA and _magika_detector:
68
+ result = _magika_detector.detect(content)
69
+ # Map Magika ContentType to router's expected format
70
+ type_map = {
71
+ "json": ContentType.JSON_ARRAY,
72
+ "code": ContentType.SOURCE_CODE,
73
+ "log": ContentType.BUILD_OUTPUT,
74
+ "markdown": ContentType.PLAIN_TEXT,
75
+ "text": ContentType.PLAIN_TEXT,
76
+ "unknown": ContentType.PLAIN_TEXT,
77
+ }
78
+ mapped_type = type_map.get(result.content_type.value, ContentType.PLAIN_TEXT)
79
+ return DetectionResult(
80
+ content_type=mapped_type,
81
+ confidence=result.confidence,
82
+ metadata={"language": result.language, "raw_label": result.raw_label},
83
+ )
84
+ else:
85
+ return detect_content_type(content)
86
+
87
 
88
  def _create_content_signature(
89
  content_type: str,
 
630
  return CompressionStrategy.MIXED
631
 
632
  # 2. Detect content type from content itself
633
+ detection = _detect_content(content)
634
  return self._strategy_from_detection(detection)
635
 
636
  def _strategy_from_detection(self, detection: Any) -> CompressionStrategy:
 
1218
  continue
1219
 
1220
  # Detect content type for protection decisions
1221
+ detection = _detect_content(content)
1222
  is_code = detection.content_type == ContentType.SOURCE_CODE
1223
 
1224
  # Protection 2: Don't compress recent CODE