chopratejas commited on
Commit
d90c6b8
·
1 Parent(s): 050d435

Commit message:

Browse files

Fix OpenAI streaming with backends and /v1 double-path bug

Add stream_openai_message() to LiteLLM and any-llm backends so
/v1/chat/completions with stream:true returns SSE events instead
of a JSON blob. Clients (Kilo Code, Cursor, etc.) were hanging
because the proxy ignored the stream flag when routing through
a backend.

Also strip trailing /v1 from OPENAI_TARGET_API_URL to prevent
double-path URLs like /v1/v1/models.

headroom/backends/anyllm.py CHANGED
@@ -454,6 +454,61 @@ class AnyLLMBackend(Backend):
454
 
455
  return BackendResponse(body=body, status_code=status_code, error=str(e))
456
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
457
  async def close(self) -> None:
458
  """Clean up (no-op for any-llm)."""
459
  pass
 
454
 
455
  return BackendResponse(body=body, status_code=status_code, error=str(e))
456
 
457
+ async def stream_openai_message(
458
+ self,
459
+ body: dict[str, Any],
460
+ headers: dict[str, str],
461
+ ) -> AsyncIterator[str]:
462
+ """Stream OpenAI-format chat completion via any-llm.
463
+
464
+ Yields SSE-formatted strings ready to send to the client.
465
+ """
466
+ original_model = body.get("model", "gpt-4o")
467
+
468
+ try:
469
+ kwargs: dict[str, Any] = {
470
+ "model": original_model,
471
+ "messages": body.get("messages", []),
472
+ "stream": True,
473
+ }
474
+
475
+ for param in [
476
+ "max_tokens",
477
+ "temperature",
478
+ "top_p",
479
+ "stop",
480
+ "tools",
481
+ "tool_choice",
482
+ "response_format",
483
+ "seed",
484
+ "n",
485
+ ]:
486
+ if param in body:
487
+ kwargs[param] = body[param]
488
+
489
+ if "stream_options" in body:
490
+ kwargs["stream_options"] = body["stream_options"]
491
+
492
+ stream_response = await self.llm.acompletion(**kwargs)
493
+
494
+ async for chunk in cast(AsyncIterator[Any], stream_response):
495
+ chunk_dict = chunk.model_dump(exclude_none=True, exclude_unset=True)
496
+ yield f"data: {json.dumps(chunk_dict)}\n\n"
497
+
498
+ yield "data: [DONE]\n\n"
499
+
500
+ except Exception as e:
501
+ logger.error(f"any-llm OpenAI streaming error: {e}")
502
+ error_data = {
503
+ "error": {
504
+ "message": str(e),
505
+ "type": "api_error",
506
+ "code": "backend_error",
507
+ }
508
+ }
509
+ yield f"data: {json.dumps(error_data)}\n\n"
510
+ yield "data: [DONE]\n\n"
511
+
512
  async def close(self) -> None:
513
  """Clean up (no-op for any-llm)."""
514
  pass
headroom/backends/base.py CHANGED
@@ -140,6 +140,30 @@ class Backend(ABC):
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
 
