chopratejas commited on
Commit
f99f02a
·
1 Parent(s): cb7b00d

Fix streaming tool calls, compressed request bodies, beacon field names; bump to 0.5.13

Browse files

- Fix LiteLLM stream_message: emit tool_use blocks from delta.tool_calls,
set stop_reason from finish_reason (fixes silent MCP tool call failures)
- Fix _convert_messages_for_litellm: convert Anthropic tool_result/tool_use
to OpenAI role=tool/tool_calls format (fixes 500 on tool round-trips)
- Fix proxy request body parsing: decompress zstd/gzip/deflate/brotli
Content-Encoding before JSON decode (fixes Codex UnicodeDecodeError crash)
- Fix telemetry beacon field names to match /stats endpoint
(tokens.total_before_compression, not tokens.original)
- Add dedupe_telemetry.py script for hourly Supabase row deduplication

headroom/__init__.py CHANGED
@@ -153,7 +153,7 @@ from .transforms import (
153
  TransformPipeline,
154
  )
155
 
156
- __version__ = "0.5.12"
157
 
158
  __all__ = [
159
  # Main client
 
153
  TransformPipeline,
154
  )
155
 
156
+ __version__ = "0.5.13"
157
 
158
  __all__ = [
159
  # Main client
headroom/backends/litellm.py CHANGED
@@ -474,8 +474,12 @@ class LiteLLMBackend(Backend):
474
  def _convert_messages_for_litellm(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
475
  """Convert Anthropic message format to LiteLLM/OpenAI format.
476
 
477
- LiteLLM expects OpenAI-style messages but handles most Anthropic
478
- content blocks automatically.
 
 
 
 
479
  """
480
  converted = []
481
  for msg in messages:
@@ -489,24 +493,68 @@ class LiteLLMBackend(Backend):
489
 
490
  # Handle content blocks (Anthropic style)
491
  if isinstance(content, list):
492
- # Check if it's simple text blocks only
493
  text_parts = []
494
- has_complex_content = False
 
495
 
496
  for block in content:
497
- if isinstance(block, dict):
498
- if block.get("type") == "text":
499
- text_parts.append(block.get("text", ""))
500
- elif block.get("type") in ("tool_use", "tool_result", "image"):
501
- has_complex_content = True
502
- break
503
-
504
- if not has_complex_content and text_parts:
505
- # Simple text - join into single string
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
506
  converted.append({"role": role, "content": "\n".join(text_parts)})
507
  else:
508
- # Complex content - pass through (LiteLLM handles it)
509
- converted.append({"role": role, "content": content})
510
 
511
  return converted
512
 
@@ -662,7 +710,12 @@ class LiteLLMBackend(Backend):
662
  body: dict[str, Any],
663
  headers: dict[str, str],
664
  ) -> AsyncIterator[StreamEvent]:
665
- """Stream message via LiteLLM."""
 
 
 
 
 
666
  original_model = body.get("model", "claude-3-5-sonnet-20241022")
667
  litellm_model = self.map_model_id(original_model)
668
 
@@ -691,6 +744,11 @@ class LiteLLMBackend(Backend):
691
  system = body["system"]
692
  if isinstance(system, str):
693
  kwargs["messages"].insert(0, {"role": "system", "content": system})
 
 
 
 
 
694
 
695
  if self.provider == "bedrock" and self.region:
696
  kwargs["aws_region_name"] = self.region
@@ -715,46 +773,126 @@ class LiteLLMBackend(Backend):
715
  },
716
  )
717
 
718
- # Emit content_block_start
719
- yield StreamEvent(
720
- event_type="content_block_start",
721
- data={
722
- "type": "content_block_start",
723
- "index": 0,
724
- "content_block": {"type": "text", "text": ""},
725
- },
726
- )
727
-
728
- # Stream content
729
  response = await acompletion(**kwargs)
730
  output_tokens = 0
 
 
 
 
731
 
732
  async for chunk in response:
733
- if hasattr(chunk, "choices") and chunk.choices:
734
- delta = chunk.choices[0].delta
735
- if hasattr(delta, "content") and delta.content:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
736
  yield StreamEvent(
737
- event_type="content_block_delta",
738
  data={
739
- "type": "content_block_delta",
740
- "index": 0,
741
- "delta": {"type": "text_delta", "text": delta.content},
742
  },
743
  )
744
- output_tokens += 1 # Rough estimate
745
 
746
- # Emit content_block_stop
747
- yield StreamEvent(
748
- event_type="content_block_stop",
749
- data={"type": "content_block_stop", "index": 0},
750
- )
 
 
 
 
 
 
 
 
 
 
 
751
 
752
- # Emit message_delta with stop reason
753
  yield StreamEvent(
754
  event_type="message_delta",
755
  data={
756
  "type": "message_delta",
757
- "delta": {"stop_reason": "end_turn", "stop_sequence": None},
758
  "usage": {"output_tokens": output_tokens},
759
  },
760
  )
 
