-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathstream.go
More file actions
153 lines (132 loc) · 3.99 KB
/
Copy pathstream.go
File metadata and controls
153 lines (132 loc) · 3.99 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
package proxy
import (
"bufio"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"github.com/HeyJiqingCode/CopilotGateway/internal/sse"
"github.com/HeyJiqingCode/CopilotGateway/internal/upstream"
)
// headersSentError wraps an error that occurred after response headers were
// already written to the client. The caller must not attempt to write an HTTP
// error response because the status line has already been flushed.
type headersSentError struct{ err error }
func (e *headersSentError) Error() string { return e.err.Error() }
func (e *headersSentError) Unwrap() error { return e.err }
// HandleStreamingRequest handles streaming requests to Copilot API
func (h *Handler) HandleStreamingRequest(w http.ResponseWriter, r *http.Request, endpoint string) error {
var body interface{}
if r.Body != nil {
body = r.Body
}
resp, _, err := h.upstream.Do(r.Context(), upstream.Request{
Method: r.Method,
Endpoint: endpoint,
Body: body,
QueryString: r.URL.RawQuery,
Stream: true,
ExtraHeaders: collectForwardHeaders(r),
})
if err != nil {
var upstreamErr *upstream.UpstreamError
if errors.As(err, &upstreamErr) {
upstreamErr.WriteRawError(w)
return nil
}
return fmt.Errorf("upstream request failed: %w", err)
}
defer resp.Body.Close()
// Set up streaming response headers
sse.BeginSSE(w)
// Flush headers
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
// Stream the response — headers are already sent at this point.
if err := h.streamResponse(w, resp.Body, endpoint); err != nil {
return &headersSentError{err: err}
}
return nil
}
// streamResponse streams the SSE response line by line.
// endpoint is used to select the correct termination strategy.
func (h *Handler) streamResponse(w http.ResponseWriter, body io.ReadCloser, endpoint string) error {
scanner := bufio.NewScanner(body)
// Increase buffer size to handle large SSE lines (default is 64KB, increase to 1MB)
buf := make([]byte, 0, 64*1024)
scanner.Buffer(buf, 1024*1024)
flusher, canFlush := w.(http.Flusher)
isResponses := endpoint == "/responses"
for scanner.Scan() {
line := scanner.Text()
// Write the line
if _, err := io.WriteString(w, line); err != nil {
return err
}
if _, err := io.WriteString(w, "\n"); err != nil {
return err
}
// SSE events are delimited by blank lines; flush at event boundaries
if canFlush && line == "" {
flusher.Flush()
}
// Chat Completions: terminate on data: [DONE]
// Skip this check for the Responses API — it uses event-based
// termination and may send data: [DONE] before response.completed.
if !isResponses && strings.TrimSpace(line) == "data: [DONE]" {
slog.Debug("stream done", "endpoint", endpoint)
if canFlush {
flusher.Flush()
}
break
}
// Responses API termination events
if isResponses && isResponsesTerminationEvent(line) {
slog.Debug("stream termination event", "endpoint", endpoint, "event", strings.TrimSpace(line))
// Write remaining data lines for this event, then stop
for scanner.Scan() {
remaining := scanner.Text()
if _, err := io.WriteString(w, remaining); err != nil {
return err
}
if _, err := io.WriteString(w, "\n"); err != nil {
return err
}
if remaining == "" {
break
}
}
if canFlush {
flusher.Flush()
}
break
}
}
if err := scanner.Err(); err != nil {
return err
}
return nil
}
// isResponsesTerminationEvent checks if an SSE event line indicates
// stream termination for the Responses API.
func isResponsesTerminationEvent(line string) bool {
line = strings.TrimSpace(line)
return line == "event: response.completed" ||
line == "event: response.incomplete" ||
line == "event: response.failed" ||
line == "event: error"
}
// isStreamingRequest checks if the request body wants streaming.
func isStreamingRequest(body []byte) bool {
var top struct {
Stream bool `json:"stream"`
}
if err := json.Unmarshal(body, &top); err != nil {
return false
}
return top.Stream
}