Jack1808/Claude_Code
0
1"""SSE event builder for Anthropic-format streaming responses."""2 3import json4from collections.abc import Iterator5from dataclasses import dataclass, field6from typing import Any7 8from loguru import logger9 10try:11 import tiktoken12 13 ENCODER = tiktoken.get_encoding("cl100k_base")14except Exception:15 ENCODER = None16 17 18# Map OpenAI finish_reason to Anthropic stop_reason19STOP_REASON_MAP = {20 "stop": "end_turn",21 "length": "max_tokens",22 "tool_calls": "tool_use",23 "content_filter": "end_turn",24}25 26 27def map_stop_reason(openai_reason: str | None) -> str:28 """Map OpenAI finish_reason to Anthropic stop_reason."""29 return (30 STOP_REASON_MAP.get(openai_reason, "end_turn") if openai_reason else "end_turn"31 )32 33 34@dataclass35class ToolCallState:36 """State for a single streaming tool call."""37 38 block_index: int # -1 if not yet allocated39 tool_id: str40 name: str41 contents: list[str] = field(default_factory=list)42 started: bool = False43 task_arg_buffer: str = ""44 task_args_emitted: bool = False45 46 47@dataclass48class ContentBlockManager:49 """Manages content block indices and state."""50 51 next_index: int = 052 thinking_index: int = -153 text_index: int = -154 thinking_started: bool = False55 text_started: bool = False56 tool_states: dict[int, ToolCallState] = field(default_factory=dict)57 58 def allocate_index(self) -> int:59 """Allocate and return the next block index."""60 idx = self.next_index61 self.next_index += 162 return idx63 64 def register_tool_name(self, index: int, name: str) -> None:65 """Register or merge a streaming tool name fragment.66 67 Handles providers that stream names as fragments and those that68 resend the full name on every chunk.69 """70 if index not in self.tool_states:71 self.tool_states[index] = ToolCallState(72 block_index=-1, tool_id="", name=name73 )74 return75 state = self.tool_states[index]76 prev = state.name77 if not prev or name.startswith(prev):78 state.name = name79 elif not prev.startswith(name):80 state.name = prev + name81 82 def buffer_task_args(self, index: int, args: str) -> dict | None:83 """Buffer Task tool args and return parsed JSON when complete.84 85 Returns the parsed (and patched) args dict once the buffer forms86 valid JSON, or None if still accumulating.87 """88 state = self.tool_states.get(index)89 if state is None or state.task_args_emitted:90 return None91 92 state.task_arg_buffer += args93 try:94 args_json = json.loads(state.task_arg_buffer)95 except Exception:96 return None97 98 if args_json.get("run_in_background") is not False:99 args_json["run_in_background"] = False100 101 state.task_args_emitted = True102 state.task_arg_buffer = ""103 return args_json104 105 def flush_task_arg_buffers(self) -> list[tuple[int, str]]:106 """Flush any remaining Task arg buffers. Returns (tool_index, json_str) pairs."""107 results: list[tuple[int, str]] = []108 for tool_index, state in list(self.tool_states.items()):109 if not state.task_arg_buffer or state.task_args_emitted:110 continue111 112 out = "{}"113 try:114 args_json = json.loads(state.task_arg_buffer)115 if args_json.get("run_in_background") is not False:116 args_json["run_in_background"] = False117 out = json.dumps(args_json)118 except Exception as e:119 prefix = state.task_arg_buffer[:120]120 logger.warning(121 "Task args invalid JSON (id={} len={} prefix={!r}): {}",122 state.tool_id or "unknown",123 len(state.task_arg_buffer),124 prefix,125 e,126 )127 128 state.task_args_emitted = True129 state.task_arg_buffer = ""130 results.append((tool_index, out))131 return results132 133 134class SSEBuilder:135 """Builder for Anthropic SSE streaming events."""136 137 def __init__(self, message_id: str, model: str, input_tokens: int = 0):138 self.message_id = message_id139 self.model = model140 self.input_tokens = input_tokens141 self.blocks = ContentBlockManager()142 self._accumulated_text_parts: list[str] = []143 self._accumulated_reasoning_parts: list[str] = []144 145 def _format_event(self, event_type: str, data: dict[str, Any]) -> str:146 """Format as SSE string."""147 event_str = f"event: {event_type}\ndata: {json.dumps(data)}\n\n"148 logger.debug("SSE_EVENT: {} - {}", event_type, event_str.strip())149 return event_str150 151 # Message lifecycle events152 def message_start(self) -> str:153 """Generate message_start event."""154 usage = {"input_tokens": self.input_tokens, "output_tokens": 1}155 return self._format_event(156 "message_start",157 {158 "type": "message_start",159 "message": {160 "id": self.message_id,161 "type": "message",162 "role": "assistant",163 "content": [],164 "model": self.model,165 "stop_reason": None,166 "stop_sequence": None,167 "usage": usage,168 },169 },170 )171 172 def message_delta(self, stop_reason: str, output_tokens: int) -> str:173 """Generate message_delta event with stop reason."""174 return self._format_event(175 "message_delta",176 {177 "type": "message_delta",178 "delta": {"stop_reason": stop_reason, "stop_sequence": None},179 "usage": {180 "input_tokens": self.input_tokens,181 "output_tokens": output_tokens,182 },183 },184 )185 186 def message_stop(self) -> str:187 """Generate message_stop event."""188 return self._format_event("message_stop", {"type": "message_stop"})189 190 # Content block events191 def content_block_start(self, index: int, block_type: str, **kwargs) -> str:192 """Generate content_block_start event."""193 content_block: dict[str, Any] = {"type": block_type}194 if block_type == "thinking":195 content_block["thinking"] = kwargs.get("thinking", "")196 elif block_type == "text":197 content_block["text"] = kwargs.get("text", "")198 elif block_type == "tool_use":199 content_block["id"] = kwargs.get("id", "")200 content_block["name"] = kwargs.get("name", "")201 content_block["input"] = kwargs.get("input", {})202 203 return self._format_event(204 "content_block_start",205 {206 "type": "content_block_start",207 "index": index,208 "content_block": content_block,209 },210 )211 212 def content_block_delta(self, index: int, delta_type: str, content: str) -> str:213 """Generate content_block_delta event."""214 delta: dict[str, Any] = {"type": delta_type}215 if delta_type == "thinking_delta":216 delta["thinking"] = content217 elif delta_type == "text_delta":218 delta["text"] = content219 elif delta_type == "input_json_delta":220 delta["partial_json"] = content221 222 return self._format_event(223 "content_block_delta",224 {225 "type": "content_block_delta",226 "index": index,227 "delta": delta,228 },229 )230 231 def content_block_stop(self, index: int) -> str:232 """Generate content_block_stop event."""233 return self._format_event(234 "content_block_stop",235 {236 "type": "content_block_stop",237 "index": index,238 },239 )240 241 # High-level helpers for thinking blocks242 def start_thinking_block(self) -> str:243 """Start a thinking block, allocating index."""244 self.blocks.thinking_index = self.blocks.allocate_index()245 self.blocks.thinking_started = True246 return self.content_block_start(self.blocks.thinking_index, "thinking")247 248 def emit_thinking_delta(self, content: str) -> str:249 """Emit thinking content delta."""250 self._accumulated_reasoning_parts.append(content)251 return self.content_block_delta(252 self.blocks.thinking_index, "thinking_delta", content253 )254 255 def stop_thinking_block(self) -> str:256 """Stop the current thinking block."""257 self.blocks.thinking_started = False258 return self.content_block_stop(self.blocks.thinking_index)259 260 # High-level helpers for text blocks261 def start_text_block(self) -> str:262 """Start a text block, allocating index."""263 self.blocks.text_index = self.blocks.allocate_index()264 self.blocks.text_started = True265 return self.content_block_start(self.blocks.text_index, "text")266 267 def emit_text_delta(self, content: str) -> str:268 """Emit text content delta."""269 self._accumulated_text_parts.append(content)270 return self.content_block_delta(self.blocks.text_index, "text_delta", content)271 272 def stop_text_block(self) -> str:273 """Stop the current text block."""274 self.blocks.text_started = False275 return self.content_block_stop(self.blocks.text_index)276 277 # High-level helpers for tool blocks278 def start_tool_block(self, tool_index: int, tool_id: str, name: str) -> str:279 """Start a tool_use block."""280 block_idx = self.blocks.allocate_index()281 if tool_index in self.blocks.tool_states:282 state = self.blocks.tool_states[tool_index]283 state.block_index = block_idx284 state.tool_id = tool_id285 state.started = True286 else:287 self.blocks.tool_states[tool_index] = ToolCallState(288 block_index=block_idx,289 tool_id=tool_id,290 name=name,291 started=True,292 )293 return self.content_block_start(block_idx, "tool_use", id=tool_id, name=name)294 295 def emit_tool_delta(self, tool_index: int, partial_json: str) -> str:296 """Emit tool input delta."""297 state = self.blocks.tool_states[tool_index]298 state.contents.append(partial_json)299 return self.content_block_delta(300 state.block_index, "input_json_delta", partial_json301 )302 303 def stop_tool_block(self, tool_index: int) -> str:304 """Stop a tool block."""305 block_idx = self.blocks.tool_states[tool_index].block_index306 return self.content_block_stop(block_idx)307 308 # State management helpers309 def ensure_thinking_block(self) -> Iterator[str]:310 """Ensure a thinking block is started, closing text block if needed."""311 if self.blocks.text_started:312 yield self.stop_text_block()313 if not self.blocks.thinking_started:314 yield self.start_thinking_block()315 316 def ensure_text_block(self) -> Iterator[str]:317 """Ensure a text block is started, closing thinking block if needed."""318 if self.blocks.thinking_started:319 yield self.stop_thinking_block()320 if not self.blocks.text_started:321 yield self.start_text_block()322 323 def close_content_blocks(self) -> Iterator[str]:324 """Close thinking and text blocks (before tool calls)."""325 if self.blocks.thinking_started:326 yield self.stop_thinking_block()327 if self.blocks.text_started:328 yield self.stop_text_block()329 330 def close_all_blocks(self) -> Iterator[str]:331 """Close all open blocks (thinking, text, tools)."""332 if self.blocks.thinking_started:333 yield self.stop_thinking_block()334 if self.blocks.text_started:335 yield self.stop_text_block()336 for tool_index, state in list(self.blocks.tool_states.items()):337 if state.started:338 yield self.stop_tool_block(tool_index)339 340 # Error handling341 def emit_error(self, error_message: str) -> Iterator[str]:342 """Emit an error as a text block."""343 error_index = self.blocks.allocate_index()344 yield self.content_block_start(error_index, "text")345 yield self.content_block_delta(error_index, "text_delta", error_message)346 yield self.content_block_stop(error_index)347 348 # Accumulated content access349 @property350 def accumulated_text(self) -> str:351 """Get accumulated text content."""352 return "".join(self._accumulated_text_parts)353 354 @property355 def accumulated_reasoning(self) -> str:356 """Get accumulated reasoning content."""357 return "".join(self._accumulated_reasoning_parts)358 359 def estimate_output_tokens(self) -> int:360 """Estimate output tokens from accumulated content."""361 accumulated_text = self.accumulated_text362 accumulated_reasoning = self.accumulated_reasoning363 if ENCODER:364 text_tokens = len(ENCODER.encode(accumulated_text))365 reasoning_tokens = len(ENCODER.encode(accumulated_reasoning))366 # Tool calls are harder to tokenize exactly without reconstruction, but we can approximate367 # by tokenizing the json dumps of tool contents368 tool_tokens = 0369 started_tool_count = 0370 for state in self.blocks.tool_states.values():371 tool_tokens += len(ENCODER.encode(state.name))372 tool_tokens += len(ENCODER.encode("".join(state.contents)))373 tool_tokens += 15 # Control tokens overhead per tool374 if state.started:375 started_tool_count += 1376 377 # Per-block overhead (~4 tokens per content block)378 block_count = (379 (1 if accumulated_reasoning else 0)380 + (1 if accumulated_text else 0)381 + started_tool_count382 )383 block_overhead = block_count * 4384 385 return text_tokens + reasoning_tokens + tool_tokens + block_overhead386 387 text_tokens = len(accumulated_text) // 4388 reasoning_tokens = len(accumulated_reasoning) // 4389 tool_tokens = sum(1 for s in self.blocks.tool_states.values() if s.started) * 50390 return text_tokens + reasoning_tokens + tool_tokens391 