474
  def _convert_messages_for_litellm(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
475
  """Convert Anthropic message format to LiteLLM/OpenAI format.
476
 
477
+ Anthropic and OpenAI have different representations for tool calls:
478
+ - Anthropic: assistant content blocks with type=tool_use, user content blocks with type=tool_result
479
+ - OpenAI: assistant message with tool_calls field, separate role=tool messages
480
+
481
+ This method converts Anthropic-style messages to OpenAI-style so LiteLLM
482
+ can send them to any provider.
483
  """
484
  converted = []
485
  for msg in messages:
 
493
 
494
  # Handle content blocks (Anthropic style)
495
  if isinstance(content, list):
496
+ # Separate blocks by type
497
  text_parts = []
498
+ tool_use_blocks = []
499
+ tool_result_blocks = []
500
 
501
  for block in content:
502
+ if not isinstance(block, dict):
503
+ continue
504
+ block_type = block.get("type", "")
505
+ if block_type == "text":
506
+ text_parts.append(block.get("text", ""))
507
+ elif block_type == "tool_use":
508
+ tool_use_blocks.append(block)
509
+ elif block_type == "tool_result":
510
+ tool_result_blocks.append(block)
511
+
512
+ # tool_result blocks → OpenAI "tool" role messages
513
+ if tool_result_blocks:
514
+ # Include any text from the user message first
515
+ if text_parts:
516
+ converted.append({"role": "user", "content": "\n".join(text_parts)})
517
+ for tr in tool_result_blocks:
518
+ tr_content = tr.get("content", "")
519
+ if isinstance(tr_content, list):
520
+ tr_content = "\n".join(
521
+ b.get("text", "") for b in tr_content if b.get("type") == "text"
522
+ )
523
+ converted.append(
524
+ {
525
+ "role": "tool",
526
+ "tool_call_id": tr["tool_use_id"],
527
+ "content": str(tr_content),
528
+ }
529
+ )
530
+ continue
531
+
532
+ # tool_use blocks → OpenAI assistant message with tool_calls
533
+ if tool_use_blocks:
534
+ assistant_msg: dict[str, Any] = {"role": "assistant"}
535
+ if text_parts:
536
+ assistant_msg["content"] = "\n".join(text_parts)
537
+ else:
538
+ assistant_msg["content"] = None
539
+ assistant_msg["tool_calls"] = [
540
+ {
541
+ "id": tu["id"],
542
+ "type": "function",
543
+ "function": {
544
+ "name": tu["name"],
545
+ "arguments": json.dumps(tu.get("input", {})),
546
+ },
547
+ }
548
+ for tu in tool_use_blocks
549
+ ]
550
+ converted.append(assistant_msg)
551
+ continue
552
+
553
+ # Simple text only
554
+ if text_parts:
555
  converted.append({"role": role, "content": "\n".join(text_parts)})
556
  else:
557
+ converted.append({"role": role, "content": ""})
 
558
 
559
  return converted
560
 
 
710
  body: dict[str, Any],
711
  headers: dict[str, str],
712
  ) -> AsyncIterator[StreamEvent]:
713
+ """Stream message via LiteLLM.
714
+
715
+ Translates OpenAI streaming chunks into Anthropic SSE events.
716
+ Handles both text content and tool_calls dynamically — block types
717
+ are emitted based on what LiteLLM actually returns, not hardcoded.
718
+ """
719
  original_model = body.get("model", "claude-3-5-sonnet-20241022")
720
  litellm_model = self.map_model_id(original_model)
721
 
 
744
  system = body["system"]
745
  if isinstance(system, str):
746
  kwargs["messages"].insert(0, {"role": "system", "content": system})
747
+ elif isinstance(system, list):
748
+ system_text = " ".join(
749
+ s.get("text", "") if isinstance(s, dict) else str(s) for s in system
750
+ )
751
+ kwargs["messages"].insert(0, {"role": "system", "content": system_text})
752
 
753
  if self.provider == "bedrock" and self.region:
754
  kwargs["aws_region_name"] = self.region
 
773
  },
774
  )
775
 
776
+ # Stream content — blocks emitted dynamically based on response
 
 
 
 
 
 
 
 
 
 
777
  response = await acompletion(**kwargs)
778
  output_tokens = 0
779
+ current_block_index = -1
780
+ active_block_type: str | None = None # "text" or "tool_use"
781
+ tool_block_map: dict[int, int] = {} # litellm tc.index → SSE block index
782
+ stop_reason = "end_turn"
783
 
784
  async for chunk in response:
785
+ if not hasattr(chunk, "choices") or not chunk.choices:
786
+ continue
787
+
788
+ choice = chunk.choices[0]
789
+ delta = choice.delta
790
+
791
+ # Check finish_reason to set stop_reason
792
+ if choice.finish_reason == "tool_calls":
793
+ stop_reason = "tool_use"
794
+ elif choice.finish_reason == "stop":
795
+ stop_reason = "end_turn"
796
+ elif choice.finish_reason == "length":
797
+ stop_reason = "max_tokens"
798
+
799
+ # Handle tool_calls in the delta
800
+ if hasattr(delta, "tool_calls") and delta.tool_calls:
801
+ for tc in delta.tool_calls:
802
+ idx = tc.index if tc.index is not None else 0
803
+ if idx not in tool_block_map:
804
+ # Close previous block if open
805
+ if active_block_type is not None:
806
+ yield StreamEvent(
807
+ event_type="content_block_stop",
808
+ data={
809
+ "type": "content_block_stop",
810
+ "index": current_block_index,
811
+ },
812
+ )
813
+ # Open a new tool_use block
814
+ current_block_index += 1
815
+ tool_block_map[idx] = current_block_index
816
+ active_block_type = "tool_use"
817
+ tool_id = tc.id or f"toolu_{uuid.uuid4().hex[:24]}"
818
+ tool_name = tc.function.name if tc.function and tc.function.name else ""
819
+ yield StreamEvent(
820
+ event_type="content_block_start",
821
+ data={
822
+ "type": "content_block_start",
823
+ "index": current_block_index,
824
+ "content_block": {
825
+ "type": "tool_use",
826
+ "id": tool_id,
827
+ "name": tool_name,
828
+ "input": {},
829
+ },
830
+ },
831
+ )
832
+
833
+ # Emit argument deltas
834
+ if tc.function and tc.function.arguments:
835
+ block_idx = tool_block_map[idx]
836
+ yield StreamEvent(
837
+ event_type="content_block_delta",
838
+ data={
839
+ "type": "content_block_delta",
840
+ "index": block_idx,
841
+ "delta": {
842
+ "type": "input_json_delta",
843
+ "partial_json": tc.function.arguments,
844
+ },
845
+ },
846
+ )
847
+ output_tokens += 1
848
+
849
+ # Handle text content in the delta
850
+ elif hasattr(delta, "content") and delta.content:
851
+ if active_block_type != "text":
852
+ # Close previous block if open
853
+ if active_block_type is not None:
854
+ yield StreamEvent(
855
+ event_type="content_block_stop",
856
+ data={
857
+ "type": "content_block_stop",
858
+ "index": current_block_index,
859
+ },
860
+ )
861
+ # Open a new text block
862
+ current_block_index += 1
863
+ active_block_type = "text"
864
  yield StreamEvent(
865
+ event_type="content_block_start",
866
  data={
867
+ "type": "content_block_start",
868
+ "index": current_block_index,
869
+ "content_block": {"type": "text", "text": ""},
870
  },
871
  )
 
872
 
