Files
ModelRouter/internal/config/config_test.go

249 lines
6.6 KiB
Go

package config
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestLoadAndApplyDefaults(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "cfg.yaml")
content := `
listen: 127.0.0.1:9999
gateway_keys: [sk-1]
default_model: AUTO
adapter_dir: adapters
runtime_file: runtime.json
sources:
- name: deepseek
base_url: https://api.deepseek.com
api_key: sk-d
adapter: deepseek
models:
- id: deepseek-v4-flash
priority: 100
- name: ollama
base_url: http://127.0.0.1:11434
adapter: ollama
endpoint: /api/chat
models:
- id: llama3
priority: 50
`
if err := os.WriteFile(path, []byte(content), 0644); err != nil {
t.Fatal(err)
}
cfg, err := Load(path)
if err != nil {
t.Fatalf("load: %v", err)
}
if len(cfg.Sources) != 2 {
t.Fatalf("sources = %d", len(cfg.Sources))
}
if cfg.Sources[1].Endpoint != "/api/chat" {
t.Fatalf("endpoint = %q", cfg.Sources[1].Endpoint)
}
if cfg.Sources[1].Timeout == 0 {
t.Fatal("default timeout not applied")
}
if cfg.Sources[1].MaxConcurrent == 0 {
t.Fatal("default max_concurrent not applied")
}
if cfg.Sources[0].Models[0].Priority != 100 {
t.Fatalf("priority = %d", cfg.Sources[0].Models[0].Priority)
}
if cfg.DefaultModel != "AUTO" {
t.Fatalf("default model = %q", cfg.DefaultModel)
}
}
func TestApplyDefaultsDuplicateSource(t *testing.T) {
cfg := Config{Sources: []Source{
{Name: "a", BaseURL: "http://x", Models: []Model{{ID: "m1"}}},
{Name: "a", BaseURL: "http://y", Models: []Model{{ID: "m2"}}},
}}
if err := cfg.ApplyDefaults(); err == nil {
t.Fatal("expected duplicate source error")
}
}
func TestApplyDefaultsNoSources(t *testing.T) {
// empty source list is valid (sources may be added later via the Web UI)
cfg := Config{}
if err := cfg.ApplyDefaults(); err != nil {
t.Fatal(err)
}
if cfg.Listen != ":8080" || cfg.DefaultModel != "AUTO" {
t.Fatalf("defaults not applied: %+v", cfg)
}
}
func TestApplyDefaultsDuplicateModel(t *testing.T) {
cfg := Config{Sources: []Source{
{Name: "a", BaseURL: "http://x", Models: []Model{{ID: "m1"}}},
{Name: "b", BaseURL: "http://y", Models: []Model{{ID: "m1"}}},
}}
if err := cfg.ApplyDefaults(); err == nil {
t.Fatal("expected duplicate model error")
}
}
func TestStoreUpsertRemove(t *testing.T) {
path := filepath.Join(t.TempDir(), "runtime.json")
s := NewStore(path)
if err := s.Load(); err != nil {
t.Fatal(err)
}
if err := s.Upsert(Source{Name: "a", BaseURL: "http://a", APIKey: "sk-a", Models: []Model{{ID: "m"}}}); err != nil {
t.Fatal(err)
}
if err := s.Upsert(Source{Name: "b", BaseURL: "http://b", Models: []Model{{ID: "m2"}}}); err != nil {
t.Fatal(err)
}
if len(s.List()) != 2 {
t.Fatalf("list = %d", len(s.List()))
}
removed, err := s.Remove("a")
if err != nil || !removed {
t.Fatalf("remove: %v %v", removed, err)
}
if len(s.List()) != 1 {
t.Fatalf("after remove list = %d", len(s.List()))
}
// reload from disk
s2 := NewStore(path)
if err := s2.Load(); err != nil {
t.Fatal(err)
}
if len(s2.List()) != 1 {
t.Fatalf("reloaded list = %d", len(s2.List()))
}
}
func TestStoreSecretEncryption(t *testing.T) {
t.Setenv("LLMS_PROXY_MASTER_KEY", "")
dir := t.TempDir()
path := filepath.Join(dir, "runtime.json")
s := NewStore(path)
if s.box == nil {
t.Fatal("expected secret box")
}
if err := s.Load(); err != nil {
t.Fatal(err)
}
headers := map[string]string{"Authorization": "Bearer sk-hdr", "X-Custom": "plain"}
if err := s.Upsert(Source{Name: "a", BaseURL: "http://a", APIKey: "sk-secret-123", Headers: headers}); err != nil {
t.Fatal(err)
}
if err := s.SaveKey(GWKey{Key: "gw-secret", Role: "admin"}); err != nil {
t.Fatal(err)
}
// file on disk must not contain plaintext secrets
raw, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
for _, plain := range []string{"sk-secret-123", "Bearer sk-hdr", "gw-secret"} {
if strings.Contains(string(raw), plain) {
t.Fatalf("secret %q stored in plaintext on disk", plain)
}
}
// in-memory stays plaintext after the writes
src := s.List()[0]
if src.APIKey != "sk-secret-123" {
t.Fatalf("in-memory api_key = %q", src.APIKey)
}
if src.Headers["Authorization"] != "Bearer sk-hdr" {
t.Fatal("in-memory header not plaintext")
}
// reload: decrypted back
s2 := NewStore(path)
if err := s2.Load(); err != nil {
t.Fatal(err)
}
if s2.List()[0].APIKey != "sk-secret-123" {
t.Fatalf("reloaded api_key = %q", s2.List()[0].APIKey)
}
if k, ok := s2.KeyByValue("gw-secret"); !ok || k.Role != "admin" {
t.Fatalf("reloaded key lookup failed: %+v %v", k, ok)
}
}
func TestSecretBoxRoundTrip(t *testing.T) {
dir := t.TempDir()
box, err := NewSecretBox(filepath.Join(dir, "runtime.json"))
if err != nil {
t.Fatal(err)
}
v, err := box.Encrypt("sk-abc-xyz")
if err != nil {
t.Fatal(err)
}
if strings.HasPrefix(v, "sk-") || strings.Contains(v, "abc-xyz") {
t.Fatalf("ciphertext leaked plaintext: %q", v)
}
out, err := box.Decrypt(v)
if err != nil || out != "sk-abc-xyz" {
t.Fatalf("roundtrip: %q %v", out, err)
}
if plain, err := box.Decrypt("sk-plain"); err != nil || plain != "sk-plain" {
t.Fatalf("plain passthrough: %q %v", plain, err)
}
// wrong key must error
bad, _ := NewSecretBox(filepath.Join(t.TempDir(), "runtime.json"))
if _, err := bad.Decrypt(v); err == nil {
t.Fatal("expected decrypt failure with wrong key")
}
}
func TestRemoveSourceFromYAML(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "cfg.yaml")
content := `
listen: 127.0.0.1:9999
gateway_keys: [sk-1]
default_model: AUTO
adapter_dir: adapters
runtime_file: runtime.json
sources:
- name: deepseek
base_url: https://api.deepseek.com
api_key: sk-d
adapter: deepseek
models:
- id: deepseek-v4-flash
- name: ollama
base_url: http://127.0.0.1:11434
adapter: ollama
models:
- id: llama3
`
if err := os.WriteFile(path, []byte(content), 0644); err != nil {
t.Fatal(err)
}
if err := RemoveSourceFromYAML(path, "deepseek"); err != nil {
t.Fatalf("remove: %v", err)
}
reloaded, err := Load(path)
if err != nil {
t.Fatalf("reload: %v", err)
}
if len(reloaded.Sources) != 1 {
t.Fatalf("sources = %d, want 1", len(reloaded.Sources))
}
if reloaded.Sources[0].Name != "ollama" {
t.Fatalf("remaining = %q, want ollama", reloaded.Sources[0].Name)
}
if reloaded.Sources[0].Models[0].ID != "llama3" {
t.Fatalf("remaining models broken: %+v", reloaded.Sources[0].Models)
}
// removing a non-existent name is a no-op that keeps the file valid
if err := RemoveSourceFromYAML(path, "nope"); err != nil {
t.Fatalf("remove missing: %v", err)
}
if cfg, err := Load(path); err != nil || len(cfg.Sources) != 1 {
t.Fatalf("after no-op: %v %v", len(cfg.Sources), err)
}
}