proxy: propagate cloud stream failures

This commit is contained in:
ParthSareen
2026-09-06 22:45:30 -07:00
parent 7cbab9882d
commit fdc64eed40
2 changed files with 181 additions and 71 deletions
+3 -1
View File
@@ -264,7 +264,9 @@ func proxyCloudRequestWithPath(c *gin.Context, body []byte, path string, disable
"request_context_err", ctxErr,
"error", err,
)
return
// Do not finish an incomplete upstream response as a successful stream.
// Propagate the abort through recovery middleware to net/http.
panic(http.ErrAbortHandler)
}
}
+178 -70
View File
@@ -1,8 +1,10 @@
package server
import (
"encoding/json"
"errors"
"io"
"log/slog"
"net"
"net/http"
"net/http/httptest"
@@ -32,80 +34,186 @@ func TestCodexProxyHealthRoute(t *testing.T) {
}
}
func TestCodexProxyUpstreamDisconnect(t *testing.T) {
for _, http2 := range []bool{false, true} {
name := "HTTP1"
if http2 {
name = "HTTP2"
func TestCodexProxyStreamTermination(t *testing.T) {
for _, cloudHop := range []bool{false, true} {
route := "direct"
if cloudHop {
route = "cloud-hop"
}
t.Run(name, func(t *testing.T) {
const partial = "data: {\"type\":\"response.created\"}\n\n"
disconnect := make(chan struct{})
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = io.WriteString(w, partial)
w.(http.Flusher).Flush()
select {
case <-disconnect:
case <-r.Context().Done():
}
panic(http.ErrAbortHandler)
}))
defer upstream.Close()
defer close(disconnect)
home := t.TempDir()
t.Setenv("HOME", home)
t.Setenv("OLLAMA_HOST", upstream.URL)
catalogDir := filepath.Join(home, ".codex")
if err := os.MkdirAll(catalogDir, 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(catalogDir, proxy.CodexDesktopRoutingCatalogFilename), []byte(`{"models":[{"slug":"glm-5.3-flash:cloud"}]}`), 0o600); err != nil {
t.Fatal(err)
}
handler, err := (&Server{}).GenerateRoutes()
if err != nil {
t.Fatal(err)
}
server := httptest.NewUnstartedServer(handler)
server.EnableHTTP2 = http2
server.StartTLS()
defer server.Close()
client := server.Client()
client.Timeout = 5 * time.Second
resp, err := client.Post(server.URL+"/api/codex/v1/responses", "application/json", strings.NewReader(`{"model":"glm-5.3-flash:cloud","stream":true,"input":[{"role":"user","content":"List files"}]}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
wantProto := 1
for _, http2 := range []bool{false, true} {
protocol := "HTTP1"
if http2 {
wantProto = 2
protocol = "HTTP2"
}
if resp.StatusCode != http.StatusOK || resp.ProtoMajor != wantProto {
t.Fatalf("response = %s %s", resp.Proto, resp.Status)
for _, abort := range []bool{true, false} {
outcome := "complete"
if abort {
outcome = "disconnect"
}
t.Run(route+"/"+protocol+"/"+outcome, func(t *testing.T) {
const partial = "data: {\"type\":\"response.created\"}\n\n"
const completed = "data: {\"type\":\"response.completed\"}\n\ndata: [DONE]\n\n"
release := make(chan struct{})
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = io.WriteString(w, partial)
w.(http.Flusher).Flush()
select {
case <-release:
case <-r.Context().Done():
return
}
if abort {
panic(http.ErrAbortHandler)
}
_, _ = io.WriteString(w, completed)
}))
defer upstream.Close()
defer close(release)
home := t.TempDir()
t.Setenv("HOME", home)
t.Setenv("OLLAMA_NO_CLOUD", "false")
t.Setenv("OLLAMA_HOST", upstream.URL)
catalogDir := filepath.Join(home, ".codex")
if err := os.MkdirAll(catalogDir, 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(catalogDir, proxy.CodexDesktopRoutingCatalogFilename), []byte(`{"models":[{"slug":"stream-probe:cloud"}]}`), 0o600); err != nil {
t.Fatal(err)
}
serverLog, err := os.Create(filepath.Join(home, "server.log"))
if err != nil {
t.Fatal(err)
}
defer serverLog.Close()
oldLogger, oldWriter, oldErrorWriter := slog.Default(), gin.DefaultWriter, gin.DefaultErrorWriter
slog.SetDefault(slog.New(slog.NewTextHandler(serverLog, nil)))
gin.DefaultWriter, gin.DefaultErrorWriter = serverLog, serverLog
defer func() {
slog.SetDefault(oldLogger)
gin.DefaultWriter, gin.DefaultErrorWriter = oldWriter, oldErrorWriter
}()
// The extra HTTP listener exercises the same two-hop route used by
// ollama serve: /api/codex/v1/responses -> /v1/responses -> cloud.
var inner *httptest.Server
if cloudHop {
inner = httptest.NewUnstartedServer(nil)
defer inner.Close()
t.Setenv("OLLAMA_HOST", "http://"+inner.Listener.Addr().String())
original := cloudProxyBaseURL
cloudProxyBaseURL = upstream.URL
defer func() { cloudProxyBaseURL = original }()
}
handler, err := (&Server{}).GenerateRoutes()
if err != nil {
t.Fatal(err)
}
if inner != nil {
inner.Config.Handler = handler
inner.Start()
}
done := make(chan struct{})
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/api/codex/v1/responses" {
defer close(done)
}
handler.ServeHTTP(w, r)
}))
server.EnableHTTP2 = http2
server.StartTLS()
defer server.Close()
client := server.Client()
client.Timeout = 5 * time.Second
resp, err := client.Post(server.URL+"/api/codex/v1/responses", "application/json", strings.NewReader(`{"model":"stream-probe:cloud","stream":true,"input":[{"role":"user","content":"Synthetic stream probe"}]}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
wantProto := 1
if http2 {
wantProto = 2
}
if resp.StatusCode != http.StatusOK || resp.ProtoMajor != wantProto {
t.Fatalf("response = %s %s", resp.Proto, resp.Status)
}
prefix := make([]byte, len(partial))
if _, err := io.ReadFull(resp.Body, prefix); err != nil || string(prefix) != partial {
t.Fatalf("partial response = %q, error = %v", prefix, err)
}
// Disconnect only after the real client has received partial output.
select {
case release <- struct{}{}:
case <-time.After(5 * time.Second):
t.Fatal("upstream did not accept release")
}
rest, readErr := io.ReadAll(resp.Body)
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("proxy handler did not finish")
}
t.Logf("client: protocol=%s status=%d partial=%q remaining=%q read_error=%v", resp.Proto, resp.StatusCode, prefix, rest, readErr)
if abort {
if !http2 && !errors.Is(readErr, io.ErrUnexpectedEOF) {
t.Errorf("HTTP/1 read error = %v, want unexpected EOF", readErr)
} else if http2 && (readErr == nil || !strings.Contains(readErr.Error(), "INTERNAL_ERROR")) {
t.Errorf("HTTP/2 read error = %v, want stream reset", readErr)
}
if len(rest) != 0 {
t.Errorf("unexpected output after abort: %q", rest)
}
} else if readErr != nil || string(rest) != completed {
t.Errorf("complete stream: remaining=%q error=%v", rest, readErr)
}
wantResult, wantErrors := "ok", 0
if abort {
wantResult, wantErrors = "stream_error", 1
}
logData, err := os.ReadFile(filepath.Join(home, ".ollama", "logs", codexDesktopLogFilename))
if err != nil {
t.Fatal(err)
}
t.Logf("activity: %s", logData)
if !strings.Contains(string(logData), "status=200 ") || !strings.Contains(string(logData), "result="+wantResult) {
t.Errorf("activity log did not record %s: %s", wantResult, logData)
}
status, err := client.Get(server.URL + "/api/codex/_status")
if err != nil {
t.Fatal(err)
}
defer status.Body.Close()
var metrics struct {
UpstreamErrors int `json:"upstream_errors"`
}
if err := json.NewDecoder(status.Body).Decode(&metrics); err != nil {
t.Fatal(err)
}
if metrics.UpstreamErrors != wantErrors {
t.Errorf("upstream_errors=%d, want %d", metrics.UpstreamErrors, wantErrors)
}
t.Logf("status: upstream_errors=%d", metrics.UpstreamErrors)
logs, err := os.ReadFile(serverLog.Name())
if err != nil {
t.Fatal(err)
}
for _, line := range strings.Split(string(logs), "\n") {
if strings.Contains(line, "level=WARN") || strings.Contains(line, "level=ERROR") || strings.Contains(line, "[Recovery]") {
t.Logf("server: %s", line)
}
}
if strings.Contains(string(logs), "panic recovered") || strings.Contains(string(logs), "override status code 200 with 500") {
t.Errorf("middleware treated a stream abort as an application panic:\n%s", logs)
}
if abort && !strings.Contains(string(logs), "Codex proxy response stream aborted") {
t.Error("server log lost the stream failure")
}
})
}
prefix := make([]byte, len(partial))
if _, err := io.ReadFull(resp.Body, prefix); err != nil || string(prefix) != partial {
t.Fatalf("partial response = %q, error = %v", prefix, err)
}
disconnect <- struct{}{}
rest, err := io.ReadAll(resp.Body)
if err == nil {
t.Errorf("truncated response ended successfully: %q", rest)
} else if !http2 && !errors.Is(err, io.ErrUnexpectedEOF) {
t.Errorf("HTTP/1 read error = %v, want unexpected EOF", err)
} else if http2 && !strings.Contains(err.Error(), "INTERNAL_ERROR") {
t.Errorf("HTTP/2 read error = %v, want stream reset", err)
}
logData, err := os.ReadFile(filepath.Join(home, ".ollama", "logs", codexDesktopLogFilename))
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(logData), "status=200 ") || !strings.Contains(string(logData), "result=stream_error") {
t.Fatalf("activity log did not record failed stream: %s", logData)
}
})
}
}
}