873
+ yield StreamEvent(
874
+ event_type="content_block_delta",
875
+ data={
876
+ "type": "content_block_delta",
877
+ "index": current_block_index,
878
+ "delta": {"type": "text_delta", "text": delta.content},
879
+ },
880
+ )
881
+ output_tokens += 1
882
+
883
+ # Close the last open block
884
+ if active_block_type is not None:
885
+ yield StreamEvent(
886
+ event_type="content_block_stop",
887
+ data={"type": "content_block_stop", "index": current_block_index},
888
+ )
889
 
890
+ # Emit message_delta with correct stop reason
891
  yield StreamEvent(
892
  event_type="message_delta",
893
  data={
894
  "type": "message_delta",
895
+ "delta": {"stop_reason": stop_reason, "stop_sequence": None},
896
  "usage": {"output_tokens": output_tokens},
897
  },
898
  )
headroom/proxy/server.py CHANGED
@@ -538,6 +538,73 @@ MAX_SSE_BUFFER_SIZE = 10 * 1024 * 1024
538
  # Maximum message array length (prevents DoS from deeply nested payloads)
539
  MAX_MESSAGE_ARRAY_LENGTH = 10000
540
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
541
  # Maximum compression cache sessions (prevents unbounded memory growth)
542
  MAX_COMPRESSION_CACHE_SESSIONS = 500
543
 
@@ -2195,7 +2262,13 @@ class HeadroomProxy:
2195
  """
2196
  # Check bypass header
2197
  if request.headers.get("x-headroom-bypass", "").lower() == "true":
2198
- body = await request.json()
 
 
 
 
 
 
2199
  messages = body.get("messages", [])
2200
  return JSONResponse(
2201
  {
@@ -2210,7 +2283,7 @@ class HeadroomProxy:
2210
  )
2211
 
2212
  try:
2213
- body = await request.json()
2214
  except Exception:
2215
  return JSONResponse(
2216
  status_code=400,
@@ -2327,15 +2400,15 @@ class HeadroomProxy:
2327
 
2328
  # Parse request
2329
  try:
2330
- body = await request.json()
2331
- except json.JSONDecodeError as e:
2332
  return JSONResponse(
2333
  status_code=400,
2334
  content={
2335
  "type": "error",
2336
  "error": {
2337
  "type": "invalid_request_error",
2338
- "message": f"Invalid JSON in request body: {e!s}",
2339
  },
2340
  },
2341
  )
@@ -3288,15 +3361,15 @@ class HeadroomProxy:
3288
 
3289
  # Parse request
3290
  try:
3291
- body = await request.json()
3292
- except json.JSONDecodeError as e:
3293
  return JSONResponse(
3294
  status_code=400,
3295
  content={
3296
  "type": "error",
3297
  "error": {
3298
  "type": "invalid_request_error",
3299
- "message": f"Invalid JSON in request body: {e!s}",
3300
  },
3301
  },
3302
  )
@@ -3726,14 +3799,14 @@ class HeadroomProxy:
3726
 
3727
  # Parse request
3728
  try:
3729
- body = await request.json()
3730
- except json.JSONDecodeError as e:
3731
  return JSONResponse(
3732
  status_code=400,
3733
  content={
3734
  "error": {
3735
  "code": 400,
3736
- "message": f"Invalid JSON in request body: {e!s}",
3737
  "status": "INVALID_ARGUMENT",
3738
  }
3739
  },
@@ -5138,13 +5211,13 @@ class HeadroomProxy:
5138
 
5139
  # Parse request
5140
  try:
5141
- body = await request.json()
5142
- except json.JSONDecodeError as e:
5143
  return JSONResponse(
5144
  status_code=400,
5145
  content={
5146
  "error": {
5147
- "message": f"Invalid JSON in request body: {e!s}",
5148
  "type": "invalid_request_error",
5149
  "code": "invalid_json",
5150
  }
@@ -5735,7 +5808,7 @@ class HeadroomProxy:
5735
  request_id = await self._next_request_id()
5736
 
5737
  try:
5738
- body = await request.json()
5739
  except Exception as e:
5740
  logger.error(f"[{request_id}] Failed to parse Databricks request body: {e}")
5741
  return JSONResponse(
@@ -5789,13 +5862,13 @@ class HeadroomProxy:
5789
 
5790
  # Parse request
5791
  try:
5792
- body = await request.json()
5793
- except json.JSONDecodeError as e:
5794
  return JSONResponse(
5795
  status_code=400,
5796
  content={
5797
  "error": {
5798
- "message": f"Invalid JSON in request body: {e!s}",
5799
  "type": "invalid_request_error",
5800
  "code": "invalid_json",
5801
  }
@@ -6272,13 +6345,13 @@ class HeadroomProxy:
6272
 
6273
  # Parse request
6274
  try:
6275
- body = await request.json()
6276
- except json.JSONDecodeError as e:
6277
  return JSONResponse(
6278
  status_code=400,
6279
  content={
6280
  "error": {
6281
- "message": f"Invalid JSON in request body: {e!s}",
6282
  "type": "invalid_request_error",
6283
  "code": "invalid_json",
6284
  }
@@ -6434,13 +6507,13 @@ class HeadroomProxy:
6434
 
6435
  # Parse request
6436
  try:
6437
- body = await request.json()
6438
- except json.JSONDecodeError as e:
6439
  return JSONResponse(
6440
  status_code=400,
6441
  content={
6442
  "error": {
6443
- "message": f"Invalid JSON in request body: {e!s}",
6444
  "code": 400,
6445
  }
6446
  },
@@ -6698,13 +6771,13 @@ class HeadroomProxy:
6698
 
6699
  # Parse request
6700
  try:
6701
- body = await request.json()
6702
- except json.JSONDecodeError as e:
6703
  return JSONResponse(
6704
  status_code=400,
6705
  content={
6706
  "error": {
6707
- "message": f"Invalid JSON in request body: {e!s}",
6708
  "code": 400,
6709
  }
6710
  },
@@ -6767,13 +6840,13 @@ class HeadroomProxy:
6767
 
6768
  # Parse request
6769
  try:
6770
- body = await request.json()
6771
- except json.JSONDecodeError as e:
6772
  return JSONResponse(
6773
  status_code=400,
6774
  content={
6775
  "error": {
6776
- "message": f"Invalid JSON in request body: {e!s}",
6777
  "code": 400,
6778
  }
6779
  },
 
538
  # Maximum message array length (prevents DoS from deeply nested payloads)
539
  MAX_MESSAGE_ARRAY_LENGTH = 10000
540
 
541
+
542
+ async def _read_request_json(request: Request) -> dict[str, Any]:
543
+ """Read and parse JSON from a request, handling compressed bodies.
544
+
545
+ Clients like OpenAI Codex may send zstd, gzip, or deflate-compressed
546
+ request bodies. Starlette's ``request.json()`` does not decompress
547
+ automatically, causing a UnicodeDecodeError on compressed bytes.
548
+
549
+ This helper inspects ``Content-Encoding``, decompresses if needed,
550
+ then JSON-decodes the result. It raises ``ValueError`` on any
551
+ decompression or parse failure so callers can return a clean 400.
552
+ """
553
+ encoding = (request.headers.get("content-encoding") or "").lower().strip()
554
+ raw = await request.body()
555
+
556
+ if encoding in ("zstd", "zstandard"):
557
+ try:
558
+ import zstandard
559
+
560
+ dctx = zstandard.ZstdDecompressor()
561
+ raw = dctx.decompress(raw)
562
+ except ImportError:
563
+ # Auto-detect: if bytes start with zstd magic (0x28 0xb5 0x2f 0xfd), fail clearly
564
+ raise ValueError(
565
+ "Request body is zstd-compressed but the 'zstandard' package is not installed. "
566
+ "Install it with: pip install zstandard"
567
+ ) from None
568
+ except Exception as exc:
569
+ raise ValueError(f"Failed to decompress zstd request body: {exc}") from exc
570
+ elif encoding == "gzip":
571
+ import gzip as _gzip
572
+
573
+ try:
574
+ raw = _gzip.decompress(raw)
575
+ except Exception as exc:
576
+ raise ValueError(f"Failed to decompress gzip request body: {exc}") from exc
577
+ elif encoding == "deflate":
578
+ import zlib
579
+
580
+ try:
581
+ raw = zlib.decompress(raw)
582
+ except Exception as exc:
583
+ raise ValueError(f"Failed to decompress deflate request body: {exc}") from exc
584
+ elif encoding == "br":
585
+ try:
586
+ import brotli
587
+
588
+ raw = brotli.decompress(raw)
589
+ except ImportError:
590
+ raise ValueError(
591
+ "Request body is brotli-compressed but the 'brotli' package is not installed."
592
+ ) from None
593
+ except Exception as exc:
594
+ raise ValueError(f"Failed to decompress brotli request body: {exc}") from exc
595
+ elif encoding and encoding != "identity":
596
+ raise ValueError(f"Unsupported Content-Encoding: {encoding}")
597
+
598
+ # Decode and parse JSON
599
+ try:
600
+ text = raw.decode("utf-8")
601
+ except UnicodeDecodeError as exc:
602
+ raise ValueError(f"Request body is not valid UTF-8 (possibly compressed?): {exc}") from exc
603
+
604
+ result: dict[str, Any] = json.loads(text)
605
+ return result
606
+
607
+
608
  # Maximum compression cache sessions (prevents unbounded memory growth)
609
  MAX_COMPRESSION_CACHE_SESSIONS = 500
610
 
 
2262
  """
