scripts/replay_trusted_request_dump.py
scripts/replay_trusted_request_dump.pyBrowse 7 files
1,555 tokens
6,895 bytes
Token encoding: o200k_base
Snapshot a9fb1c3
← Back to SKILL.md
1#!/usr/bin/env python32"""Replay a trusted SGLang request dump directly over HTTP.3 4Use this only for locally captured or otherwise trusted dump files.5It uses plain pickle loading to bypass SafeUnpickler restrictions that may block6the stock replay helper on newer SGLang builds.7"""8 9from __future__ import annotations10 11import argparse12import glob13import json14import pickle15import time16from concurrent.futures import ThreadPoolExecutor17from dataclasses import asdict, is_dataclass18from datetime import datetime19from pathlib import Path20from typing import Any, Sequence21 22import requests23 24Record = tuple[object, dict[str, Any], float, float]25 26 27def normalize_mm_data_item(item: Any) -> Any:28 if isinstance(item, dict) and "url" in item:29 return item["url"]30 return item31 32 33def normalize_mm_data(data: Any) -> Any:34 if data is None:35 return None36 if isinstance(data, list):37 return [38 (39 [normalize_mm_data_item(item) for item in sublist]40 if isinstance(sublist, list)41 else normalize_mm_data_item(sublist)42 )43 for sublist in data44 ]45 return normalize_mm_data_item(data)46 47 48def normalize_request_data(json_data: dict[str, Any]) -> dict[str, Any]:49 for field in ["image_data", "video_data", "audio_data"]:50 if field in json_data and json_data[field] is not None:51 json_data[field] = normalize_mm_data(json_data[field])52 return json_data53 54 55def to_plain_dict(obj: Any) -> dict[str, Any]:56 if obj is None:57 return {}58 if isinstance(obj, dict):59 return dict(obj)60 if is_dataclass(obj):61 return asdict(obj)62 63 model_dump = getattr(obj, "model_dump", None)64 if callable(model_dump):65 dumped = model_dump()66 if isinstance(dumped, dict):67 return dumped68 69 dict_method = getattr(obj, "dict", None)70 if callable(dict_method):71 dumped = dict_method()72 if isinstance(dumped, dict):73 return dumped74 75 obj_dict = getattr(obj, "__dict__", None)76 if isinstance(obj_dict, dict):77 return {78 key: value for key, value in obj_dict.items() if not key.startswith("_")79 }80 81 raise TypeError(f"Unsupported request object type: {type(obj)!r}")82 83 84def request_to_json_data(req: Any) -> dict[str, Any]:85 json_data = normalize_request_data(to_plain_dict(req))86 sampling_params = json_data.get("sampling_params")87 if sampling_params is not None and not isinstance(sampling_params, dict):88 json_data["sampling_params"] = to_plain_dict(sampling_params)89 return json_data90 91 92def load_records(path: Path) -> list[Record]:93 with path.open("rb") as fh:94 payload = pickle.load(fh)95 if isinstance(payload, dict) and "requests" in payload:96 return payload["requests"]97 return payload98 99 100def iter_files(args: argparse.Namespace) -> Sequence[Path]:101 if args.input_file:102 return [Path(args.input_file)]103 if args.input_folder:104 return [105 Path(p)106 for p in sorted(glob.glob(f"{args.input_folder}/*.pkl"))[: args.file_number]107 ]108 raise SystemExit("Either --input-file or --input-folder must be provided.")109 110 111def run_one_request(112 record: Record,113 args: argparse.Namespace,114 replay_init_time: float,115 base_time: float,116 idx: int,117) -> None:118 req, output, start_time, end_time = record119 relative_start = start_time - base_time120 delay = max(0.0, (relative_start - (time.time() - replay_init_time)) / args.speed)121 if delay:122 time.sleep(delay)123 124 json_data = request_to_json_data(req)125 if args.ignore_eos:126 json_data.setdefault("sampling_params", {})["ignore_eos"] = True127 completion_tokens = output.get("meta_info", {}).get("completion_tokens")128 if completion_tokens:129 json_data["sampling_params"]["max_new_tokens"] = completion_tokens130 131 t0 = time.time()132 response = requests.post(133 f"http://{args.host}:{args.port}/generate",134 json=json_data,135 timeout=args.timeout,136 stream=bool(json_data.get("stream")),137 )138 elapsed = time.time() - t0139 140 if json_data.get("stream"):141 last = None142 for chunk in response.iter_lines(decode_unicode=False):143 decoded = chunk.decode("utf-8")144 if decoded and decoded.startswith("data:"):145 if decoded == "data: [DONE]":146 break147 last = json.loads(decoded[5:].strip())148 result = last or {}149 else:150 result = response.json()151 152 meta = result.get("meta_info", {})153 print(154 json.dumps(155 {156 "idx": idx,157 "status_code": response.status_code,158 "elapsed_seconds": round(elapsed, 3),159 "prompt_tokens": meta.get("prompt_tokens"),160 "completion_tokens": meta.get("completion_tokens"),161 "rid": meta.get("id"),162 },163 ensure_ascii=False,164 )165 )166 167 168def main() -> int:169 parser = argparse.ArgumentParser(170 description="Replay a trusted SGLang request dump or crash dump directly over HTTP."171 )172 parser.add_argument("--host", default="127.0.0.1")173 parser.add_argument("--port", type=int, default=30000)174 parser.add_argument("--input-folder", default=None)175 parser.add_argument("--input-file", default=None)176 parser.add_argument("--file-number", type=int, default=1)177 parser.add_argument("--req-number", type=int, default=1_000_000)178 parser.add_argument("--req-start", type=int, default=0)179 parser.add_argument("--parallel", type=int, default=1)180 parser.add_argument("--ignore-eos", action="store_true")181 parser.add_argument("--speed", type=float, default=1.0)182 parser.add_argument("--timeout", type=float, default=120.0)183 args = parser.parse_args()184 185 files = iter_files(args)186 print(f"Replay files: {[str(p) for p in files]}")187 188 records: list[Record] = []189 for path in files:190 records.extend(load_records(path))191 192 if not records:193 print("No requests found.")194 return 0195 196 records.sort(key=lambda x: x[-2])197 records = records[args.req_start : args.req_start + args.req_number]198 print(f"Replay requests: {len(records)}")199 base_time = records[0][-2]200 print(201 "Base time: " + datetime.fromtimestamp(base_time).strftime("%Y-%m-%d %H:%M:%S")202 )203 204 replay_init_time = time.time()205 with ThreadPoolExecutor(max_workers=args.parallel) as executor:206 futures = []207 for idx, record in enumerate(records):208 futures.append(209 executor.submit(210 run_one_request, record, args, replay_init_time, base_time, idx211 )212 )213 for future in futures:214 future.result()215 return 0216 217 218if __name__ == "__main__":219 raise SystemExit(main())220 Referenced from SKILL.md
SKILL.mdView in source ↗
Source excerpt starting at line 174.SKILL.mdView in source ↗174```bash175python3 scripts/replay_trusted_request_dump.py \176 --input-file /path/to/request_dump.pkl \
Source excerpt starting at line 285.285 - summarize a trusted request dump or crash dump before replay286- [scripts/replay_trusted_request_dump.py](scripts/replay_trusted_request_dump.py)287 - replay a trusted request dump when `safe_pickle_load` blocks stock replay