summaryrefslogtreecommitdiff
path: root/internal/httpapi/api.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/httpapi/api.go62
1 files changed, 59 insertions, 3 deletions
diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go
index e66b170..ea4c0c2 100644
--- a/internal/httpapi/api.go
+++ b/internal/httpapi/api.go
@@ -150,6 +150,8 @@ func (a *API) registerInference(mux *http.ServeMux) {
mux.HandleFunc("POST /api/v1/chat/completions", a.openAIChat)
mux.HandleFunc("POST /v1/responses", a.openAIResponses)
mux.HandleFunc("POST /api/v1/responses", a.openAIResponses)
+ mux.HandleFunc("POST /v1/embeddings", a.openAIEmbeddings)
+ mux.HandleFunc("POST /api/v1/embeddings", a.openAIEmbeddings)
mux.HandleFunc("GET /anthropic/v1/models", a.anthropicModels)
mux.HandleFunc("GET /api/anthropic/v1/models", a.anthropicModels)
@@ -171,6 +173,10 @@ func (a *API) openAIResponses(w http.ResponseWriter, r *http.Request) {
a.serveInference(w, r, domain.ProtocolOpenAIResponses)
}
+func (a *API) openAIEmbeddings(w http.ResponseWriter, r *http.Request) {
+ a.serveInference(w, r, domain.ProtocolOpenAIEmbeddings)
+}
+
func (a *API) anthropicMessages(w http.ResponseWriter, r *http.Request) {
a.serveInference(w, r, domain.ProtocolAnthropic)
}
@@ -243,6 +249,16 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
return
}
publicModel := model.ID
+ if protocol == domain.ProtocolOpenAIEmbeddings {
+ if envelope.Stream {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Embeddings does not support streaming"}, requestID)
+ return
+ }
+ if !modelDeclaresCapability(model, "embeddings") {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support embeddings"}, requestID)
+ return
+ }
+ }
requestedOutput := max(envelope.MaxTokens, envelope.MaxCompletionTokens, envelope.MaxOutputTokens)
if model.MaxOutputTokens > 0 && requestedOutput > model.MaxOutputTokens {
apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Requested output exceeds the model maximum"}, requestID)
@@ -284,7 +300,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
policy, _ = a.limiter.Policy(principal.ProjectID)
}
if err := a.billingMeter.Authorize(r.Context(), billing.Authorization{
- RequestID: requestID, Principal: principal, Model: model, Body: body, Policy: policy,
+ RequestID: requestID, Principal: principal, Model: model, Protocol: protocol, Body: body, Policy: policy,
}); err != nil {
if errors.Is(err, billing.ErrInsufficientBalance) {
apierror.Write(w, apierror.Error{Status: http.StatusPaymentRequired, Type: "insufficient_balance", Message: "Account balance is insufficient"}, requestID)
@@ -293,6 +309,13 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
ErrorType: "insufficient_balance", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()})
return
}
+ if errors.Is(err, billing.ErrDailyQuotaExceeded) {
+ apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "daily_quota_exceeded", Message: "API key daily spend quota exceeded"}, requestID)
+ a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID,
+ PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusTooManyRequests, Success: false,
+ ErrorType: "daily_quota_exceeded", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()})
+ return
+ }
if errors.Is(err, billing.ErrQuotaExceeded) {
apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "monthly_quota_exceeded", Message: "Project monthly spend quota exceeded"}, requestID)
a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID,
@@ -341,6 +364,16 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
})
return
}
+ if protocol == domain.ProtocolOpenAIEmbeddings && !strings.HasPrefix(strings.ToLower(result.Response.Header.Get("Content-Type")), "application/json") {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadGateway, Type: "invalid_provider_response", Message: "Upstream provider returned a non-JSON Embeddings response"}, requestID)
+ a.finishUsage(r, domain.UsageEvent{
+ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID,
+ PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel,
+ Protocol: protocol, StatusCode: http.StatusBadGateway, Success: false, ErrorType: "invalid_provider_response",
+ Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(),
+ })
+ return
+ }
stream := envelope.Stream || strings.HasPrefix(strings.ToLower(result.Response.Header.Get("Content-Type")), "text/event-stream")
copyResponseHeaders(w.Header(), result.Response.Header, stream)
@@ -354,6 +387,18 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
copyErr := copyResponse(w, result.Response.Body, observer, stream)
usageResult := observer.Usage()
usageReported := observer.Reported()
+ ttftMS := int64(0)
+ if firstOutputAt := observer.FirstOutputAt(); !firstOutputAt.IsZero() {
+ ttftMS = firstOutputAt.Sub(startedAt).Milliseconds()
+ if ttftMS < 1 {
+ ttftMS = 1
+ }
+ providerStartedAt := result.AttemptStartedAt
+ if providerStartedAt.IsZero() {
+ providerStartedAt = startedAt
+ }
+ a.forwarder.ObserveTTFT(model.ID, result.Route, firstOutputAt.Sub(providerStartedAt))
+ }
success = copyErr == nil
errorType := ""
if copyErr != nil {
@@ -363,7 +408,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID,
PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel,
Protocol: protocol, Stream: stream, StatusCode: result.Response.StatusCode, Success: success,
- ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), Usage: usageResult,
+ ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), TTFTMS: ttftMS, Usage: usageResult,
UsageReported: usageReported,
})
a.logger.Info("inference_request",
@@ -376,6 +421,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
"status", result.Response.StatusCode,
"attempts", result.Attempts,
"duration_ms", time.Since(startedAt).Milliseconds(),
+ "ttft_ms", ttftMS,
)
}
@@ -428,6 +474,7 @@ func modelProviderDescriptors(model domain.Model) []map[string]string {
func (a *API) availableOpenAIModels(principal domain.Principal) []domain.Model {
combined := append(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal), a.catalog.ModelsFor(domain.ProtocolOpenAIResponses, principal)...)
+ combined = append(combined, a.catalog.ModelsFor(domain.ProtocolOpenAIEmbeddings, principal)...)
seen := make(map[string]struct{}, len(combined))
result := make([]domain.Model, 0, len(combined))
for _, model := range combined {
@@ -527,6 +574,15 @@ func modelHasCapability(model domain.Model, wanted string) bool {
}
return false
}
+
+func modelDeclaresCapability(model domain.Model, wanted string) bool {
+ for _, value := range model.Capabilities {
+ if value == wanted || value == "*" {
+ return true
+ }
+ }
+ return false
+}
func modelCreated(model domain.Model) int64 {
if model.ReleasedAt != nil {
return model.ReleasedAt.Unix()
@@ -653,7 +709,7 @@ func copyResponse(w http.ResponseWriter, body io.Reader, observer io.Writer, str
if stream {
destination = &flushingWriter{writer: w, controller: http.NewResponseController(w)}
}
- _, err := io.CopyBuffer(io.MultiWriter(destination, observer), body, make([]byte, 32<<10))
+ _, err := io.CopyBuffer(io.MultiWriter(observer, destination), body, make([]byte, 32<<10))
return err
}