128 lines
3.3 KiB
Go
128 lines
3.3 KiB
Go
package openai_compat
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"strings"
|
|
"time"
|
|
|
|
runtime "iop/packages/go/execution"
|
|
)
|
|
|
|
// TunnelProvider relays an arbitrary HTTP request to the OpenAI-compatible
|
|
// endpoint and streams the response back through the sink.
|
|
func (a *Adapter) TunnelProvider(ctx context.Context, req runtime.ProviderTunnelRequest, sink runtime.ProviderTunnelSink) error {
|
|
var seq int64 = 0
|
|
|
|
if a.endpoint == "" && a.profile == nil {
|
|
err := fmt.Errorf("openai_compat adapter: endpoint or profile is required")
|
|
_ = emitTunnelError(ctx, sink, req, seq, err)
|
|
return err
|
|
}
|
|
|
|
isJSON := req.Method == http.MethodPost
|
|
httpReq, err := a.prepareRequest(ctx, req.Method, req.Operation, req.Path, req.Body, req.Headers, isJSON)
|
|
if err != nil {
|
|
_ = emitTunnelError(ctx, sink, req, seq, err)
|
|
return err
|
|
}
|
|
if req.Credential != nil {
|
|
defer req.Credential.Zero()
|
|
for header := range httpReq.Header {
|
|
if strings.EqualFold(header, req.Credential.HeaderName) {
|
|
err := fmt.Errorf("provider credential header collision")
|
|
_ = emitTunnelError(ctx, sink, req, seq, err)
|
|
return err
|
|
}
|
|
}
|
|
value := string(req.Credential.Secret)
|
|
if scheme := strings.TrimSpace(req.Credential.Scheme); scheme != "" {
|
|
value = scheme + " " + value
|
|
}
|
|
httpReq.Header.Set(req.Credential.HeaderName, value)
|
|
defer httpReq.Header.Del(req.Credential.HeaderName)
|
|
}
|
|
|
|
resp, err := a.client.Do(httpReq)
|
|
if err != nil {
|
|
err = fmt.Errorf("request failed: %w", err)
|
|
_ = emitTunnelError(ctx, sink, req, seq, err)
|
|
return err
|
|
}
|
|
defer resp.Body.Close()
|
|
|
|
respHeaders := make(map[string]string)
|
|
for k, values := range resp.Header {
|
|
if len(values) > 0 {
|
|
respHeaders[k] = values[0]
|
|
}
|
|
}
|
|
|
|
err = sink.EmitTunnelFrame(ctx, runtime.ProviderTunnelFrame{
|
|
RunID: req.RunID,
|
|
TunnelID: req.TunnelID,
|
|
Sequence: seq,
|
|
Kind: runtime.ProviderTunnelFrameKindResponseStart,
|
|
StatusCode: resp.StatusCode,
|
|
Headers: respHeaders,
|
|
Timestamp: time.Now(),
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
seq++
|
|
|
|
buf := make([]byte, 4096)
|
|
for {
|
|
n, readErr := resp.Body.Read(buf)
|
|
if n > 0 {
|
|
err = sink.EmitTunnelFrame(ctx, runtime.ProviderTunnelFrame{
|
|
RunID: req.RunID,
|
|
TunnelID: req.TunnelID,
|
|
Sequence: seq,
|
|
Kind: runtime.ProviderTunnelFrameKindBody,
|
|
Body: append([]byte(nil), buf[:n]...),
|
|
Timestamp: time.Now(),
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
seq++
|
|
}
|
|
if readErr != nil {
|
|
if readErr == io.EOF {
|
|
break
|
|
}
|
|
if ctx.Err() != nil {
|
|
_ = emitTunnelError(ctx, sink, req, seq, ctx.Err())
|
|
return ctx.Err()
|
|
}
|
|
err = fmt.Errorf("read response body: %w", readErr)
|
|
_ = emitTunnelError(ctx, sink, req, seq, err)
|
|
return err
|
|
}
|
|
}
|
|
|
|
err = sink.EmitTunnelFrame(ctx, runtime.ProviderTunnelFrame{
|
|
RunID: req.RunID,
|
|
TunnelID: req.TunnelID,
|
|
Sequence: seq,
|
|
Kind: runtime.ProviderTunnelFrameKindEnd,
|
|
End: true,
|
|
Timestamp: time.Now(),
|
|
})
|
|
return err
|
|
}
|
|
|
|
func emitTunnelError(ctx context.Context, sink runtime.ProviderTunnelSink, req runtime.ProviderTunnelRequest, seq int64, err error) error {
|
|
return sink.EmitTunnelFrame(ctx, runtime.ProviderTunnelFrame{
|
|
RunID: req.RunID,
|
|
TunnelID: req.TunnelID,
|
|
Sequence: seq,
|
|
Kind: runtime.ProviderTunnelFrameKindError,
|
|
Error: err.Error(),
|
|
Timestamp: time.Now(),
|
|
})
|
|
}
|