Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
128 changes: 127 additions & 1 deletion internal/server/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ import (
"os/exec"
"strings"
"sync"
"time"

"github.com/devrail-dev/devrail-router/internal/config"
)
Expand Down Expand Up @@ -144,15 +145,29 @@ func (s *Server) proxyOpenAI(w http.ResponseWriter, r *http.Request) {
}

proxy := httputil.NewSingleHostReverseProxy(target)
started := time.Now()
originalDirector := proxy.Director
proxy.Director = func(req *http.Request) {
originalDirector(req)
req.URL.Path = joinOpenAIPath(target.Path, r.URL.Path)
req.Host = target.Host
setBackendAuth(req, backend)
}
proxy.ModifyResponse = func(resp *http.Response) error {
instrumentBackendResponse(resp, started, model, backend)
return nil
}
proxy.ErrorHandler = func(rw http.ResponseWriter, req *http.Request, proxyErr error) {
slog.Error("backend request failed", "method", req.Method, "path", req.URL.Path, "backend", backend.ID, "error", proxyErr)
slog.Error(
"backend request failed",
"method", req.Method,
"path", req.URL.Path,
"alias", model.ID,
"target_model", model.TargetModel,
"backend", backend.ID,
"duration_ms", time.Since(started).Milliseconds(),
"error", proxyErr,
)
http.Error(rw, "backend request failed", http.StatusBadGateway)
}

Expand Down Expand Up @@ -267,6 +282,117 @@ func (s *Server) ensureModelReadyWithCommand(w http.ResponseWriter, r *http.Requ
return true
}

type responseTelemetry struct {
Alias string
TargetModel string
Backend string
UpstreamModel string
Status int
PromptTokens int
CompletionTokens int
TotalTokens int
Bytes int64
Started time.Time
}

type telemetryReadCloser struct {
body io.ReadCloser
telemetry *responseTelemetry
once sync.Once
}

func (body *telemetryReadCloser) Read(p []byte) (int, error) {
n, err := body.body.Read(p)
body.telemetry.Bytes += int64(n)
if errors.Is(err, io.EOF) {
body.log()
}
return n, err
}

func (body *telemetryReadCloser) Close() error {
err := body.body.Close()
body.log()
return err
}

func (body *telemetryReadCloser) log() {
body.once.Do(func() {
slog.Info(
"backend response completed",
"alias", body.telemetry.Alias,
"target_model", body.telemetry.TargetModel,
"upstream_model", body.telemetry.UpstreamModel,
"backend", body.telemetry.Backend,
"status", body.telemetry.Status,
"duration_ms", time.Since(body.telemetry.Started).Milliseconds(),
"bytes", body.telemetry.Bytes,
"prompt_tokens", body.telemetry.PromptTokens,
"completion_tokens", body.telemetry.CompletionTokens,
"total_tokens", body.telemetry.TotalTokens,
)
})
}

func instrumentBackendResponse(resp *http.Response, started time.Time, model config.ModelConfig, backend config.BackendConfig) {
if resp.Body == nil {
return
}

telemetry := &responseTelemetry{
Alias: model.ID,
TargetModel: model.TargetModel,
Backend: backend.ID,
Status: resp.StatusCode,
Started: started,
}

if isJSONResponse(resp) {
raw, err := io.ReadAll(resp.Body)
if closeErr := resp.Body.Close(); closeErr != nil {
slog.Warn("close backend response body", "alias", model.ID, "backend", backend.ID, "error", closeErr)
}
if err == nil {
applyOpenAIUsageTelemetry(raw, telemetry)
resp.Body = io.NopCloser(bytes.NewReader(raw))
resp.ContentLength = int64(len(raw))
resp.Header.Set("Content-Length", fmt.Sprintf("%d", len(raw)))
} else {
slog.Warn("read backend response body for telemetry", "alias", model.ID, "backend", backend.ID, "error", err)
resp.Body = io.NopCloser(bytes.NewReader(nil))
resp.ContentLength = 0
resp.Header.Set("Content-Length", "0")
}
}

resp.Body = &telemetryReadCloser{body: resp.Body, telemetry: telemetry}
}

func isJSONResponse(resp *http.Response) bool {
contentType := strings.ToLower(resp.Header.Get("Content-Type"))
return strings.Contains(contentType, "application/json")
}

func applyOpenAIUsageTelemetry(raw []byte, telemetry *responseTelemetry) {
var payload struct {
Model string `json:"model"`
Usage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
}

if err := json.Unmarshal(raw, &payload); err != nil {
return
}

telemetry.UpstreamModel = payload.Model
telemetry.PromptTokens = payload.Usage.PromptTokens
telemetry.CompletionTokens = payload.Usage.CompletionTokens
telemetry.TotalTokens = payload.Usage.TotalTokens
}

func requestModel(r *http.Request) (string, []byte, error) {
if r.Body == nil {
return "", nil, fmt.Errorf("request body is required")
Expand Down
68 changes: 68 additions & 0 deletions internal/server/server_test.go
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
package server

import (
"bytes"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"os"
Expand Down Expand Up @@ -327,6 +329,72 @@ func TestEnsureCommandFailureReturnsServiceUnavailable(t *testing.T) {
}
}

func TestBackendResponseTelemetryLogsUsage(t *testing.T) {
var logs bytes.Buffer
originalLogger := slog.Default()
slog.SetDefault(slog.New(slog.NewJSONHandler(&logs, nil)))
t.Cleanup(func() {
slog.SetDefault(originalLogger)
})

backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var payload struct {
Model string `json:"model"`
}
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
t.Errorf("decode backend request: %v", err)
http.Error(w, "bad request", http.StatusBadRequest)
return
}
writeJSON(w, http.StatusOK, map[string]any{
"id": "chatcmpl-test",
"object": "chat.completion",
"created": time.Now().Unix(),
"model": payload.Model,
"choices": []map[string]any{{
"index": 0,
"message": map[string]string{
"role": "assistant",
"content": "ok",
},
"finish_reason": "stop",
}},
"usage": map[string]int{
"prompt_tokens": 51,
"completion_tokens": 4,
"total_tokens": 55,
},
})
}))
t.Cleanup(backend.Close)

srv := testServerWithBackend(t, backend.URL, config.ModelConfig{
ID: "local-coder",
Backend: "lmstudio",
TargetModel: "target-model",
})

if status := serveChat(t, srv, "local-coder"); status != http.StatusOK {
t.Fatalf("unexpected status: %d", status)
}

logText := logs.String()
for _, want := range []string{
`"msg":"backend response completed"`,
`"alias":"local-coder"`,
`"target_model":"target-model"`,
`"upstream_model":"target-model"`,
`"status":200`,
`"prompt_tokens":51`,
`"completion_tokens":4`,
`"total_tokens":55`,
} {
if !strings.Contains(logText, want) {
t.Fatalf("expected log to contain %s, got logs:\n%s", want, logText)
}
}
}

func testServer(t *testing.T) *Server {
t.Helper()

Expand Down
5 changes: 5 additions & 0 deletions test/mock-openai-backend/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,11 @@ func main() {
},
"finish_reason": "stop",
}},
"usage": map[string]int{
"prompt_tokens": 12,
"completion_tokens": 4,
"total_tokens": 16,
},
})
})

Expand Down
Loading