120 lines
2.9 KiB
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(),
|
|
})
|
|
}
|