CoolFace
Apppublic

Jack1808/Claude_Code

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
sse_builder.py391 linesDownload Raw Back to common
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