iop/apps/node/internal/adapters/vllm/provider_tunnel.go

120 lines
2.9 KiB
Go

package vllm
import (
"bytes"
"context"
"fmt"
"io"
"net/http"
"time"
"iop/apps/node/internal/runtime"
)
// TunnelProvider relays an arbitrary HTTP request to the vLLM endpoint and
// streams the response back through the sink.
func (v *Vllm) TunnelProvider(ctx context.Context, req runtime.ProviderTunnelRequest, sink runtime.ProviderTunnelSink) error {
var seq int64 = 0
if v.endpoint == "" {
err := fmt.Errorf("vllm adapter: endpoint is required")
_ = emitTunnelError(ctx, sink, req, seq, err)
return err
}
urlStr := joinURL(v.endpoint, req.Path)
httpReq, err := http.NewRequestWithContext(ctx, req.Method, urlStr, bytes.NewReader(req.Body))
if err != nil {
err = fmt.Errorf("build request: %w", err)
_ = emitTunnelError(ctx, sink, req, seq, err)
return err
}
if req.Method == http.MethodPost {
httpReq.Header.Set("Content-Type", "application/json")
}
for k, val := range req.Headers {
httpReq.Header.Set(k, val)
}
resp, err := v.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(),
})
}