iop/apps/edge/internal/openai/server.go

627 lines
16 KiB
Go

package openai
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
"go.uber.org/zap"
edgeservice "iop/apps/edge/internal/service"
"iop/packages/config"
)
type runService interface {
SubmitRun(context.Context, edgeservice.SubmitRunRequest) (*edgeservice.RunHandle, error)
OllamaAPI(context.Context, edgeservice.OllamaAPIRequest) (edgeservice.OllamaAPIView, error)
}
type Server struct {
cfg config.EdgeOpenAIConf
service runService
logger *zap.Logger
server *http.Server
}
func NewServer(cfg config.EdgeOpenAIConf, svc runService, logger *zap.Logger) *Server {
if logger == nil {
logger = zap.NewNop()
}
return &Server{cfg: cfg, service: svc, logger: logger}
}
func (s *Server) Enabled() bool {
return s != nil && s.cfg.Enabled
}
func (s *Server) Start(ctx context.Context) error {
if !s.Enabled() {
return nil
}
if s.cfg.Listen == "" {
s.cfg.Listen = "0.0.0.0:8080"
}
mux := http.NewServeMux()
mux.HandleFunc("/healthz", s.handleHealthz)
mux.HandleFunc("/v1/models", s.handleModels)
mux.HandleFunc("/v1/chat/completions", s.handleChatCompletions)
mux.HandleFunc("/api/", s.handleOllamaAPI)
s.server = &http.Server{
Addr: s.cfg.Listen,
Handler: mux,
ReadHeaderTimeout: 5 * time.Second,
}
ln, err := net.Listen("tcp", s.cfg.Listen)
if err != nil {
return fmt.Errorf("openai server listen %s: %w", s.cfg.Listen, err)
}
go func() {
<-ctx.Done()
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_ = s.server.Shutdown(shutdownCtx)
}()
go func() {
s.logger.Info("openai-compatible server listening", zap.String("addr", s.cfg.Listen))
if err := s.server.Serve(ln); err != nil && !errors.Is(err, http.ErrServerClosed) {
s.logger.Warn("openai-compatible server exited", zap.Error(err))
}
}()
return nil
}
func (s *Server) Stop(ctx context.Context) error {
if s == nil || s.server == nil {
return nil
}
return s.server.Shutdown(ctx)
}
func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
writeJSON(w, http.StatusOK, map[string]string{"status": "ok"})
}
func (s *Server) handleModels(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
return
}
models := s.cfg.Models
if len(models) == 0 && s.cfg.Target != "" {
models = []string{s.cfg.Target}
}
data := make([]openAIModel, 0, len(models))
for _, model := range models {
model = strings.TrimSpace(model)
if model == "" {
continue
}
data = append(data, openAIModel{
ID: model,
Object: "model",
Created: time.Now().Unix(),
OwnedBy: "iop",
})
}
writeJSON(w, http.StatusOK, openAIModelsResponse{Object: "list", Data: data})
}
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
}
target := s.resolveTarget(req.Model)
if target == "" {
writeError(w, http.StatusBadRequest, "invalid_request_error", "model is required")
return
}
prompt := promptFromMessages(req.Messages)
if strings.TrimSpace(prompt) == "" {
writeError(w, http.StatusBadRequest, "invalid_request_error", "messages are required")
return
}
handle, err := s.service.SubmitRun(r.Context(), edgeservice.SubmitRunRequest{
NodeRef: s.cfg.NodeRef,
Adapter: s.resolveAdapter(),
Target: target,
SessionID: s.resolveSessionID(),
Prompt: prompt,
Input: req.runInput(prompt),
TimeoutSec: s.resolveTimeoutSec(),
Metadata: map[string]string{
"source": "openai",
"openai_model": req.Model,
"openai_stream": fmt.Sprintf("%t", req.Stream),
},
})
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)
return
}
s.completeChatCompletion(w, r, req, handle)
}
func (s *Server) handleOllamaAPI(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodPost && r.Method != http.MethodDelete {
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
return
}
defer r.Body.Close()
body, err := io.ReadAll(io.LimitReader(r.Body, 16<<20))
if err != nil {
writeError(w, http.StatusBadRequest, "invalid_request_error", "read request body failed")
return
}
resp, err := s.service.OllamaAPI(r.Context(), edgeservice.OllamaAPIRequest{
NodeRef: s.cfg.NodeRef,
Adapter: s.resolveAdapter(),
Method: r.Method,
Path: r.URL.RequestURI(),
Body: string(body),
TimeoutSec: s.resolveTimeoutSec(),
})
if err != nil {
writeError(w, http.StatusBadGateway, "ollama_passthrough_error", err.Error())
return
}
contentType := resp.ContentType
if contentType == "" {
contentType = "application/json"
}
w.Header().Set("Content-Type", contentType)
status := resp.StatusCode
if status < 100 || status > 999 {
status = http.StatusOK
}
w.WriteHeader(status)
_, _ = w.Write([]byte(resp.Body))
}
func (s *Server) completeChatCompletion(w http.ResponseWriter, r *http.Request, req chatCompletionRequest, handle *edgeservice.RunHandle) {
text, reasoning, usage, err := collectRunResult(r.Context(), handle)
if err != nil {
writeError(w, httpStatusForRunError(err), "run_error", err.Error())
return
}
created := time.Now().Unix()
writeJSON(w, http.StatusOK, chatCompletionResponse{
ID: "chatcmpl-" + handle.RunID,
Object: "chat.completion",
Created: created,
Model: responseModel(req.Model, handle.Target),
Choices: []chatCompletionChoice{{
Index: 0,
Message: chatMessage{Role: "assistant", Content: text, ReasoningContent: reasoning},
FinishReason: "stop",
}},
Usage: usage,
})
}
func (s *Server) streamChatCompletion(w http.ResponseWriter, r *http.Request, req chatCompletionRequest, handle *edgeservice.RunHandle) {
flusher, ok := w.(http.Flusher)
if !ok {
writeError(w, http.StatusInternalServerError, "streaming_not_supported", "response writer does not support streaming")
return
}
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
created := time.Now().Unix()
id := "chatcmpl-" + handle.RunID
model := responseModel(req.Model, handle.Target)
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{Role: "assistant"},
}},
})
for {
select {
case <-r.Context().Done():
return
case nodeEvent := <-handle.NodeEvents:
if edgeservice.IsNodeDisconnected(nodeEvent) {
writeSSEError(w, flusher, "node disconnected")
return
}
case event := <-handle.Events:
if event == nil {
continue
}
switch event.GetType() {
case "delta":
if event.GetDelta() == "" {
continue
}
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{Content: event.GetDelta()},
}},
})
case "reasoning_delta":
if event.GetDelta() == "" {
continue
}
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{ReasoningContent: event.GetDelta()},
}},
})
case "complete":
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{},
FinishReason: "stop",
}},
})
fmt.Fprint(w, "data: [DONE]\n\n")
flusher.Flush()
return
case "error", "cancelled":
msg := event.GetError()
if msg == "" {
msg = event.GetMessage()
}
if msg == "" {
msg = "run failed"
}
writeSSEError(w, flusher, msg)
return
}
case <-time.After(handle.WaitTimeout()):
writeSSEError(w, flusher, "run timed out")
return
}
}
}
func collectRunResult(ctx context.Context, handle *edgeservice.RunHandle) (string, string, *openAIUsage, error) {
var contentBuilder strings.Builder
var reasoningBuilder strings.Builder
var usage *openAIUsage
timer := time.NewTimer(handle.WaitTimeout())
defer timer.Stop()
for {
select {
case <-ctx.Done():
return "", "", nil, ctx.Err()
case <-timer.C:
return "", "", nil, fmt.Errorf("run timed out")
case nodeEvent := <-handle.NodeEvents:
if edgeservice.IsNodeDisconnected(nodeEvent) {
return "", "", nil, fmt.Errorf("node disconnected")
}
case event := <-handle.Events:
if event == nil {
continue
}
switch event.GetType() {
case "delta":
contentBuilder.WriteString(event.GetDelta())
case "reasoning_delta":
reasoningBuilder.WriteString(event.GetDelta())
case "complete":
if u := event.GetUsage(); u != nil {
usage = &openAIUsage{
PromptTokens: int(u.GetInputTokens()),
CompletionTokens: int(u.GetOutputTokens()),
TotalTokens: int(u.GetInputTokens() + u.GetOutputTokens()),
}
}
return contentBuilder.String(), reasoningBuilder.String(), usage, nil
case "error", "cancelled":
msg := event.GetError()
if msg == "" {
msg = event.GetMessage()
}
if msg == "" {
msg = "run failed"
}
return "", "", nil, fmt.Errorf("%s", msg)
}
}
}
}
func (s *Server) resolveAdapter() string {
if s.cfg.Adapter != "" {
return s.cfg.Adapter
}
return "ollama"
}
func (s *Server) resolveTarget(model string) string {
if s.cfg.Target != "" {
return s.cfg.Target
}
return strings.TrimSpace(model)
}
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 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 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}})
}
func writeSSE(w http.ResponseWriter, flusher http.Flusher, v any) {
b, err := json.Marshal(v)
if err != nil {
return
}
fmt.Fprintf(w, "data: %s\n\n", b)
flusher.Flush()
}
func writeSSEError(w http.ResponseWriter, flusher http.Flusher, message string) {
writeSSE(w, flusher, errorResponse{Error: errorBody{Type: "run_error", Message: message}})
fmt.Fprint(w, "data: [DONE]\n\n")
flusher.Flush()
}
type chatCompletionRequest struct {
Model string `json:"model"`
Messages []chatMessage `json:"messages"`
Stream bool `json:"stream"`
Options map[string]any `json:"options,omitempty"`
Format any `json:"format,omitempty"`
KeepAlive any `json:"keep_alive,omitempty"`
Think any `json:"think,omitempty"`
Tools []any `json:"tools,omitempty"`
}
type chatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
ReasoningContent string `json:"reasoning_content,omitempty"`
Images []any `json:"images,omitempty"`
ToolCalls []any `json:"tool_calls,omitempty"`
ToolName string `json:"tool_name,omitempty"`
}
func (m *chatMessage) UnmarshalJSON(b []byte) error {
var raw struct {
Role string `json:"role"`
Content any `json:"content"`
ReasoningContent string `json:"reasoning_content,omitempty"`
Images []any `json:"images,omitempty"`
ToolCalls []any `json:"tool_calls,omitempty"`
ToolName string `json:"tool_name,omitempty"`
}
if err := json.Unmarshal(b, &raw); err != nil {
return err
}
m.Role = raw.Role
m.Content = contentToString(raw.Content)
m.ReasoningContent = raw.ReasoningContent
m.Images = raw.Images
m.ToolCalls = raw.ToolCalls
m.ToolName = raw.ToolName
return nil
}
func (req chatCompletionRequest) runInput(prompt string) map[string]any {
input := map[string]any{
"prompt": prompt,
"messages": chatMessagesInput(req.Messages),
}
if len(req.Options) > 0 {
input["options"] = req.Options
}
if req.Format != nil {
input["format"] = req.Format
}
if req.KeepAlive != nil {
input["keep_alive"] = req.KeepAlive
}
if req.Think != nil {
input["think"] = req.Think
}
if len(req.Tools) > 0 {
input["tools"] = req.Tools
}
return input
}
func chatMessagesInput(messages []chatMessage) []any {
out := make([]any, 0, len(messages))
for _, msg := range messages {
m := map[string]any{
"role": msg.Role,
"content": msg.Content,
}
if msg.ReasoningContent != "" {
m["thinking"] = msg.ReasoningContent
}
if len(msg.Images) > 0 {
m["images"] = msg.Images
}
if len(msg.ToolCalls) > 0 {
m["tool_calls"] = msg.ToolCalls
}
if msg.ToolName != "" {
m["tool_name"] = msg.ToolName
}
out = append(out, m)
}
return out
}
func contentToString(v any) string {
switch t := v.(type) {
case string:
return t
case []any:
parts := make([]string, 0, len(t))
for _, item := range t {
if m, ok := item.(map[string]any); ok {
if text, ok := m["text"].(string); ok {
parts = append(parts, text)
}
}
}
return strings.Join(parts, "\n")
default:
b, _ := json.Marshal(t)
return string(b)
}
}
type chatCompletionResponse struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []chatCompletionChoice `json:"choices"`
Usage *openAIUsage `json:"usage,omitempty"`
}
type chatCompletionChoice struct {
Index int `json:"index"`
Message chatMessage `json:"message"`
FinishReason string `json:"finish_reason"`
}
type chatCompletionChunk struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []chatCompletionChunkChoice `json:"choices"`
}
type chatCompletionChunkChoice struct {
Index int `json:"index"`
Delta chatDelta `json:"delta"`
FinishReason string `json:"finish_reason,omitempty"`
}
type chatDelta struct {
Role string `json:"role,omitempty"`
Content string `json:"content,omitempty"`
ReasoningContent string `json:"reasoning_content,omitempty"`
}
type openAIUsage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
}
type openAIModel struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
OwnedBy string `json:"owned_by"`
}
type openAIModelsResponse struct {
Object string `json:"object"`
Data []openAIModel `json:"data"`
}
type errorResponse struct {
Error errorBody `json:"error"`
}
type errorBody struct {
Type string `json:"type"`
Message string `json:"message"`
}