@@ -10,12 +10,14 @@ from __future__ import annotations
1010
1111import argparse
1212import base64
13+ import gzip
1314import json
1415import os
1516import re
1617import tempfile
1718import threading
1819import uuid
20+ import zlib
1921from datetime import datetime , timezone
2022from http import HTTPStatus
2123from http .server import BaseHTTPRequestHandler , ThreadingHTTPServer
@@ -44,6 +46,7 @@ SAFE_BROWSER_HEADERS = {
4446 "content-security-policy" : "default-src 'none'; sandbox" ,
4547 "x-content-type-options" : "nosniff" ,
4648}
49+ ZSTD_FRAME_MAGIC = b"\x28 \xb5 \x2f \xfd "
4750
4851
4952class RecordingDatabase :
@@ -52,9 +55,11 @@ class RecordingDatabase:
5255 self .lock = threading .RLock ()
5356 self .shards : dict [tuple [str , str ], dict [str , Any ]] = {}
5457 self .shard_paths : dict [tuple [str , str ], Path ] = {}
58+ self .request_plans : dict [tuple [str , str , str ], dict [str , Any ]] = {}
5559 self .sessions : dict [str , dict [str , Any ]] = {}
5660 self .fallback_consumed : set [tuple [str , str , str , int ]] = set ()
5761 self ._load ()
62+ self ._load_request_plans ()
5863
5964 def _load (self ) -> None :
6065 manifest_path = self .root / "manifest.json"
@@ -70,6 +75,17 @@ class RecordingDatabase:
7075 self .shards [key ] = shard
7176 self .shard_paths [key ] = path
7277
78+ def _load_request_plans (self ) -> None :
79+ root = self .root .parent / "test-runner-data"
80+ manifest_path = root / "manifest.json"
81+ if not manifest_path .exists ():
82+ return
83+ manifest = _read_json (manifest_path )
84+ for item in manifest .get ("scenarios" , []):
85+ plan = _read_json (root / item ["file" ])
86+ key = (item ["version" ], item ["feature" ], item ["scenario" ])
87+ self .request_plans [key ] = plan .get ("request" , {})
88+
7389 def start (self , version : str , feature : str , scenario : str , mode : str ) -> dict [str , Any ]:
7490 with self .lock :
7591 key = (version , feature )
@@ -91,6 +107,7 @@ class RecordingDatabase:
91107 "key" : key ,
92108 "scenario" : scenario ,
93109 "recording" : recording ,
110+ "request_plan" : self .request_plans .get ((version , feature , scenario )),
94111 "cursor" : 0 ,
95112 "captures" : [],
96113 "frozen_at" : frozen_at ,
@@ -109,7 +126,7 @@ class RecordingDatabase:
109126 if cursor >= len (interactions ):
110127 raise LookupError (f"Recording has no interaction #{ cursor + 1 } " )
111128 expected = interactions [cursor ]
112- if not _requests_match (expected ["request" ], actual ):
129+ if not _requests_match (expected ["request" ], actual , session [ "request_plan" ] ):
113130 raise RequestMismatchError (expected ["request" ], actual , cursor )
114131 session ["cursor" ] += 1
115132 return expected ["response" ]
@@ -298,7 +315,13 @@ class TestRequestHandler(BaseHTTPRequestHandler):
298315
299316 def _handle_api_request (self ) -> None :
300317 body = self ._read_body ()
301- actual = _normalise_request (self .command , self .path , self .headers .get ("content-type" , "" ), body )
318+ actual = _normalise_request (
319+ self .command ,
320+ self .path ,
321+ self .headers .get ("content-type" , "" ),
322+ self .headers .get ("content-encoding" , "" ),
323+ body ,
324+ )
302325 session_id = self .headers .get (SESSION_HEADER )
303326 if self .server .mode == "replay" :
304327 response = self .server .database .replay (session_id , actual )
@@ -372,21 +395,50 @@ class TestRequestHandler(BaseHTTPRequestHandler):
372395 self .wfile .write (body )
373396
374397
375- def _normalise_request (method : str , raw_path : str , content_type : str , body : bytes ) -> dict [str , Any ]:
398+ def _normalise_request (
399+ method : str ,
400+ raw_path : str ,
401+ content_type : str ,
402+ content_encoding : str ,
403+ body : bytes ,
404+ ) -> dict [str , Any ]:
376405 parsed = urlsplit (raw_path )
377- normalised_body = _normalise_body (body , content_type )
378- return {
406+ compression = content_encoding .strip ().casefold ()
407+ normalised_body = _normalise_body (body , content_type , compression )
408+ request = {
379409 "method" : method .upper (),
380410 "path" : _normalise_path (parsed .path ),
381411 "query" : sorted ([list (pair ) for pair in parse_qsl (parsed .query , keep_blank_values = True )]),
382412 "content_type" : _normalise_content_type (content_type , normalised_body ),
383413 "body" : normalised_body ,
384414 }
415+ if compression :
416+ request ["compression" ] = compression
417+ return request
385418
386419
387- def _normalise_body (body : bytes , content_type : str ) -> dict [str , Any ]:
420+ def _normalise_body (body : bytes , content_type : str , compression : str = "" ) -> dict [str , Any ]:
388421 if not body :
389422 return {"type" : "empty" , "value" : None }
423+ if compression == "zstd1" :
424+ if len (body ) > len (ZSTD_FRAME_MAGIC ) and body .startswith (ZSTD_FRAME_MAGIC ):
425+ # The generated server has no third-party dependencies, so validate
426+ # the Zstandard frame on the wire without attempting to decode it.
427+ return {"type" : "zstd1" , "value" : None }
428+ return {
429+ "type" : "invalid-compression" ,
430+ "value" : base64 .b64encode (body ).decode ("ascii" ),
431+ }
432+ try :
433+ if compression == "gzip" :
434+ body = gzip .decompress (body )
435+ elif compression == "deflate" :
436+ body = zlib .decompress (body )
437+ except (EOFError , OSError , zlib .error ):
438+ return {
439+ "type" : "invalid-compression" ,
440+ "value" : base64 .b64encode (body ).decode ("ascii" ),
441+ }
390442 media_type = _media_type (content_type )
391443 text = body .decode ("utf-8" , errors = "surrogateescape" )
392444 if media_type .endswith ("json" ):
@@ -415,16 +467,38 @@ def _normalise_content_type(content_type: str, body: dict[str, Any]) -> str:
415467 return "" if body ["type" ] == "empty" else _media_type (content_type )
416468
417469
418- def _requests_match (expected : dict [str , Any ], actual : dict [str , Any ]) -> bool :
470+ def _requests_match (
471+ expected : dict [str , Any ],
472+ actual : dict [str , Any ],
473+ request_plan : dict [str , Any ] | None = None ,
474+ ) -> bool :
419475 comparable_fields = ("method" , "path" , "query" , "content_type" )
420476 if any (expected [field ] != actual [field ] for field in comparable_fields ):
421477 return False
478+ expected_compression = expected .get ("compression" )
479+ if request_plan and _request_matches_plan (expected , request_plan ):
480+ expected_compression = request_plan .get ("compression" , expected_compression )
481+ if expected_compression is not None and actual .get ("compression" , "" ) != expected_compression :
482+ return False
422483 return _bodies_match (expected ["body" ], actual ["body" ])
423484
424485
486+ def _request_matches_plan (request : dict [str , Any ], plan : dict [str , Any ]) -> bool :
487+ if request ["method" ] != plan .get ("method" ):
488+ return False
489+ path = plan .get ("path" )
490+ if not path :
491+ return False
492+ parts = re .split (r"(\{[^/{}]+\})" , path )
493+ pattern = "" .join (r"[^/]+" if part .startswith ("{" ) else re .escape (part ) for part in parts )
494+ return re .fullmatch (pattern , request ["path" ]) is not None
495+
496+
425497def _bodies_match (expected : dict [str , Any ], actual : dict [str , Any ]) -> bool :
426498 if expected == actual :
427499 return True
500+ if actual ["type" ] == "zstd1" :
501+ return expected ["type" ] != "empty"
428502 if expected ["type" ] != "json" or actual ["type" ] != "json" :
429503 return False
430504 return _json_contains (actual ["value" ], expected ["value" ])
0 commit comments