140
  """
141
  raise NotImplementedError(f"{self.name} backend does not support OpenAI format")
142
 
143
+ async def stream_openai_message(
144
+ self,
145
+ body: dict[str, Any],
146
+ headers: dict[str, str],
147
+ ) -> AsyncIterator[str]:
148
+ """Stream an OpenAI-format chat completion.
149
+
150
+ Yields SSE-formatted strings: 'data: {...}\\n\\n' for each chunk,
151
+ ending with 'data: [DONE]\\n\\n'.
152
+
153
+ Args:
154
+ body: Request body in OpenAI chat completion format (stream: true).
155
+ headers: Request headers.
156
+
157
+ Yields:
158
+ SSE-formatted strings ready to send to client.
159
+
160
+ Raises:
161
+ NotImplementedError: If backend doesn't support OpenAI streaming.
162
+ """
163
+ raise NotImplementedError(f"{self.name} backend does not support OpenAI streaming")
164
+ # Make this an async generator (yield never reached but needed for type)
165
+ yield "" # type: ignore[misc] # pragma: no cover
166
+
167
  async def close(self) -> None: # noqa: B027
168
  """Clean up resources (e.g., close HTTP clients)."""
169
  pass
headroom/backends/litellm.py CHANGED
@@ -819,3 +819,62 @@ class LiteLLMBackend(Backend):
819
  status_code=status_code,
820
  error=str(e),
821
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
819
  status_code=status_code,
820
  error=str(e),
821
  )
822
+
823
+ async def stream_openai_message(
824
+ self,
825
+ body: dict[str, Any],
826
+ headers: dict[str, str],
827
+ ) -> AsyncIterator[str]:
828
+ """Stream OpenAI-format chat completion via LiteLLM.
829
+
830
+ Yields SSE-formatted strings ready to send to the client.
831
+ """
832
+ original_model = body.get("model", "gpt-4")
833
+ litellm_model = self.map_model_id(original_model)
834
+
835
+ try:
836
+ kwargs: dict[str, Any] = {
837
+ "model": litellm_model,
838
+ "messages": body.get("messages", []),
839
+ "stream": True,
840
+ }
841
+
842
+ for param in [
843
+ "max_tokens",
844
+ "temperature",
845
+ "top_p",
846
+ "stop",
847
+ "tools",
848
+ "tool_choice",
849
+ "response_format",
850
+ "seed",
851
+ "n",
852
+ ]:
853
+ if param in body:
854
+ kwargs[param] = body[param]
855
+
856
+ if "stream_options" in body:
857
+ kwargs["stream_options"] = body["stream_options"]
858
+
859
+ if self.provider == "bedrock" and self.region:
860
+ kwargs["aws_region_name"] = self.region
861
+
862
+ response = await acompletion(**kwargs)
863
+
864
+ async for chunk in response:
865
+ chunk_dict = chunk.model_dump(exclude_none=True, exclude_unset=True)
866
+ yield f"data: {json.dumps(chunk_dict)}\n\n"
867
+
868
+ yield "data: [DONE]\n\n"
869
+
870
+ except Exception as e:
871
+ logger.error(f"LiteLLM OpenAI streaming error: {e}")
872
+ error_data = {
873
+ "error": {
874
+ "message": str(e),
875
+ "type": "api_error",
876
+ "code": "backend_error",
877
+ }
878
+ }
879
+ yield f"data: {json.dumps(error_data)}\n\n"
880
+ yield "data: [DONE]\n\n"
headroom/proxy/server.py CHANGED
@@ -1313,12 +1313,17 @@ class HeadroomProxy:
1313
  self.config = config
1314
 
1315
  # Override OPENAI_API_URL with config if set
 
1316
  if config.openai_api_url:
1317
- HeadroomProxy.OPENAI_API_URL = config.openai_api_url
 
 
 
1318
 
1319
  # Override GEMINI_API_URL with config if set
1320
  if config.gemini_api_url:
1321
- HeadroomProxy.GEMINI_API_URL = config.gemini_api_url
 
1322
 
1323
  # Initialize providers
1324
  self.anthropic_provider = AnthropicProvider()
@@ -1651,26 +1656,45 @@ class HeadroomProxy:
1651
  else:
1652
  logger.info("Smart Routing: DISABLED (legacy sequential mode)")
1653
 
1654
- # Eagerly load LLMLingua model at startup (avoids 5s delay on first request)
1655
- if self.config.llmlingua_enabled:
 
 
 
 
 
1656
  for transform in self.anthropic_pipeline.transforms:
1657
  if hasattr(transform, "eager_load_compressors"):
1658
  transform.eager_load_compressors()
1659
- self._llmlingua_status = "enabled"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1660
  break
1661
 
1662
- # LLMLingua status with helpful hint
1663
  if self._llmlingua_status == "enabled":
1664
  logger.info(
1665
  f"LLMLingua: ENABLED (device={self.config.llmlingua_device}, "
1666
  f"rate={self.config.llmlingua_target_rate})"
1667
  )
 
 
1668
  elif self._llmlingua_status == "lazy":
1669
  logger.info("LLMLingua: LAZY (will load when prose content detected)")
1670
- elif self._llmlingua_status == "available":
1671
- logger.info("LLMLingua: available but disabled (use --llmlingua)")
1672
- elif self._llmlingua_status == "unavailable":
1673
- logger.info("LLMLingua: not installed (pip install headroom-ai[llmlingua])")
1674
  elif self._llmlingua_status == "disabled":
1675
  logger.info("LLMLingua: DISABLED")
1676
 
@@ -4340,6 +4364,67 @@ class HeadroomProxy:
4340
  media_type="text/event-stream",
4341
  )
4342
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4343
  async def handle_openai_chat(
4344
  self,
4345
  request: Request,
@@ -4533,46 +4618,65 @@ class HeadroomProxy:
4533
  if tools is not None:
4534
  body["tools"] = tools
4535
 
4536
- # Route through LiteLLM backend if configured (Databricks, Bedrock, etc.)
4537
  if self.anthropic_backend is not None:
4538
  try:
4539
- # Use the backend's OpenAI-format method
4540
- backend_response = await self.anthropic_backend.send_openai_message(body, headers)
4541
-
4542
- if backend_response.error:
4543
- return JSONResponse(
4544
- status_code=backend_response.status_code,
4545
- content=backend_response.body,
 
 
 
 
 
 
 
 
 
 
 
 
 
4546
  )
4547
 
4548
- # Track metrics
4549
- total_latency = (time.time() - start_time) * 1000
4550
- usage = backend_response.body.get("usage", {})
4551
- output_tokens = usage.get("completion_tokens", 0)
4552
- total_input_tokens = usage.get("prompt_tokens", optimized_tokens)
4553
 
4554
- await self.metrics.record_request(
4555
- provider=self.anthropic_backend.name,
4556
- model=model,
4557
- input_tokens=total_input_tokens,
4558
- output_tokens=output_tokens,
4559
- tokens_saved=tokens_saved,
4560
- latency_ms=total_latency,
4561
- cached=False,
4562
- overhead_ms=optimization_latency,
4563
- pipeline_timing=pipeline_timing,
4564
- )
4565
 
4566
- if tokens_saved > 0:
4567
- logger.info(
4568
- f"[{request_id}] {model}: {original_tokens:,} → {optimized_tokens:,} "
4569
- f"(saved {tokens_saved:,} tokens) via {self.anthropic_backend.name}"
 
 
 
 
 
 
4570
  )
4571
 
4572
- return JSONResponse(
4573
- status_code=backend_response.status_code,
4574
- content=backend_response.body,
4575
- )
 
 
 
 
 
 
4576
  except Exception as e:
4577
  logger.error(f"[{request_id}] Backend error: {e}")
4578
  return JSONResponse(
 
1313
  self.config = config
1314
 
1315
  # Override OPENAI_API_URL with config if set
1316
+ # Strip trailing /v1 or /v1/ to avoid double-path (e.g., .../v1/v1/models)
1317
  if config.openai_api_url:
1318
+ url = config.openai_api_url.rstrip("/")
1319
+ if url.endswith("/v1"):
1320
+ url = url[:-3]
1321
+ HeadroomProxy.OPENAI_API_URL = url
1322
 
1323
  # Override GEMINI_API_URL with config if set
1324
  if config.gemini_api_url:
1325
+ gurl = config.gemini_api_url.rstrip("/")
1326
+ HeadroomProxy.GEMINI_API_URL = gurl
1327
 
1328
  # Initialize providers
1329
  self.anthropic_provider = AnthropicProvider()
 
1656
  else:
1657
  logger.info("Smart Routing: DISABLED (legacy sequential mode)")
1658
 
1659
+ # Eagerly load ML compressors at startup (avoids download on first request)
1660
+ # Kompress requires [ml] extra (torch + transformers). If not installed, skip.
1661
+ self._kompress_status = "not installed"
1662
+ from headroom.transforms.kompress_compressor import is_kompress_available
1663
+
1664
+ if is_kompress_available() and self.config.optimize:
1665
+ logger.info("Kompress: Downloading model (first-time only)...")
1666
  for transform in self.anthropic_pipeline.transforms:
1667
  if hasattr(transform, "eager_load_compressors"):
1668
  transform.eager_load_compressors()
1669
+ self._kompress_status = "enabled"
1670
+ break
1671
+ if self._kompress_status == "enabled":
1672
+ logger.info("Kompress: ENABLED (ModernBERT token compressor)")
1673
+ else:
1674
+ if self.config.optimize:
1675
+ logger.info(
1676
+ "Kompress: not installed (pip install headroom-ai[ml] for ML compression)"
1677
+ )
1678
+
1679
+ # LLMLingua fallback (only loads if Kompress is not available)
1680
+ if self._kompress_status != "enabled" and self.config.llmlingua_enabled:
1681
+ for transform in self.anthropic_pipeline.transforms:
1682
+ if hasattr(transform, "_get_llmlingua"):
1683
+ llmlingua = transform._get_llmlingua()
1684
+ if llmlingua:
1685
+ self._llmlingua_status = "enabled"
1686
  break
1687
 
1688
+ # LLMLingua status
1689
  if self._llmlingua_status == "enabled":
1690
  logger.info(
1691
  f"LLMLingua: ENABLED (device={self.config.llmlingua_device}, "
1692
  f"rate={self.config.llmlingua_target_rate})"
1693
  )
1694
+ elif self._kompress_status == "enabled":
1695
+ logger.info("LLMLingua: skipped (Kompress is active)")
1696
  elif self._llmlingua_status == "lazy":
1697
  logger.info("LLMLingua: LAZY (will load when prose content detected)")
 
 
 
 
1698
  elif self._llmlingua_status == "disabled":
1699
  logger.info("LLMLingua: DISABLED")
1700
 
 
4364
  media_type="text/event-stream",
4365
  )
4366
 
4367
+ async def _stream_openai_via_backend(
4368
+ self,
4369
+ body: dict,
4370
+ headers: dict,
4371
+ model: str,
4372
+ request_id: str,
4373
+ start_time: float,
4374
+ original_tokens: int,
4375
+ optimized_tokens: int,
4376
+ tokens_saved: int,
4377
+ transforms_applied: list[str],
4378
+ tags: dict[str, str],
4379
+ optimization_latency: float,
4380
+ pipeline_timing: dict[str, float] | None = None,
4381
+ ) -> StreamingResponse:
4382
+ """Stream OpenAI chat completion response from backend.
4383
+
4384
+ Routes stream:true requests through the backend's stream_openai_message(),
4385
+ yielding SSE events to the client.
4386
+ """
4387
+ assert self.anthropic_backend is not None
4388
+
4389
+ async def generate():
4390
+ try:
4391
+ async for sse_chunk in self.anthropic_backend.stream_openai_message(body, headers):
4392
+ yield sse_chunk.encode() if isinstance(sse_chunk, str) else sse_chunk
4393
+ except Exception as e:
4394
+ logger.error(f"[{request_id}] Backend streaming error: {e}")
4395
+ error_data = {
4396
+ "error": {
4397
+ "message": str(e),
4398
+ "type": "api_error",
4399
+ "code": "backend_error",
4400
+ }
4401
+ }
4402
+ yield f"data: {json.dumps(error_data)}\n\n".encode()
4403
+ yield b"data: [DONE]\n\n"
4404
+ finally:
4405
+ total_latency = (time.time() - start_time) * 1000
4406
+ await self.metrics.record_request(
4407
+ provider=self.anthropic_backend.name,
4408
+ model=model,
4409
+ input_tokens=optimized_tokens,
4410
+ output_tokens=0, # Unknown in streaming
4411
+ tokens_saved=tokens_saved,
4412
+ latency_ms=total_latency,
4413
+ cached=False,
4414
+ overhead_ms=optimization_latency,
4415
+ pipeline_timing=pipeline_timing,
4416
+ )
4417
+ if tokens_saved > 0:
4418
+ logger.info(
4419
+ f"[{request_id}] {model}: {original_tokens:,} → {optimized_tokens:,} "
4420
+ f"(saved {tokens_saved:,} tokens) via {self.anthropic_backend.name} [stream]"
4421
+ )
4422
+
4423
+ return StreamingResponse(
4424
+ generate(),
4425
+ media_type="text/event-stream",
4426
+ )
4427
+
4428
  async def handle_openai_chat(
4429
  self,
4430
  request: Request,
 
4618
  if tools is not None:
4619
  body["tools"] = tools
4620
 
4621
+ # Route through LiteLLM/any-llm backend if configured
4622
  if self.anthropic_backend is not None:
4623
  try:
4624
+ if stream:
4625
+ # Streaming: use stream_openai_message() → SSE events
4626
+ return await self._stream_openai_via_backend(
4627
+ body,
4628
+ headers,
4629
+ model,
4630
+ request_id,
4631
+ start_time,
4632
+ original_tokens,
4633
+ optimized_tokens,
4634
+ tokens_saved,
4635
+ transforms_applied,
4636
+ tags,
4637
+ optimization_latency,
4638
+ pipeline_timing=pipeline_timing,
4639
+ )
4640
+ else:
4641
+ # Non-streaming: use send_openai_message() → JSON
4642
+ backend_response = await self.anthropic_backend.send_openai_message(
4643
+ body, headers
4644
  )
4645
 
4646
+ if backend_response.error:
4647
+ return JSONResponse(
4648
+ status_code=backend_response.status_code,
4649
+ content=backend_response.body,
4650
+ )
4651
 
4652
+ # Track metrics
4653
+ total_latency = (time.time() - start_time) * 1000
4654
+ usage = backend_response.body.get("usage", {})
4655
+ output_tokens = usage.get("completion_tokens", 0)
4656
+ total_input_tokens = usage.get("prompt_tokens", optimized_tokens)
 
 
 
 
 
 
4657
 
4658
+ await self.metrics.record_request(
4659
+ provider=self.anthropic_backend.name,
4660
+ model=model,
4661
+ input_tokens=total_input_tokens,
4662
+ output_tokens=output_tokens,
4663
+ tokens_saved=tokens_saved,
4664
+ latency_ms=total_latency,
4665
+ cached=False,
4666
+ overhead_ms=optimization_latency,
4667
+ pipeline_timing=pipeline_timing,
4668
  )
4669
 
4670
+ if tokens_saved > 0:
4671
+ logger.info(
4672
+ f"[{request_id}] {model}: {original_tokens:,} → {optimized_tokens:,} "
4673
+ f"(saved {tokens_saved:,} tokens) via {self.anthropic_backend.name}"
4674
+ )
4675
+
4676
+ return JSONResponse(
4677
+ status_code=backend_response.status_code,
4678
+ content=backend_response.body,
4679
+ )
4680
  except Exception as e:
4681
  logger.error(f"[{request_id}] Backend error: {e}")
4682
  return JSONResponse(
headroom/transforms/kompress_compressor.py CHANGED
@@ -140,10 +140,12 @@ def _load_kompress(device: str = "auto") -> tuple[HeadroomCompressorModel, Any]:
140
 
141
 
142
  def is_kompress_available() -> bool:
143
- """Check if Kompress dependencies are available."""
144
  try:
145
  import huggingface_hub # noqa: F401
146
  import safetensors # noqa: F401
 
 
147
 
148
  return True
149
  except ImportError:
 
140
 
141
 
142
  def is_kompress_available() -> bool:
143
+ """Check if Kompress dependencies are available (requires [ml] extra)."""
144
  try:
145
  import huggingface_hub # noqa: F401
146
  import safetensors # noqa: F401
147
+ import torch # noqa: F401
148
+ import transformers # noqa: F401
149
 
150
  return True
151
  except ImportError:
tests/test_backend_bugs.py CHANGED
@@ -290,3 +290,62 @@ class TestVertexModelMap:
290
 
291
  def test_claude_3_legacy(self):
292
  assert "claude-3-haiku-20240307" in _VERTEX_MODEL_MAP
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
290
 
291
  def test_claude_3_legacy(self):
292
  assert "claude-3-haiku-20240307" in _VERTEX_MODEL_MAP
293
+
294
+
295
+ # =============================================================================
296
+ # URL Normalization (trailing /v1 stripping)
297
+ # =============================================================================
298
+
299
+ pytest.importorskip("fastapi")
300
+
301
+
302
+ class TestOpenAIURLNormalization:
303
+ """Test that OPENAI_TARGET_API_URL with /v1 suffix is normalized."""
304
+
305
+ def test_v1_suffix_stripped(self):
306
+ from headroom.proxy.server import HeadroomProxy, ProxyConfig
307
+
308
+ original = HeadroomProxy.OPENAI_API_URL
309
+ try:
310
+ config = ProxyConfig(
311
+ openai_api_url="http://localhost:4000/v1",
312
+ optimize=False,
313
+ cache_enabled=False,
314
+ rate_limit_enabled=False,
315
+ )
316
+ proxy = HeadroomProxy(config)
317
+ assert proxy.OPENAI_API_URL == "http://localhost:4000"
318
+ finally:
319
+ HeadroomProxy.OPENAI_API_URL = original
320
+
321
+ def test_v1_slash_suffix_stripped(self):
322
+ from headroom.proxy.server import HeadroomProxy, ProxyConfig
323
+
324
+ original = HeadroomProxy.OPENAI_API_URL
325
+ try:
326
+ config = ProxyConfig(
327
+ openai_api_url="http://localhost:4000/v1/",
328
+ optimize=False,
329
+ cache_enabled=False,
330
+ rate_limit_enabled=False,
331
+ )
332
+ proxy = HeadroomProxy(config)
333
+ assert proxy.OPENAI_API_URL == "http://localhost:4000"
334
+ finally:
335
+ HeadroomProxy.OPENAI_API_URL = original
336
+
337
+ def test_no_v1_unchanged(self):
338
+ from headroom.proxy.server import HeadroomProxy, ProxyConfig
339
+
340
+ original = HeadroomProxy.OPENAI_API_URL
341
+ try:
342
+ config = ProxyConfig(
343
+ openai_api_url="http://localhost:4000",
344
+ optimize=False,
345
+ cache_enabled=False,
346
+ rate_limit_enabled=False,
347
+ )
348
+ proxy = HeadroomProxy(config)
349
+ assert proxy.OPENAI_API_URL == "http://localhost:4000"
350
+ finally:
351
+ HeadroomProxy.OPENAI_API_URL = original
tests/test_openai_streaming_backend.py ADDED
@@ -0,0 +1,262 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Test OpenAI /v1/chat/completions streaming through headroom proxy backends.
2
+
3
+ Proves that streaming works end-to-end: client → headroom proxy → backend → OpenAI API.
4
+
5
+ Two test modes:
6
+ 1. Real API test (requires OPENAI_API_KEY): hits actual OpenAI with gpt-4o-mini
7
+ 2. Mock test: proves the proxy returns SSE when stream:true with a backend configured
8
+
9
+ Run with:
10
+ OPENAI_API_KEY=sk-... pytest tests/test_openai_streaming_backend.py -v
11
+ """
12
+
13
+ import os
14
+ from unittest.mock import AsyncMock, MagicMock, patch
15
+
16
+ import pytest
17
+
18
+ fastapi = pytest.importorskip("fastapi")
19
+ httpx = pytest.importorskip("httpx")
20
+
21
+ from fastapi.testclient import TestClient # noqa: E402
22
+
23
+ from headroom.backends.base import BackendResponse # noqa: E402
24
+ from headroom.proxy.server import ProxyConfig, create_app # noqa: E402
25
+
26
+ # =============================================================================
27
+ # Real API test (requires OPENAI_API_KEY)
28
+ # =============================================================================
29
+
30
+
31
+ @pytest.mark.skipif(not os.environ.get("OPENAI_API_KEY"), reason="OPENAI_API_KEY not set")
32
+ class TestOpenAIStreamingRealAPI:
33
+ """Test streaming with real OpenAI API calls through the proxy."""
34
+
35
+ @pytest.fixture
36
+ def openai_api_key(self):
37
+ return os.environ["OPENAI_API_KEY"]
38
+
39
+ @pytest.fixture
40
+ def direct_proxy_client(self):
41
+ """Proxy with NO backend — direct to OpenAI. This is the baseline."""
42
+ config = ProxyConfig(
43
+ optimize=False,
44
+ cache_enabled=False,
45
+ rate_limit_enabled=False,
46
+ )
47
+ app = create_app(config)
48
+ with TestClient(app) as client:
49
+ yield client
50
+
51
+ @pytest.fixture
52
+ def litellm_backend_client(self):
53
+ """Proxy with litellm-openai backend — routes through LiteLLM."""
54
+ config = ProxyConfig(
55
+ optimize=False,
56
+ cache_enabled=False,
57
+ rate_limit_enabled=False,
58
+ backend="litellm-openai",
59
+ )
60
+ app = create_app(config)
61
+ with TestClient(app) as client:
62
+ yield client
63
+
64
+ def test_baseline_streaming_works_direct(self, direct_proxy_client, openai_api_key):
65
+ """Baseline: streaming through proxy WITHOUT backend works (direct to OpenAI)."""
66
+ response = direct_proxy_client.post(
67
+ "/v1/chat/completions",
68
+ json={
69
+ "model": "gpt-4o-mini",
70
+ "messages": [{"role": "user", "content": "Say 'hello' and nothing else."}],
71
+ "stream": True,
72
+ "max_tokens": 10,
73
+ },
74
+ headers={"Authorization": f"Bearer {openai_api_key}"},
75
+ )
76
+
77
+ assert response.status_code == 200, f"Got {response.status_code}: {response.text[:200]}"
78
+
79
+ content_type = response.headers.get("content-type", "")
80
+ assert "text/event-stream" in content_type, (
81
+ f"Direct proxy streaming broken: got content-type '{content_type}'"
82
+ )
83
+
84
+ # Verify we got actual SSE chunks
85
+ body = response.text
86
+ assert "data: " in body, "No SSE data chunks in response"
87
+ assert "data: [DONE]" in body, "Missing [DONE] terminator"
88
+
89
+ def test_streaming_with_litellm_backend(self, litellm_backend_client, openai_api_key):
90
+ """CRITICAL: streaming through proxy WITH litellm backend must also stream.
91
+
92
+ This test fails before the fix — the proxy returns a JSON blob
93
+ instead of SSE events, causing clients to hang.
94
+ """
95
+ response = litellm_backend_client.post(
96
+ "/v1/chat/completions",
97
+ json={
98
+ "model": "gpt-4o-mini",
99
+ "messages": [{"role": "user", "content": "Say 'hello' and nothing else."}],
100
+ "stream": True,
101
+ "max_tokens": 10,
102
+ },
103
+ headers={"Authorization": f"Bearer {openai_api_key}"},
104
+ )
105
+
106
+ assert response.status_code == 200, f"Got {response.status_code}: {response.text[:200]}"
107
+
108
+ content_type = response.headers.get("content-type", "")
109
+ assert "text/event-stream" in content_type, (
110
+ f"STREAMING BUG: litellm backend returned '{content_type}' instead of "
111
+ f"'text/event-stream'. Client sees a JSON blob, not SSE events.\n"
112
+ f"Response body (first 300 chars): {response.text[:300]}"
113
+ )
114
+
115
+ # Verify SSE format
116
+ body = response.text
117
+ assert "data: " in body, "No SSE data chunks in streaming response"
118
+
119
+ def test_non_streaming_with_litellm_backend(self, litellm_backend_client, openai_api_key):
120
+ """Non-streaming with backend should return normal JSON (sanity check)."""
121
+ response = litellm_backend_client.post(
122
+ "/v1/chat/completions",
123
+ json={
124
+ "model": "gpt-4o-mini",
125
+ "messages": [{"role": "user", "content": "Say 'hello' and nothing else."}],
126
+ "stream": False,
127
+ "max_tokens": 10,
128
+ },
129
+ headers={"Authorization": f"Bearer {openai_api_key}"},
130
+ )
131
+
132
+ assert response.status_code == 200, f"Got {response.status_code}: {response.text[:200]}"
133
+
134
+ content_type = response.headers.get("content-type", "")
135
+ assert "application/json" in content_type
136
+
137
+ data = response.json()
138
+ assert "choices" in data
139
+ assert data["choices"][0]["message"]["content"]
140
+
141
+
142
+ # =============================================================================
143
+ # Mock test (no API key needed — proves the routing bug)
144
+ # =============================================================================
145
+
146
+
147
+ class TestOpenAIStreamingMock:
148
+ """Prove the streaming bug with mocks — no API key needed."""
149
+
150
+ def test_streaming_request_returns_sse_not_json(self):
151
+ """When stream:true with a backend, content-type MUST be text/event-stream.
152
+
153
+ This test FAILS before the fix: the proxy calls send_openai_message()
154
+ (non-streaming) and returns application/json even though stream:true.
155
+ """
156
+ config = ProxyConfig(
157
+ optimize=False,
158
+ cache_enabled=False,
159
+ rate_limit_enabled=False,
160
+ backend="anyllm",
161
+ anyllm_provider="openai",
162
+ )
163
+
164
+ mock_backend = MagicMock()
165
+ mock_backend.name = "anyllm-openai"
166
+ mock_backend.send_openai_message = AsyncMock(
167
+ return_value=BackendResponse(
168
+ body={
169
+ "id": "chatcmpl-123",
170
+ "object": "chat.completion",
171
+ "model": "test-model",
172
+ "choices": [
173
+ {
174
+ "index": 0,
175
+ "message": {"role": "assistant", "content": "Hello!"},
176
+ "finish_reason": "stop",
177
+ }
178
+ ],
179
+ "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
180
+ },
181
+ status_code=200,
182
+ headers={"content-type": "application/json"},
183
+ )
184
+ )
185
+
186
+ with patch("headroom.proxy.server.AnyLLMBackend", return_value=mock_backend):
187
+ app = create_app(config)
188
+
189
+ with TestClient(app) as client:
190
+ response = client.post(
191
+ "/v1/chat/completions",
192
+ json={
193
+ "model": "test-model",
194
+ "messages": [{"role": "user", "content": "hello"}],
195
+ "stream": True,
196
+ },
197
+ headers={"Authorization": "Bearer test-key"},
198
+ )
199
+
200
+ assert response.status_code == 200, (
201
+ f"Got {response.status_code}: {response.text[:200]}"
202
+ )
203
+
204
+ content_type = response.headers.get("content-type", "")
205
+ assert "text/event-stream" in content_type, (
206
+ f"STREAMING BUG: stream:true with backend returned '{content_type}' "
207
+ f"instead of 'text/event-stream'. The proxy ignored the stream flag "
208
+ f"and returned a JSON blob. Clients expecting SSE will hang.\n"
209
+ f"Response: {response.text[:300]}"
210
+ )
211
+
212
+ def test_non_streaming_still_returns_json(self):
213
+ """Sanity: stream:false with backend should return JSON as before."""
214
+ config = ProxyConfig(
215
+ optimize=False,
216
+ cache_enabled=False,
217
+ rate_limit_enabled=False,
218
+ backend="anyllm",
219
+ anyllm_provider="openai",
220
+ )
221
+
222
+ mock_backend = MagicMock()
223
+ mock_backend.name = "anyllm-openai"
224
+ mock_backend.send_openai_message = AsyncMock(
225
+ return_value=BackendResponse(
226
+ body={
227
+ "id": "chatcmpl-123",
228
+ "object": "chat.completion",
229
+ "model": "test-model",
230
+ "choices": [
231
+ {
232
+ "index": 0,
233
+ "message": {"role": "assistant", "content": "Hello!"},
234
+ "finish_reason": "stop",
235
+ }
236
+ ],
237
+ "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
238
+ },
239
+ status_code=200,
240
+ headers={"content-type": "application/json"},
241
+ )
242
+ )
243
+
244
+ with patch("headroom.proxy.server.AnyLLMBackend", return_value=mock_backend):
245
+ app = create_app(config)
246
+
247
+ with TestClient(app) as client:
248
+ response = client.post(
249
+ "/v1/chat/completions",
250
+ json={
251
+ "model": "test-model",
252
+ "messages": [{"role": "user", "content": "hello"}],
253
+ "stream": False,
254
+ },
255
+ headers={"Authorization": "Bearer test-key"},
256
+ )
257
+
258
+ assert response.status_code == 200
259
+ content_type = response.headers.get("content-type", "")
260
+ assert "application/json" in content_type
261
+ data = response.json()
262
+ assert data["choices"][0]["message"]["content"] == "Hello!"