package proxy_test import ( "io" "net/http" "net/http/httptest" "net/url" "testing" "time" "git.misaka.ren/M1saka/token_thief/config" "git.misaka.ren/M1saka/token_thief/logger" "git.misaka.ren/M1saka/token_thief/proxy" ) type sliceSubmitter struct{ entries []*logger.LogEntry } func (s *sliceSubmitter) Submit(e *logger.LogEntry) { s.entries = append(s.entries, e) } // TestUpstreamErrorRecorded 验证上游不可达时 502 响应被记录、错误信息进入 LogEntry.Error。 func TestUpstreamErrorRecorded(t *testing.T) { // 指向一个一定不可用的端口 badURL, _ := url.Parse("http://127.0.0.1:1") // port 1 几乎肯定 connection refused sub := &sliceSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(badURL, filter, sub, 1024) srv := httptest.NewServer(h) defer srv.Close() resp, err := http.Get(srv.URL + "/v1/chat/completions") if err != nil { t.Fatal(err) } defer resp.Body.Close() body, _ := io.ReadAll(resp.Body) if resp.StatusCode != http.StatusBadGateway { t.Errorf("status=%d want 502, body=%q", resp.StatusCode, body) } if resp.Header.Get("X-Request-Id") == "" { t.Errorf("missing X-Request-Id header") } deadline := time.Now().Add(2 * time.Second) for len(sub.entries) == 0 && time.Now().Before(deadline) { time.Sleep(10 * time.Millisecond) } if len(sub.entries) == 0 { t.Fatal("no log entry captured") } e := sub.entries[0] if e.StatusCode != http.StatusBadGateway { t.Errorf("LogEntry.StatusCode=%d want 502", e.StatusCode) } if e.Error == "" { t.Errorf("LogEntry.Error should be set, got empty") } if string(e.ResponseBody) != "bad gateway\n" { t.Errorf("response_body should be fixed bad gateway, got %q", e.ResponseBody) } if e.RequestID == "" { t.Errorf("LogEntry.RequestID empty") } } // TestModifyResponseSetsRequestID 验证正常上游响应也会带上 X-Request-Id。 func TestModifyResponseSetsRequestID(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(200) _, _ = w.Write([]byte(`{"ok":true}`)) })) defer upstream.Close() u, _ := url.Parse(upstream.URL) sub := &sliceSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, sub, 1024) srv := httptest.NewServer(h) defer srv.Close() resp, err := http.Get(srv.URL + "/v1/anything") if err != nil { t.Fatal(err) } defer resp.Body.Close() rid := resp.Header.Get("X-Request-Id") if len(rid) != 32 { t.Errorf("X-Request-Id length=%d want 32, value=%q", len(rid), rid) } deadline := time.Now().Add(2 * time.Second) for len(sub.entries) == 0 && time.Now().Before(deadline) { time.Sleep(10 * time.Millisecond) } if len(sub.entries) == 0 { t.Fatal("no log entry") } if sub.entries[0].RequestID != rid { t.Errorf("LogEntry.RequestID=%q response header=%q (should match)", sub.entries[0].RequestID, rid) } } func TestUpstreamTLSInsecureSkipVerifyAllowsSelfSigned(t *testing.T) { upstream := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(`{"ok":true}`)) })) defer upstream.Close() u, _ := url.Parse(upstream.URL) sub := &sliceSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.NewWithOptions(u, filter, sub, 1024, proxy.Options{UpstreamTLSInsecureSkipVerify: true}) srv := httptest.NewServer(h) defer srv.Close() resp, err := http.Get(srv.URL + "/v1/anything") if err != nil { t.Fatal(err) } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { body, _ := io.ReadAll(resp.Body) t.Fatalf("status=%d want 200, body=%q", resp.StatusCode, body) } }