Spaces:
Build error
Build error
File size: 6,087 Bytes
9c7d451 bf779b5 9c7d451 bf779b5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 | """Shared utilities for Headroom SDK."""
from __future__ import annotations
import hashlib
import json
import re
import uuid
from datetime import datetime
from typing import Any
# Marker format for Headroom modifications
MARKER_PREFIX = "<headroom:"
MARKER_SUFFIX = ">"
def generate_request_id() -> str:
"""Generate a unique request ID."""
return str(uuid.uuid4())
def compute_hash(data: str | bytes) -> str:
"""Compute SHA256 hash, returning hex string."""
if isinstance(data, str):
data = data.encode("utf-8")
return hashlib.sha256(data).hexdigest()
def compute_short_hash(data: str | bytes, length: int = 16) -> str:
"""Compute truncated SHA256 hash."""
return compute_hash(data)[:length]
def compute_messages_hash(messages: list[dict[str, Any]]) -> str:
"""Compute hash of messages list for deduplication."""
# Serialize deterministically
serialized = json.dumps(messages, sort_keys=True, separators=(",", ":"))
return compute_short_hash(serialized)
def compute_prefix_hash(messages: list[dict[str, Any]], prefix_count: int | None = None) -> str:
"""
Compute hash of message prefix for cache alignment.
Args:
messages: List of messages.
prefix_count: Number of messages to include (default: all system messages + 1).
Returns:
Hash of the prefix content.
"""
if not messages:
return compute_short_hash("")
if prefix_count is None:
# Default: system messages + first non-system
prefix_count = 1
for i, msg in enumerate(messages):
if msg.get("role") == "system":
prefix_count = i + 2
else:
break
prefix_messages = messages[:prefix_count]
serialized = json.dumps(prefix_messages, sort_keys=True, separators=(",", ":"))
return compute_short_hash(serialized)
def format_timestamp(dt: datetime | None = None) -> str:
"""Format datetime as ISO8601 string."""
if dt is None:
dt = datetime.utcnow()
return dt.isoformat() + "Z"
def parse_timestamp(ts: str) -> datetime:
"""Parse ISO8601 timestamp string."""
# Handle both with and without Z suffix
ts = ts.rstrip("Z")
return datetime.fromisoformat(ts)
def create_marker(marker_type: str, **kwargs: Any) -> str:
"""
Create a Headroom marker string.
Args:
marker_type: Type of marker (e.g., "tool_digest", "dropped_context").
**kwargs: Attributes to include in the marker.
Returns:
Formatted marker string.
"""
attrs = " ".join(f'{k}="{v}"' for k, v in kwargs.items())
if attrs:
return f"{MARKER_PREFIX}{marker_type} {attrs}{MARKER_SUFFIX}"
return f"{MARKER_PREFIX}{marker_type}{MARKER_SUFFIX}"
def create_tool_digest_marker(original_hash: str) -> str:
"""Create marker for crushed tool output."""
return create_marker("tool_digest", sha256=original_hash)
def create_dropped_context_marker(reason: str, count: int | None = None) -> str:
"""Create marker for dropped context."""
if count is not None:
return create_marker("dropped_context", reason=reason, count=str(count))
return create_marker("dropped_context", reason=reason)
def create_truncated_marker(original_length: int, truncated_to: int) -> str:
"""Create marker for truncated content."""
return create_marker(
"truncated",
original=str(original_length),
truncated_to=str(truncated_to),
)
def extract_markers(text: str) -> list[dict[str, Any]]:
"""
Extract Headroom markers from text.
Returns:
List of dicts with marker_type and attributes.
"""
pattern = re.compile(r"<headroom:(\w+)([^>]*)>")
markers = []
for match in pattern.finditer(text):
marker_type = match.group(1)
attrs_str = match.group(2).strip()
# Parse attributes
attrs: dict[str, str] = {}
if attrs_str:
attr_pattern = re.compile(r'(\w+)="([^"]*)"')
for attr_match in attr_pattern.finditer(attrs_str):
attrs[attr_match.group(1)] = attr_match.group(2)
markers.append({"type": marker_type, "attributes": attrs})
return markers
def safe_json_loads(text: str) -> tuple[Any | None, bool]:
"""
Safely parse JSON, returning (result, success).
Args:
text: JSON string to parse.
Returns:
Tuple of (parsed_result or None, success_bool).
"""
try:
return json.loads(text), True
except (json.JSONDecodeError, ValueError):
return None, False
def safe_json_dumps(obj: Any, **kwargs: Any) -> str:
"""
Safely serialize to JSON with defaults.
Args:
obj: Object to serialize.
**kwargs: Additional json.dumps arguments.
Returns:
JSON string.
"""
kwargs.setdefault("ensure_ascii", False)
kwargs.setdefault("separators", (",", ":")) # Compact by default
return json.dumps(obj, **kwargs)
def estimate_cost(
input_tokens: int,
output_tokens: int,
model: str,
cached_tokens: int = 0,
provider: Any = None,
) -> float | None:
"""
Estimate API cost in USD using provider.
Args:
input_tokens: Number of input tokens.
output_tokens: Number of output tokens.
model: Model name.
cached_tokens: Number of cached input tokens.
provider: Provider instance for cost estimation.
Returns:
Estimated cost in USD, or None if not available.
"""
if provider is None:
return None
result = provider.estimate_cost(input_tokens, output_tokens, model, cached_tokens)
return float(result) if result is not None else None
def format_cost(cost: float) -> str:
"""Format cost as human-readable string."""
if cost < 0.01:
return f"${cost:.4f}"
return f"${cost:.2f}"
def deep_copy_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Create a deep copy of messages list."""
result: list[dict[str, Any]] = json.loads(json.dumps(messages))
return result
|