2263
  # Check bypass header
2264
  if request.headers.get("x-headroom-bypass", "").lower() == "true":
2265
+ try:
2266
+ body = await _read_request_json(request)
2267
+ except (json.JSONDecodeError, ValueError) as e:
2268
+ return JSONResponse(
2269
+ status_code=400,
2270
+ content={"error": f"Invalid request body: {e!s}"},
2271
+ )
2272
  messages = body.get("messages", [])
2273
  return JSONResponse(
2274
  {
 
2283
  )
2284
 
2285
  try:
2286
+ body = await _read_request_json(request)
2287
  except Exception:
2288
  return JSONResponse(
2289
  status_code=400,
 
2400
 
2401
  # Parse request
2402
  try:
2403
+ body = await _read_request_json(request)
2404
+ except (json.JSONDecodeError, ValueError) as e:
2405
  return JSONResponse(
2406
  status_code=400,
2407
  content={
2408
  "type": "error",
2409
  "error": {
2410
  "type": "invalid_request_error",
2411
+ "message": f"Invalid request body: {e!s}",
2412
  },
2413
  },
2414
  )
 
3361
 
3362
  # Parse request
3363
  try:
3364
+ body = await _read_request_json(request)
3365
+ except (json.JSONDecodeError, ValueError) as e:
3366
  return JSONResponse(
3367
  status_code=400,
3368
  content={
3369
  "type": "error",
3370
  "error": {
3371
  "type": "invalid_request_error",
3372
+ "message": f"Invalid request body: {e!s}",
3373
  },
3374
  },
3375
  )
 
3799
 
3800
  # Parse request
3801
  try:
3802
+ body = await _read_request_json(request)
3803
+ except (json.JSONDecodeError, ValueError) as e:
3804
  return JSONResponse(
3805
  status_code=400,
3806
  content={
3807
  "error": {
3808
  "code": 400,
3809
+ "message": f"Invalid request body: {e!s}",
3810
  "status": "INVALID_ARGUMENT",
3811
  }
3812
  },
 
5211
 
5212
  # Parse request
5213
  try:
5214
+ body = await _read_request_json(request)
5215
+ except (json.JSONDecodeError, ValueError) as e:
5216
  return JSONResponse(
5217
  status_code=400,
5218
  content={
5219
  "error": {
5220
+ "message": f"Invalid request body: {e!s}",
5221
  "type": "invalid_request_error",
5222
  "code": "invalid_json",
5223
  }
 
5808
  request_id = await self._next_request_id()
5809
 
5810
  try:
5811
+ body = await _read_request_json(request)
5812
  except Exception as e:
5813
  logger.error(f"[{request_id}] Failed to parse Databricks request body: {e}")
5814
  return JSONResponse(
 
5862
 
5863
  # Parse request
5864
  try:
5865
+ body = await _read_request_json(request)
5866
+ except (json.JSONDecodeError, ValueError) as e:
5867
  return JSONResponse(
5868
  status_code=400,
5869
  content={
5870
  "error": {
5871
+ "message": f"Invalid request body: {e!s}",
5872
  "type": "invalid_request_error",
5873
  "code": "invalid_json",
5874
  }
 
6345
 
6346
  # Parse request
6347
  try:
6348
+ body = await _read_request_json(request)
6349
+ except (json.JSONDecodeError, ValueError) as e:
6350
  return JSONResponse(
6351
  status_code=400,
6352
  content={
6353
  "error": {
6354
+ "message": f"Invalid request body: {e!s}",
6355
  "type": "invalid_request_error",
6356
  "code": "invalid_json",
6357
  }
 
6507
 
6508
  # Parse request
6509
  try:
6510
+ body = await _read_request_json(request)
6511
+ except (json.JSONDecodeError, ValueError) as e:
6512
  return JSONResponse(
6513
  status_code=400,
6514
  content={
6515
  "error": {
6516
+ "message": f"Invalid request body: {e!s}",
6517
  "code": 400,
6518
  }
6519
  },
 
6771
 
6772
  # Parse request
6773
  try:
6774
+ body = await _read_request_json(request)
6775
+ except (json.JSONDecodeError, ValueError) as e:
6776
  return JSONResponse(
6777
  status_code=400,
6778
  content={
6779
  "error": {
6780
+ "message": f"Invalid request body: {e!s}",
6781
  "code": 400,
6782
  }
6783
  },
 
6840
 
6841
  # Parse request
6842
  try:
6843
+ body = await _read_request_json(request)
6844
+ except (json.JSONDecodeError, ValueError) as e:
6845
  return JSONResponse(
6846
  status_code=400,
6847
  content={
6848
  "error": {
6849
+ "message": f"Invalid request body: {e!s}",
6850
  "code": 400,
6851
  }
6852
  },
headroom/telemetry/beacon.py CHANGED
@@ -208,8 +208,8 @@ class TelemetryBeacon:
208
  try:
209
  tokens = stats.get("tokens", {})
210
  total_req = stats.get("requests", {}).get("total", 1)
211
- tokens_before = tokens.get("original", 0)
212
- tokens_after = tokens.get("optimized", 0)
213
  payload.update(
214
  {
215
  "avg_tokens_before": round(tokens_before / max(total_req, 1)),
@@ -226,7 +226,7 @@ class TelemetryBeacon:
226
  payload["compression_cache"] = {
227
  "hit_rate": cc.get("hit_rate", 0),
228
  "entries": cc.get("entries", 0),
229
- "avg_lookup_ns": cc.get("avg_lookup_ns", 0),
230
  }
231
  except Exception:
232
  logger.debug("Beacon: failed to extract cache stats", exc_info=True)
 
208
  try:
209
  tokens = stats.get("tokens", {})
210
  total_req = stats.get("requests", {}).get("total", 1)
211
+ tokens_before = tokens.get("total_before_compression", 0)
212
+ tokens_after = tokens_before - tokens.get("saved", 0)
213
  payload.update(
214
  {
215
  "avg_tokens_before": round(tokens_before / max(total_req, 1)),
 
226
  payload["compression_cache"] = {
227
  "hit_rate": cc.get("hit_rate", 0),
228
  "entries": cc.get("entries", 0),
229
+ "tokens_saved": cc.get("total_tokens_saved", 0),
230
  }
231
  except Exception:
232
  logger.debug("Beacon: failed to extract cache stats", exc_info=True)
pyproject.toml CHANGED
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
 
5
  [project]
6
  name = "headroom-ai"
7
- version = "0.5.12"
8
  description = "The Context Optimization Layer for LLM Applications - Cut costs by 50-90%"
9
  readme = "README.md"
10
  license = "Apache-2.0"
 
4
 
5
  [project]
6
  name = "headroom-ai"
7
+ version = "0.5.13"
8
  description = "The Context Optimization Layer for LLM Applications - Cut costs by 50-90%"
9
  readme = "README.md"
10
  license = "Apache-2.0"
scripts/dedupe_telemetry.py ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Deduplicate proxy_telemetry_v2 rows by session_id.
3
+
4
+ Keeps the latest entry (by created_at) per session and deletes older duplicates.
5
+ Run manually or via cron:
6
+
7
+ # One-off
8
+ python scripts/dedupe_telemetry.py
9
+
10
+ # Hourly via crontab -e
11
+ 0 * * * * cd /Users/tchopra/claude-projects/headroom && python scripts/dedupe_telemetry.py >> /tmp/dedupe_telemetry.log 2>&1
12
+
13
+ # Dry run (show what would be deleted)
14
+ python scripts/dedupe_telemetry.py --dry-run
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import argparse
20
+ from collections import defaultdict
21
+ from datetime import datetime
22
+
23
+ import requests
24
+
25
+ # Supabase config (same anon key as beacon — INSERT/SELECT/DELETE via RLS)
26
+ _SUPABASE_URL = "https://dtlllcsudcoasebbamcq.supabase.co"
27
+ _SUPABASE_KEY = ".".join(
28
+ [
29
+ "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9",
30
+ "eyJpc3MiOiJzdXBhYmFzZSIsInJlZiI6ImR0bGxsY3N1ZGNvYXNlYmJhbWNxIiwicm9sZSI6ImFub24iLCJpYXQiOjE3NzM3MDc4NDUsImV4cCI6MjA4OTI4Mzg0NX0",
31
+ "h_C6dLQKa8BVc3upgEvulR4E0K4eiEViyddRMIylKjU",
32
+ ]
33
+ )
34
+ _TABLE = "proxy_telemetry_v2"
35
+ _ENDPOINT = f"{_SUPABASE_URL}/rest/v1/{_TABLE}"
36
+ _HEADERS = {
37
+ "apikey": _SUPABASE_KEY,
38
+ "Authorization": f"Bearer {_SUPABASE_KEY}",
39
+ }
40
+
41
+
42
+ def fetch_all_rows() -> list[dict]:
43
+ """Fetch id, session_id, created_at for all rows."""
44
+ rows = []
45
+ offset = 0
46
+ limit = 1000
47
+ while True:
48
+ resp = requests.get(
49
+ _ENDPOINT,
50
+ params={
51
+ "select": "id,session_id,created_at",
52
+ "order": "created_at.desc",
53
+ "offset": offset,
54
+ "limit": limit,
55
+ },
56
+ headers=_HEADERS,
57
+ )
58
+ resp.raise_for_status()
59
+ batch = resp.json()
60
+ if not batch:
61
+ break
62
+ rows.extend(batch)
63
+ if len(batch) < limit:
64
+ break
65
+ offset += limit
66
+ return rows
67
+
68
+
69
+ def find_duplicates(rows: list[dict]) -> list[str]:
70
+ """Find row IDs to delete (all but the latest per session_id)."""
71
+ by_session: dict[str, list[dict]] = defaultdict(list)
72
+ for row in rows:
73
+ by_session[row["session_id"]].append(row)
74
+
75
+ ids_to_delete = []
76
+ for _session_id, entries in by_session.items():
77
+ if len(entries) <= 1:
78
+ continue
79
+ # Sort by created_at descending — keep first (latest), delete rest
80
+ entries.sort(key=lambda r: r["created_at"], reverse=True)
81
+ for entry in entries[1:]:
82
+ ids_to_delete.append(entry["id"])
83
+
84
+ return ids_to_delete
85
+
86
+
87
+ def delete_rows(ids: list[str], dry_run: bool = False) -> int:
88
+ """Delete rows by ID. Returns count deleted."""
89
+ if dry_run or not ids:
90
+ return 0
91
+
92
+ deleted = 0
93
+ # PostgREST supports filtering by id in a list
94
+ batch_size = 50
95
+ for i in range(0, len(ids), batch_size):
96
+ batch = ids[i : i + batch_size]
97
+ id_filter = ",".join(batch)
98
+ resp = requests.delete(
99
+ _ENDPOINT,
100
+ params={"id": f"in.({id_filter})"},
101
+ headers={**_HEADERS, "Prefer": "return=minimal"},
102
+ )
103
+ resp.raise_for_status()
104
+ deleted += len(batch)
105
+
106
+ return deleted
107
+
108
+
109
+ def main() -> None:
110
+ parser = argparse.ArgumentParser(description="Deduplicate proxy_telemetry_v2 by session_id")
111
+ parser.add_argument(
112
+ "--dry-run",
113
+ action="store_true",
114
+ help="Show what would be deleted without deleting",
115
+ )
116
+ args = parser.parse_args()
117
+
118
+ now = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
119
+ print(f"[{now}] Fetching telemetry rows...")
120
+
121
+ rows = fetch_all_rows()
122
+ print(f" Total rows: {len(rows)}")
123
+
124
+ # Group stats
125
+ by_session: dict[str, int] = defaultdict(int)
126
+ for row in rows:
127
+ by_session[row["session_id"]] += 1
128
+ dup_sessions = {k: v for k, v in by_session.items() if v > 1}
129
+
130
+ if not dup_sessions:
131
+ print(" No duplicates found. Nothing to do.")
132
+ return
133
+
134
+ print(f" Sessions with duplicates: {len(dup_sessions)}")
135
+ total_dups = sum(v - 1 for v in dup_sessions.values())
136
+ print(f" Rows to delete: {total_dups}")
137
+
138
+ ids_to_delete = find_duplicates(rows)
139
+
140
+ if args.dry_run:
141
+ print("\n [DRY RUN] Would delete these rows:")
142
+ for row in rows:
143
+ if row["id"] in ids_to_delete:
144
+ print(
145
+ f" {row['id']} session={row['session_id'][:12]}... created={row['created_at']}"
146
+ )
147
+ print(f"\n [DRY RUN] Would keep {len(rows) - len(ids_to_delete)} rows")
148
+ return
149
+
150
+ deleted = delete_rows(ids_to_delete)
151
+ print(f" Deleted: {deleted} duplicate rows")
152
+ print(f" Remaining: {len(rows) - deleted} rows (1 per session)")
153
+
154
+
155
+ if __name__ == "__main__":
156
+ main()
tests/test_backend_bugs.py CHANGED
@@ -1,9 +1,11 @@
1
  """Tests for backend bug fixes in LiteLLM and any-llm integrations.
2
 
3
  Tests tool forwarding, tool argument parsing, streaming param forwarding,
4
- and Vertex AI model mapping.
 
5
  """
6
 
 
7
  from unittest.mock import AsyncMock, MagicMock, patch
8
 
9
  import pytest
@@ -196,6 +198,336 @@ class TestLiteLLMToolsForwarding:
196
  assert isinstance(tool_block["input"], dict)
197
 
198
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
199
  # =============================================================================
200
  # Streaming Params (Bugs 3-4)
201
  # =============================================================================
 
1
  """Tests for backend bug fixes in LiteLLM and any-llm integrations.
2
 
3
  Tests tool forwarding, tool argument parsing, streaming param forwarding,
4
+ message conversion (tool_use/tool_result), streaming tool_calls, and
5
+ Vertex AI model mapping.
6
  """
7
 
8
+ import json
9
  from unittest.mock import AsyncMock, MagicMock, patch
10
 
11
  import pytest
 
198
  assert isinstance(tool_block["input"], dict)
199
 
200
 
201
+ # =============================================================================
202
+ # Message Conversion: tool_use / tool_result (GitHub Issue — Bug 2)
203
+ # =============================================================================
204
+
205
+
206
+ class TestConvertMessagesToolBlocks:
207
+ """Test that _convert_messages_for_litellm converts Anthropic tool blocks to OpenAI format."""
208
+
209
+ def _make_backend(self):
210
+ with patch("headroom.backends.litellm._fetch_bedrock_inference_profiles", return_value={}):
211
+ return LiteLLMBackend(provider="openrouter")
212
+
213
+ def test_tool_result_converted_to_tool_role(self):
214
+ """Anthropic tool_result blocks must become role=tool messages."""
215
+ backend = self._make_backend()
216
+ messages = [
217
+ {"role": "user", "content": "Weather in Paris?"},
218
+ {
219
+ "role": "assistant",
220
+ "content": [
221
+ {
222
+ "type": "tool_use",
223
+ "id": "toolu_01",
224
+ "name": "get_weather",
225
+ "input": {"city": "Paris"},
226
+ },
227
+ ],
228
+ },
229
+ {
230
+ "role": "user",
231
+ "content": [
232
+ {"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny, 22C"},
233
+ ],
234
+ },
235
+ ]
236
+ converted = backend._convert_messages_for_litellm(messages)
237
+
238
+ # assistant message should have tool_calls
239
+ assistant = converted[1]
240
+ assert assistant["role"] == "assistant"
241
+ assert "tool_calls" in assistant
242
+ assert assistant["tool_calls"][0]["id"] == "toolu_01"
243
+ assert assistant["tool_calls"][0]["type"] == "function"
244
+ assert assistant["tool_calls"][0]["function"]["name"] == "get_weather"
245
+ assert json.loads(assistant["tool_calls"][0]["function"]["arguments"]) == {"city": "Paris"}
246
+
247
+ # tool_result should become role=tool
248
+ tool_msg = converted[2]
249
+ assert tool_msg["role"] == "tool"
250
+ assert tool_msg["tool_call_id"] == "toolu_01"
251
+ assert tool_msg["content"] == "Sunny, 22C"
252
+
253
+ def test_tool_result_with_list_content(self):
254
+ """tool_result with list content should be flattened to string."""
255
+ backend = self._make_backend()
256
+ messages = [
257
+ {
258
+ "role": "user",
259
+ "content": [
260
+ {
261
+ "type": "tool_result",
262
+ "tool_use_id": "toolu_02",
263
+ "content": [
264
+ {"type": "text", "text": "Line 1"},
265
+ {"type": "text", "text": "Line 2"},
266
+ ],
267
+ },
268
+ ],
269
+ },
270
+ ]
271
+ converted = backend._convert_messages_for_litellm(messages)
272
+ assert converted[0]["role"] == "tool"
273
+ assert converted[0]["content"] == "Line 1\nLine 2"
274
+
275
+ def test_assistant_tool_use_with_text(self):
276
+ """Assistant message with both text and tool_use blocks."""
277
+ backend = self._make_backend()
278
+ messages = [
279
+ {
280
+ "role": "assistant",
281
+ "content": [
282
+ {"type": "text", "text": "Let me check the weather."},
283
+ {
284
+ "type": "tool_use",
285
+ "id": "toolu_03",
286
+ "name": "get_weather",
287
+ "input": {"city": "Tokyo"},
288
+ },
289
+ ],
290
+ },
291
+ ]
292
+ converted = backend._convert_messages_for_litellm(messages)
293
+ assert len(converted) == 1
294
+ assert converted[0]["role"] == "assistant"
295
+ assert converted[0]["content"] == "Let me check the weather."
296
+ assert converted[0]["tool_calls"][0]["function"]["name"] == "get_weather"
297
+
298
+ def test_simple_text_messages_unchanged(self):
299
+ """Plain string messages pass through."""
300
+ backend = self._make_backend()
301
+ messages = [
302
+ {"role": "user", "content": "Hello"},
303
+ {"role": "assistant", "content": "Hi!"},
304
+ ]
305
+ converted = backend._convert_messages_for_litellm(messages)
306
+ assert converted == messages
307
+
308
+ def test_multiple_tool_results(self):
309
+ """Multiple tool_result blocks in one user message → multiple role=tool messages."""
310
+ backend = self._make_backend()
311
+ messages = [
312
+ {
313
+ "role": "user",
314
+ "content": [
315
+ {"type": "tool_result", "tool_use_id": "toolu_a", "content": "Result A"},
316
+ {"type": "tool_result", "tool_use_id": "toolu_b", "content": "Result B"},
317
+ ],
318
+ },
319
+ ]
320
+ converted = backend._convert_messages_for_litellm(messages)
321
+ assert len(converted) == 2
322
+ assert converted[0]["role"] == "tool"
323
+ assert converted[0]["tool_call_id"] == "toolu_a"
324
+ assert converted[1]["role"] == "tool"
325
+ assert converted[1]["tool_call_id"] == "toolu_b"
326
+
327
+
328
+ # =============================================================================
329
+ # Streaming tool_calls (GitHub Issue — Bug 1)
330
+ # =============================================================================
331
+
332
+
333
+ class TestStreamMessageToolCalls:
334
+ """Test that stream_message emits tool_use blocks and correct stop_reason."""
335
+
336
+ @pytest.mark.asyncio
337
+ async def test_stream_emits_tool_use_blocks(self):
338
+ """Tool calls in streaming should produce content_block_start with type=tool_use."""
339
+
340
+ async def mock_stream():
341
+ # First chunk: tool call start (id + name)
342
+ tc = MagicMock()
343
+ tc.index = 0
344
+ tc.id = "toolu_stream_01"
345
+ tc.function = MagicMock()
346
+ tc.function.name = "get_weather"
347
+ tc.function.arguments = ""
348
+
349
+ chunk1 = MagicMock()
350
+ chunk1.choices = [
351
+ MagicMock(delta=MagicMock(content=None, tool_calls=[tc]), finish_reason=None)
352
+ ]
353
+ yield chunk1
354
+
355
+ # Second chunk: arguments delta
356
+ tc2 = MagicMock()
357
+ tc2.index = 0
358
+ tc2.id = None
359
+ tc2.function = MagicMock()
360
+ tc2.function.name = None
361
+ tc2.function.arguments = '{"city":"Paris"}'
362
+
363
+ chunk2 = MagicMock()
364
+ chunk2.choices = [
365
+ MagicMock(delta=MagicMock(content=None, tool_calls=[tc2]), finish_reason=None)
366
+ ]
367
+ yield chunk2
368
+
369
+ # Final chunk: finish_reason=tool_calls
370
+ chunk3 = MagicMock()
371
+ chunk3.choices = [
372
+ MagicMock(
373
+ delta=MagicMock(content=None, tool_calls=None), finish_reason="tool_calls"
374
+ )
375
+ ]
376
+ yield chunk3
377
+
378
+ with (
379
+ patch("headroom.backends.litellm.acompletion", new_callable=AsyncMock) as mock_acomp,
380
+ patch("headroom.backends.litellm._fetch_bedrock_inference_profiles", return_value={}),
381
+ ):
382
+ mock_acomp.return_value = mock_stream()
383
+ backend = LiteLLMBackend(provider="openrouter")
384
+
385
+ events = []
386
+ async for event in backend.stream_message(
387
+ {
388
+ "model": "test",
389
+ "messages": [{"role": "user", "content": "weather?"}],
390
+ "tools": [
391
+ {
392
+ "name": "get_weather",
393
+ "description": "Get weather",
394
+ "input_schema": {"type": "object"},
395
+ }
396
+ ],
397
+ },
398
+ {},
399
+ ):
400
+ events.append(event)
401
+
402
+ # Find content_block_start events
403
+ block_starts = [e for e in events if e.event_type == "content_block_start"]
404
+ assert len(block_starts) == 1
405
+ assert block_starts[0].data["content_block"]["type"] == "tool_use"
406
+ assert block_starts[0].data["content_block"]["id"] == "toolu_stream_01"
407
+ assert block_starts[0].data["content_block"]["name"] == "get_weather"
408
+
409
+ # Find input_json_delta events
410
+ json_deltas = [
411
+ e
412
+ for e in events
413
+ if e.event_type == "content_block_delta"
414
+ and e.data.get("delta", {}).get("type") == "input_json_delta"
415
+ ]
416
+ assert len(json_deltas) == 1
417
+ assert json_deltas[0].data["delta"]["partial_json"] == '{"city":"Paris"}'
418
+
419
+ # Check stop_reason is "tool_use"
420
+ msg_delta = [e for e in events if e.event_type == "message_delta"]
421
+ assert len(msg_delta) == 1
422
+ assert msg_delta[0].data["delta"]["stop_reason"] == "tool_use"
423
+
424
+ @pytest.mark.asyncio
425
+ async def test_stream_text_still_works(self):
426
+ """Pure text streaming should still work correctly."""
427
+
428
+ async def mock_stream():
429
+ chunk = MagicMock()
430
+ chunk.choices = [
431
+ MagicMock(delta=MagicMock(content="Hello!", tool_calls=None), finish_reason=None)
432
+ ]
433
+ yield chunk
434
+
435
+ chunk2 = MagicMock()
436
+ chunk2.choices = [
437
+ MagicMock(delta=MagicMock(content=None, tool_calls=None), finish_reason="stop")
438
+ ]
439
+ yield chunk2
440
+
441
+ with (
442
+ patch("headroom.backends.litellm.acompletion", new_callable=AsyncMock) as mock_acomp,
443
+ patch("headroom.backends.litellm._fetch_bedrock_inference_profiles", return_value={}),
444
+ ):
445
+ mock_acomp.return_value = mock_stream()
446
+ backend = LiteLLMBackend(provider="openrouter")
447
+
448
+ events = []
449
+ async for event in backend.stream_message(
450
+ {"model": "test", "messages": [{"role": "user", "content": "hi"}]},
451
+ {},
452
+ ):
453
+ events.append(event)
454
+
455
+ block_starts = [e for e in events if e.event_type == "content_block_start"]
456
+ assert len(block_starts) == 1
457
+ assert block_starts[0].data["content_block"]["type"] == "text"
458
+
459
+ text_deltas = [e for e in events if e.event_type == "content_block_delta"]
460
+ assert len(text_deltas) == 1
461
+ assert text_deltas[0].data["delta"]["text"] == "Hello!"
462
+
463
+ msg_delta = [e for e in events if e.event_type == "message_delta"]
464
+ assert msg_delta[0].data["delta"]["stop_reason"] == "end_turn"
465
+
466
+ @pytest.mark.asyncio
467
+ async def test_stream_text_then_tool(self):
468
+ """Text followed by tool call should produce two blocks."""
469
+
470
+ async def mock_stream():
471
+ # Text chunk
472
+ chunk1 = MagicMock()
473
+ chunk1.choices = [
474
+ MagicMock(
475
+ delta=MagicMock(content="I'll check. ", tool_calls=None), finish_reason=None
476
+ )
477
+ ]
478
+ yield chunk1
479
+
480
+ # Tool call chunk
481
+ tc = MagicMock()
482
+ tc.index = 0
483
+ tc.id = "toolu_mixed"
484
+ tc.function = MagicMock()
485
+ tc.function.name = "search"
486
+ tc.function.arguments = '{"q":"test"}'
487
+
488
+ chunk2 = MagicMock()
489
+ chunk2.choices = [
490
+ MagicMock(delta=MagicMock(content=None, tool_calls=[tc]), finish_reason=None)
491
+ ]
492
+ yield chunk2
493
+
494
+ # Finish
495
+ chunk3 = MagicMock()
496
+ chunk3.choices = [
497
+ MagicMock(
498
+ delta=MagicMock(content=None, tool_calls=None), finish_reason="tool_calls"
499
+ )
500
+ ]
501
+ yield chunk3
502
+
503
+ with (
504
+ patch("headroom.backends.litellm.acompletion", new_callable=AsyncMock) as mock_acomp,
505
+ patch("headroom.backends.litellm._fetch_bedrock_inference_profiles", return_value={}),
506
+ ):
507
+ mock_acomp.return_value = mock_stream()
508
+ backend = LiteLLMBackend(provider="openrouter")
509
+
510
+ events = []
511
+ async for event in backend.stream_message(
512
+ {"model": "test", "messages": [{"role": "user", "content": "hi"}]},
513
+ {},
514
+ ):
515
+ events.append(event)
516
+
517
+ block_starts = [e for e in events if e.event_type == "content_block_start"]
518
+ assert len(block_starts) == 2
519
+ assert block_starts[0].data["content_block"]["type"] == "text"
520
+ assert block_starts[1].data["content_block"]["type"] == "tool_use"
521
+
522
+ # Two content_block_stop events (one per block)
523
+ block_stops = [e for e in events if e.event_type == "content_block_stop"]
524
+ assert len(block_stops) == 2
525
+
526
+ # stop_reason should be tool_use
527
+ msg_delta = [e for e in events if e.event_type == "message_delta"]
528
+ assert msg_delta[0].data["delta"]["stop_reason"] == "tool_use"
529
+
530
+
531
  # =============================================================================
532
  # Streaming Params (Bugs 3-4)
533
  # =============================================================================