@@ -49,17 +49,23 @@ def json(self):
4949
5050
5151class _RecordingSession :
52- """Captures the args of a single ``.post`` and returns a canned response ."""
52+ """Records each ``.post`` and returns canned responses, repeating the last ."""
5353
54- def __init__ (self , response ):
55- self ._response = response
54+ def __init__ (self , * responses ):
55+ self ._responses = list ( responses )
5656 self .calls = []
5757
58- def post (self , url , data = None , headers = None , timeout = None ):
58+ def post (self , url , data = None , headers = None , timeout = None , allow_redirects = True ):
5959 self .calls .append (
60- {"url" : url , "data" : data , "headers" : headers , "timeout" : timeout }
60+ {
61+ "url" : url ,
62+ "data" : data ,
63+ "headers" : headers ,
64+ "timeout" : timeout ,
65+ "allow_redirects" : allow_redirects ,
66+ }
6167 )
62- return self ._response
68+ return self ._responses [ min ( len ( self . calls ), len ( self . _responses )) - 1 ]
6369
6470
6571class _PostV1Stub :
@@ -212,6 +218,90 @@ def test_zstd_without_package_raises_actionable_error(self) -> None:
212218 self ._post (_results_response ({}), compression = CaptureCompression .ZSTD )
213219 self .assertIn ("posthog[zstd]" , str (ctx .exception ))
214220
221+ @parameterized .expand (
222+ [
223+ (
224+ "relative_location" ,
225+ "https://us.i.posthog.com" ,
226+ "/i/v1/analytics/events?retry=1" ,
227+ "https://us.i.posthog.com/i/v1/analytics/events?retry=1" ,
228+ ),
229+ (
230+ "host_path_prefix" ,
231+ "https://example.com/ingest" ,
232+ "https://example.com/ingest/i/v1/analytics/events" ,
233+ "https://example.com/ingest/i/v1/analytics/events" ,
234+ ),
235+ (
236+ "explicit_default_port" ,
237+ "https://us.i.posthog.com" ,
238+ "https://us.i.posthog.com:443/other" ,
239+ "https://us.i.posthog.com/other" ,
240+ ),
241+ ]
242+ )
243+ def test_follows_same_origin_redirect_with_same_body (
244+ self , _name , host , location , expected_url
245+ ) -> None :
246+ final = _results_response ({})
247+ session = _RecordingSession (
248+ _FakeResponse (307 , headers = {"Location" : location }), final
249+ )
250+ body = _build_v1_batch_body ([_to_v1_event (_msg ("u-1" ))])
251+ res = _post_v1 (
252+ "phc_key" , host , body , attempt = 1 , request_id = "r" , session = session
253+ )
254+
255+ self .assertIs (res , final )
256+ self .assertEqual ([c ["url" ] for c in session .calls ][1 ], expected_url )
257+ self .assertEqual (session .calls [0 ]["data" ], session .calls [1 ]["data" ])
258+ self .assertEqual (session .calls [0 ]["headers" ], session .calls [1 ]["headers" ])
259+ self .assertTrue (all (c ["allow_redirects" ] is False for c in session .calls ))
260+
261+ @parameterized .expand (
262+ [
263+ ("other_host" , 307 , "https://attacker.example.com/collect" ),
264+ ("loopback" , 308 , "http://127.0.0.1:8080/collect" ),
265+ ("https_to_http" , 307 , "http://us.i.posthog.com/i/v1/analytics/events" ),
266+ ("other_port" , 308 , "https://us.i.posthog.com:8443/i/v1/analytics/events" ),
267+ ("missing_location" , 307 , None ),
268+ ("not_307_or_308" , 302 , "/i/v1/analytics/events" ),
269+ ]
270+ )
271+ def test_does_not_follow_other_redirects (self , _name , status , location ) -> None :
272+ redirect = _FakeResponse (
273+ status , headers = {"Location" : location } if location else {}
274+ )
275+ session = _RecordingSession (redirect , _results_response ({}))
276+ body = _build_v1_batch_body ([_to_v1_event (_msg ("u-1" ))])
277+ res = _post_v1 (
278+ "phc_key" ,
279+ "https://us.i.posthog.com" ,
280+ body ,
281+ attempt = 1 ,
282+ request_id = "r" ,
283+ session = session ,
284+ )
285+
286+ self .assertIs (res , redirect )
287+ self .assertEqual (len (session .calls ), 1 )
288+
289+ def test_stops_after_max_redirects (self ) -> None :
290+ loop = _FakeResponse (307 , headers = {"Location" : "/i/v1/analytics/events" })
291+ session = _RecordingSession (loop )
292+ body = _build_v1_batch_body ([_to_v1_event (_msg ("u-1" ))])
293+ res = _post_v1 (
294+ "phc_key" ,
295+ "https://us.i.posthog.com" ,
296+ body ,
297+ attempt = 1 ,
298+ request_id = "r" ,
299+ session = session ,
300+ )
301+
302+ self .assertIs (res , loop )
303+ self .assertEqual (len (session .calls ), 6 )
304+
215305
216306class TestParseV1Response (unittest .TestCase ):
217307 def test_success_parses_results_with_details (self ) -> None :
@@ -449,7 +539,14 @@ def test_malformed_2xx_is_terminal(self) -> None:
449539 self .assertEqual (len (stub .calls ), 1 )
450540 self .assertEqual (exc .status , 200 )
451541
452- @parameterized .expand ([("bad_request" , 400 ), ("rate_limited" , 429 )])
542+ @parameterized .expand (
543+ [
544+ ("bad_request" , 400 ),
545+ ("rate_limited" , 429 ),
546+ ("unfollowed_redirect" , 307 ),
547+ ("unfollowed_permanent_redirect" , 308 ),
548+ ]
549+ )
453550 def test_terminal_status_raises_immediately (self , _name , status ) -> None :
454551 stub , exc = self ._run_expecting_error (
455552 [_msg ("u-1" )],
0 commit comments