Add support for multiple providers

This commit is contained in:
dwrz
2026-06-19 14:31:41 +00:00
parent c71300b1bf
commit f5e72b1c7a
12 changed files with 948 additions and 869 deletions

View File

@@ -1,5 +1,5 @@
// Package service orchestrates the odidere voice assistant server.
// It coordinates the HTTP server, LLM client, and handles
// It coordinates the HTTP server, LLM clients, and handles
// graceful shutdown.
package service
@@ -16,6 +16,7 @@ import (
"os/signal"
"runtime/debug"
"strings"
"sync"
"syscall"
"time"
@@ -37,6 +38,8 @@ const (
logKey key = "log"
idKey key = "id"
ipKey key = "ip"
modelsTimeout = 10 * time.Second
)
//go:embed all:static/*
@@ -45,15 +48,16 @@ var static embed.FS
// Service is the main application coordinator.
// It owns the HTTP server and all processing clients.
type Service struct {
cfg *config.Config
cron *cron.Cron
jobs []*job.Job
llm *llm.Client
log *slog.Logger
mux *http.ServeMux
server *http.Server
tmpl *template.Template
tools *tool.Registry
cfg *config.Config
cron *cron.Cron
jobs []*job.Job
llms map[string]*llm.Client
log *slog.Logger
allowedModels map[string]map[string]struct{}
mux *http.ServeMux
server *http.Server
tmpl *template.Template
tools *tool.Registry
}
// New creates a Service from the provided configuration.
@@ -72,19 +76,80 @@ func New(cfg *config.Config, log *slog.Logger) (*Service, error) {
}
svc.tools = registry
// Create LLM client.
llmClient, err := llm.NewClient(cfg.LLM, registry, log)
if err != nil {
return nil, fmt.Errorf("create LLM client: %v", err)
// Create LLM clients for each provider.
svc.llms = make(map[string]*llm.Client, len(cfg.Providers))
for _, pc := range cfg.Providers {
client, err := llm.NewClient(
llm.Config{
Key: pc.Key,
SystemMessage: cfg.SystemMessage,
Timeout: pc.Timeout,
URL: pc.URL,
},
registry,
log,
)
if err != nil {
return nil, fmt.Errorf(
"create LLM client for provider %q: %w",
pc.Name, err,
)
}
svc.llms[pc.Name] = client
log.Info(
"LLM provider registered",
slog.String("provider", pc.Name),
slog.String("url", pc.URL),
)
}
// Build model allowed list map: provider name -> set of base model IDs.
svc.allowedModels = make(
map[string]map[string]struct{}, len(cfg.Providers),
)
for _, pc := range cfg.Providers {
if len(pc.Models) == 0 {
continue
}
allowed := make(map[string]struct{}, len(pc.Models))
for _, m := range pc.Models {
base, _, _ := strings.Cut(m, ":")
allowed[base] = struct{}{}
}
svc.allowedModels[pc.Name] = allowed
}
svc.llm = llmClient
// Convert job configs to jobs.
jobs := make([]*job.Job, 0, len(cfg.Jobs))
for _, jc := range cfg.Jobs {
j, err := job.New(jc, log, svc.llm)
// Resolve the provider and model for this job.
var provider, modelID string
if jc.Model.Provider != "" {
if jc.Model.Model == "" {
return nil, fmt.Errorf(
"job %q: model.provider set but model.model is empty",
jc.Name,
)
}
provider = jc.Model.Provider
modelID = jc.Model.Model
} else {
provider = svc.cfg.DefaultModel.Provider
modelID = svc.cfg.DefaultModel.Model
}
llmc := svc.llms[provider]
if llmc == nil {
return nil, fmt.Errorf(
"create job %q: no LLM client for provider %q",
jc.Name,
provider,
)
}
j, err := job.New(jc, log, llmc, modelID)
if err != nil {
return nil, fmt.Errorf("create job %q: %w", jc.Name, err)
return nil, fmt.Errorf(
"create job %q: %w", jc.Name, err,
)
}
jobs = append(jobs, j)
}
@@ -132,8 +197,9 @@ func (svc *Service) ServeHTTP(w http.ResponseWriter, r *http.Request) {
id = uuid.NewString()
ip = func() string {
if ip := r.Header.Get("X-Forwarded-For"); ip != "" {
if idx := strings.Index(ip, ","); idx != -1 {
return strings.TrimSpace(ip[:idx])
before, _, found := strings.Cut(ip, ",")
if found {
return strings.TrimSpace(before)
}
return ip
}
@@ -196,8 +262,14 @@ func (svc *Service) Run(ctx context.Context) error {
"starting odidere",
slog.Group(
"llm",
slog.String("url", svc.cfg.LLM.URL),
slog.String("model", svc.cfg.LLM.Model),
slog.String(
"default_provider",
svc.cfg.DefaultModel.Provider,
),
slog.String(
"default_model", svc.cfg.DefaultModel.Model,
),
slog.Int("providers", len(svc.cfg.Providers)),
),
slog.Group(
"server",
@@ -237,10 +309,6 @@ func (svc *Service) Run(ctx context.Context) error {
// Register jobs with cron scheduler.
svc.cron = cron.New()
for _, j := range svc.jobs {
if !j.IsEnabled() {
svc.log.Info("job disabled", slog.String("name", j.Name()))
continue
}
entry, err := svc.cron.AddFunc(j.Schedule(), func() {
result := j.Run(ctx)
if result.Error != nil {
@@ -374,6 +442,8 @@ func (svc *Service) status(w http.ResponseWriter, r *http.Request) {
type Request struct {
// Messages is the conversation history.
Messages []openai.ChatCompletionMessage `json:"messages"`
// Provider is the LLM provider name. If empty, the default provider is used.
Provider string `json:"provider,omitempty"`
// Model is the LLM model ID. If empty, the default model is used.
Model string `json:"model,omitempty"`
// SystemMessage overrides the configured system message for this request.
@@ -387,6 +457,8 @@ type Response struct {
// Messages is the full list of messages generated during the query,
// including tool calls and tool results.
Messages []openai.ChatCompletionMessage `json:"messages,omitempty"`
// Provider is the LLM provider used for the response.
Provider string `json:"used_provider,omitempty"`
// Model is the LLM model used for the response.
Model string `json:"used_model,omitempty"`
// Voice is the voice used for TTS synthesis.
@@ -422,12 +494,53 @@ func (svc *Service) chat(w http.ResponseWriter, r *http.Request) {
slog.Any("data", req.Messages),
)
// Get LLM response.
var model = req.Model
if model == "" {
model = svc.llm.DefaultModel()
// Resolve provider and model.
provider := req.Provider
if provider == "" {
provider = svc.cfg.DefaultModel.Provider
}
msgs, err := svc.llm.Query(ctx, req.Messages, model, req.SystemMessage, 0)
llmc := svc.llms[provider]
if llmc == nil {
log.ErrorContext(
ctx,
"unknown provider",
slog.String("provider", provider),
)
http.Error(
w,
fmt.Sprintf("unknown provider: %q", provider),
http.StatusBadRequest,
)
return
}
model := req.Model
if model == "" {
model = svc.cfg.DefaultModel.Model
}
// Validate model against provider allowed list, if configured.
if allowed, ok := svc.allowedModels[provider]; ok {
m, _, _ := strings.Cut(model, ":")
if _, ok := allowed[m]; !ok {
log.ErrorContext(
ctx,
"model not in allowed list",
slog.String("provider", provider),
slog.String("model", model),
)
http.Error(
w,
fmt.Sprintf(
"model %q not allowed for provider %q",
model, provider,
),
http.StatusBadRequest,
)
return
}
}
msgs, err := llmc.Query(ctx, req.Messages, model, req.SystemMessage, 0)
if err != nil {
log.ErrorContext(
ctx,
@@ -450,12 +563,14 @@ func (svc *Service) chat(w http.ResponseWriter, r *http.Request) {
ctx,
"LLM response",
slog.String("text", final.Content),
slog.String("provider", provider),
slog.String("model", model),
)
w.Header().Set("Content-Type", "application/json")
if err := json.NewEncoder(w).Encode(Response{
Messages: msgs,
Provider: provider,
Model: model,
Voice: req.Voice,
}); err != nil {
@@ -473,6 +588,8 @@ type StreamMessage struct {
Error string `json:"error,omitempty"`
// Message is the chat completion message.
Message openai.ChatCompletionMessage `json:"message"`
// Provider is the LLM provider used for the response.
Provider string `json:"provider,omitempty"`
// Model is the LLM model used for the response.
Model string `json:"model,omitempty"`
// Voice is the voice used for TTS synthesis.
@@ -516,6 +633,52 @@ func (svc *Service) chatStream(w http.ResponseWriter, r *http.Request) {
return
}
// Resolve provider and model.
provider := req.Provider
if provider == "" {
provider = svc.cfg.DefaultModel.Provider
}
model := req.Model
if model == "" {
model = svc.cfg.DefaultModel.Model
}
llmc := svc.llms[provider]
if llmc == nil {
log.ErrorContext(
ctx,
"unknown provider",
slog.String("provider", provider),
)
http.Error(
w,
fmt.Sprintf("unknown provider: %q", provider),
http.StatusBadRequest,
)
return
}
// Validate model against provider allowed list, if configured.
if allowed, ok := svc.allowedModels[provider]; ok {
m, _, _ := strings.Cut(model, ":")
if _, ok := allowed[m]; !ok {
log.ErrorContext(
ctx,
"model not in allowed list",
slog.String("provider", provider),
slog.String("model", model),
)
http.Error(
w,
fmt.Sprintf(
"model %q not allowed for provider %q",
model, provider,
),
http.StatusBadRequest,
)
return
}
}
// Set SSE headers.
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
@@ -534,19 +697,15 @@ func (svc *Service) chatStream(w http.ResponseWriter, r *http.Request) {
flusher.Flush()
}
// Get model.
var model = req.Model
if model == "" {
model = svc.llm.DefaultModel()
}
// Start streaming LLM query.
var (
events = make(chan llm.StreamEvent)
llmErr error
)
go func() {
llmErr = svc.llm.QueryStream(ctx, req.Messages, model, req.SystemMessage, 0, events)
llmErr = llmc.QueryStream(
ctx, req.Messages, model, req.SystemMessage, 0, events,
)
}()
// Consume events and send as SSE.
@@ -578,46 +737,127 @@ func (svc *Service) chatStream(w http.ResponseWriter, r *http.Request) {
})
return
}
last.Provider = provider
last.Model = model
last.Voice = req.Voice
send(last)
}
// models returns available LLM models.
// ModelStatus represents the availability status of a model.
type ModelStatus string
const (
// ModelStatusAvailable indicates the model is available for use.
ModelStatusAvailable ModelStatus = "available"
// ModelStatusNotFound indicates the model is configured but not found
// in the provider's model list.
ModelStatusNotFound ModelStatus = "not_found"
// ModelStatusProviderError indicates the provider could not be reached
// to verify the model.
ModelStatusProviderError ModelStatus = "provider_error"
)
// Model represents a model in the /v1/models response.
type Model struct {
Model string `json:"model"`
Status ModelStatus `json:"status"`
}
// Models is the response format for the /v1/models endpoint.
type Models struct {
Providers map[string][]Model `json:"providers"`
DefaultProvider string `json:"default_provider"`
DefaultModel string `json:"default_model"`
}
// models returns available LLM models from all configured providers.
func (svc *Service) models(w http.ResponseWriter, r *http.Request) {
var (
ctx = r.Context()
log = ctx.Value(logKey).(*slog.Logger)
mu sync.Mutex
wg sync.WaitGroup
providers = make(map[string][]Model, len(svc.llms))
)
models, err := svc.llm.ListModels(ctx)
if err != nil {
log.ErrorContext(
ctx,
"failed to list models",
slog.Any("error", err),
)
http.Error(
w,
"failed to list models",
http.StatusInternalServerError,
)
return
cfg := make(map[string]config.ProviderConfig, len(svc.cfg.Providers))
for _, pc := range svc.cfg.Providers {
cfg[pc.Name] = pc
}
for name, client := range svc.llms {
wg.Add(1)
go func(name string, c *llm.Client, pc config.ProviderConfig) {
defer wg.Done()
var result = []Model{}
fetchCtx, cancel := context.WithTimeout(
ctx, modelsTimeout,
)
defer cancel()
models, err := c.ListModels(fetchCtx)
if err != nil {
log.WarnContext(
ctx,
"failed to list models for provider",
slog.String("provider", name),
slog.Any("error", err),
)
for _, m := range pc.Models {
result = append(result, Model{
Model: m,
Status: ModelStatusProviderError,
})
}
mu.Lock()
providers[name] = result
mu.Unlock()
return
}
available := make(map[string]struct{}, len(models))
for _, m := range models {
baseID, _, _ := strings.Cut(m.ID, ":")
available[baseID] = struct{}{}
}
var toCheck []string
if len(pc.Models) > 0 {
toCheck = pc.Models
} else {
toCheck = make([]string, 0, len(models))
for _, m := range models {
toCheck = append(toCheck, m.ID)
}
}
for _, m := range toCheck {
base, _, _ := strings.Cut(m, ":")
mi := Model{Model: m}
if _, ok := available[base]; ok {
mi.Status = ModelStatusAvailable
} else {
mi.Status = ModelStatusNotFound
}
result = append(result, mi)
}
mu.Lock()
providers[name] = result
mu.Unlock()
}(name, client, cfg[name])
}
wg.Wait()
w.Header().Set("Content-Type", "application/json")
if err := json.NewEncoder(w).Encode(struct {
Models []openai.Model `json:"models"`
DefaultModel string `json:"default_model"`
}{
Models: models,
DefaultModel: svc.llm.DefaultModel(),
if err := json.NewEncoder(w).Encode(Models{
Providers: providers,
DefaultProvider: svc.cfg.DefaultModel.Provider,
DefaultModel: svc.cfg.DefaultModel.Model,
}); err != nil {
log.ErrorContext(
ctx,
"failed to encode models response",
slog.Any("error", err),
)
log.ErrorContext(ctx, "failed to encode models response",
slog.Any("error", err))
}
}