iop/apps/edge/internal/openai/chat_handler.go
toki 957700c2d6 feat: model route queue policy alignment phase added and edge service updates
- Add model-route-queue-policy-alignment milestone to inference provider extension phase
- Add SDD documentation for inference provider extension
- Update edge chat/responses handlers for OpenAI compatible API
- Update edge config with new queue policy settings
- Add config tests for queue policy support
- Update task tracking for model route queue policy alignment
2026-06-16 22:30:36 +09:00

329 lines
9.5 KiB
Go

package openai
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"path/filepath"
"strings"
"time"
"go.uber.org/zap"
edgeservice "iop/apps/edge/internal/service"
"iop/packages/go/config"
)
func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
return
}
defer r.Body.Close()
var req chatCompletionRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request_error", "invalid JSON request")
return
}
runMeta, inferenceTarget, workspace, err := parseOpenAIMetadata(req.Metadata)
if err != nil {
writeError(w, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
dispatch, ok := s.resolveRouteDispatch(req.Model, inferenceTarget)
if !ok {
writeError(w, http.StatusBadRequest, "invalid_request_error", "model is required")
return
}
if err := validateWorkspaceForRoute(dispatch, workspace); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
basePrompt := promptFromMessages(req.Messages)
if strings.TrimSpace(basePrompt) == "" {
writeError(w, http.StatusBadRequest, "invalid_request_error", "messages are required")
return
}
outputPolicy := s.resolveOutputPolicy(basePrompt)
messages := req.Messages
if instruction := strictOutputContractInstruction(outputPolicy); instruction != "" {
messages = prependSystemMessage(messages, instruction)
}
prompt := promptFromMessages(messages)
input := req.runInput(prompt, messages, outputPolicy.Strict)
s.logger.Info("openai chat completion input",
zap.String("model", req.Model),
zap.String("target", dispatch.Target),
zap.String("adapter", dispatch.Adapter),
zap.Bool("strict_output", outputPolicy.Strict),
zap.Bool("strict_stream_buffer", outputPolicy.StreamBuffer),
zap.String("xml_completion_tool", outputPolicy.XMLCompletionTool),
zap.Bool("contract_instruction", outputPolicy.ContractInstruction),
zap.Bool("stream", req.Stream),
zap.Int("message_count", len(req.Messages)),
zap.Int("prompt_len", len(prompt)),
zap.String("prompt_preview", previewString(prompt, 1000)),
zap.Any("input_keys", mapKeys(input)),
)
handle, err := s.service.SubmitRun(r.Context(), edgeservice.SubmitRunRequest{
NodeRef: dispatch.NodeRef,
ModelGroupKey: strings.TrimSpace(req.Model),
Adapter: dispatch.Adapter,
Target: dispatch.Target,
SessionID: dispatch.SessionID,
Workspace: workspace,
Prompt: prompt,
Input: input,
TimeoutSec: dispatch.TimeoutSec,
MaxQueue: dispatch.MaxQueue,
QueueTimeoutMS: dispatch.QueueTimeoutMS,
Metadata: chatRunMetadata(runMeta, req, outputPolicy),
})
if err != nil {
writeError(w, http.StatusBadGateway, "node_dispatch_error", err.Error())
return
}
defer handle.Close()
if req.Stream {
s.streamChatCompletion(w, r, req, handle, outputPolicy)
return
}
s.completeChatCompletion(w, r, req, handle, outputPolicy)
}
func chatRunMetadata(runMeta map[string]string, req chatCompletionRequest, outputPolicy strictOutputPolicy) map[string]string {
if runMeta == nil {
runMeta = make(map[string]string)
}
runMeta["source"] = "openai"
runMeta["openai_model"] = req.Model
runMeta["openai_stream"] = fmt.Sprintf("%t", req.Stream)
runMeta["strict_output"] = fmt.Sprintf("%t", outputPolicy.Strict)
return runMeta
}
func (s *Server) completeChatCompletion(w http.ResponseWriter, r *http.Request, req chatCompletionRequest, handle edgeservice.RunResult, outputPolicy strictOutputPolicy) {
text, reasoning, usage, err := collectRunResult(r.Context(), handle.Stream(), handle.WaitTimeout())
if err != nil {
writeError(w, httpStatusForRunError(err), "run_error", err.Error())
return
}
text, reasoning, normalized := normalizeCompletionOutput(outputPolicy, text, reasoning)
s.logger.Info("openai chat completion output",
zap.String("run_id", handle.Dispatch().RunID),
zap.Bool("strict_output", outputPolicy.Strict),
zap.Bool("strict_stream_buffer", outputPolicy.StreamBuffer),
zap.String("xml_completion_tool", outputPolicy.XMLCompletionTool),
zap.Bool("normalized", normalized),
zap.Int("content_len", len(text)),
zap.Int("reasoning_len", len(reasoning)),
zap.String("content_preview", previewString(text, 1000)),
)
created := time.Now().Unix()
writeJSON(w, http.StatusOK, chatCompletionResponse{
ID: "chatcmpl-" + handle.Dispatch().RunID,
Object: "chat.completion",
Created: created,
Model: responseModel(req.Model, handle.Dispatch().Target),
Choices: []chatCompletionChoice{{
Index: 0,
Message: chatMessage{Role: "assistant", Content: text, ReasoningContent: reasoning},
FinishReason: "stop",
}},
Usage: usage,
})
}
func (s *Server) resolveAdapter() string {
if s.cfg.Adapter != "" {
return s.cfg.Adapter
}
return "ollama"
}
func (s *Server) resolveTarget(model string) string {
return s.resolveTargetWithOverride(model, "")
}
func (s *Server) resolveTargetWithOverride(model, override string) string {
if s.cfg.Target != "" {
return s.cfg.Target
}
if override != "" {
return strings.TrimSpace(override)
}
return strings.TrimSpace(model)
}
// routeDispatch holds fully-resolved dispatch parameters for a single request.
type routeDispatch struct {
NodeRef string
Adapter string
Target string
SessionID string
TimeoutSec int
MaxQueue int
QueueTimeoutMS int
WorkspaceRequired bool
}
// resolveRoute returns the first catalog entry whose Model matches model.
// Entries with an empty Target are skipped.
func (s *Server) resolveRoute(model string) *config.OpenAIRouteEntry {
model = strings.TrimSpace(model)
if model == "" {
return nil
}
for i := range s.cfg.ModelRoutes {
r := &s.cfg.ModelRoutes[i]
if strings.TrimSpace(r.Model) == model && r.Target != "" {
return r
}
}
return nil
}
// resolveRouteDispatch returns fully-resolved dispatch params for model.
// Route catalog entries take priority; metadataTarget is only used in the
// legacy fallback path (when no catalog entry matches).
// Returns (dispatch, true) on success; (zero, false) when no target can be resolved.
func (s *Server) resolveRouteDispatch(model, metadataTarget string) (routeDispatch, bool) {
if route := s.resolveRoute(model); route != nil {
adapter := route.Adapter
if adapter == "" {
adapter = s.resolveAdapter()
}
nodeRef := route.NodeRef
if nodeRef == "" {
nodeRef = s.cfg.NodeRef
}
sessionID := route.SessionID
if sessionID == "" {
sessionID = s.resolveSessionID()
}
timeoutSec := route.TimeoutSec
if timeoutSec <= 0 {
timeoutSec = s.resolveTimeoutSec()
}
return routeDispatch{
NodeRef: nodeRef,
Adapter: adapter,
Target: route.Target,
SessionID: sessionID,
TimeoutSec: timeoutSec,
MaxQueue: route.MaxQueue,
QueueTimeoutMS: route.QueueTimeoutMS,
WorkspaceRequired: route.WorkspaceRequired,
}, true
}
target := s.resolveTargetWithOverride(model, metadataTarget)
if target == "" {
return routeDispatch{}, false
}
return routeDispatch{
NodeRef: s.cfg.NodeRef,
Adapter: s.resolveAdapter(),
Target: target,
SessionID: s.resolveSessionID(),
TimeoutSec: s.resolveTimeoutSec(),
}, true
}
func (s *Server) resolveSessionID() string {
if s.cfg.SessionID != "" {
return s.cfg.SessionID
}
return edgeservice.DefaultSessionID
}
func (s *Server) resolveTimeoutSec() int {
if s.cfg.TimeoutSec > 0 {
return s.cfg.TimeoutSec
}
return edgeservice.DefaultTimeoutSec
}
func (s *Server) resolveStrictOutput() bool {
return s.cfg.StrictOutput
}
func (s *Server) resolveStrictStreamBuffer() bool {
return s.cfg.StrictStreamBuffer
}
func (s *Server) resolveOutputPolicy(prompt string) strictOutputPolicy {
policy := strictOutputPolicy{
Strict: s.resolveStrictOutput(),
StreamBuffer: s.resolveStrictStreamBuffer(),
}
if !policy.Strict {
return policy
}
policy.XMLCompletionTool, policy.XMLResultTag = inferXMLCompletionContract(prompt)
policy.ContractInstruction = policy.XMLCompletionTool != ""
return policy
}
func promptFromMessages(messages []chatMessage) string {
var b strings.Builder
for _, msg := range messages {
content := strings.TrimSpace(msg.Content)
if content == "" {
continue
}
role := strings.TrimSpace(msg.Role)
if role == "" {
role = "user"
}
if b.Len() > 0 {
b.WriteString("\n")
}
b.WriteString(role)
b.WriteString(": ")
b.WriteString(content)
}
return b.String()
}
func responseModel(requestModel, target string) string {
if requestModel != "" {
return requestModel
}
return target
}
func httpStatusForRunError(err error) int {
if errors.Is(err, context.Canceled) {
return http.StatusRequestTimeout
}
return http.StatusBadGateway
}
func validateWorkspaceForRoute(d routeDispatch, workspace string) error {
if !d.WorkspaceRequired {
return nil
}
if strings.TrimSpace(workspace) == "" {
return fmt.Errorf("workspace is required for this model route")
}
if !filepath.IsAbs(workspace) {
return fmt.Errorf("workspace must be an absolute path")
}
return nil
}
func writeJSON(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
func writeError(w http.ResponseWriter, status int, code, message string) {
writeJSON(w, status, errorResponse{Error: errorBody{Type: code, Message: message}})
}