diff --git a/internal/modelref/modelref_test.go b/internal/modelref/modelref_test.go new file mode 100644 index 00000000..5a627cb7 --- /dev/null +++ b/internal/modelref/modelref_test.go @@ -0,0 +1,119 @@ +package modelref_test + +import ( + "testing" + + "moonbridge/internal/modelref" +) + +func TestParse(t *testing.T) { + tests := []struct { + name string + ref string + wantProvider string + wantModel string + }{ + { + name: "provider slash model", + ref: "openai/gpt-4o", + wantProvider: "openai", + wantModel: "gpt-4o", + }, + { + name: "model paren provider", + ref: "claude-opus-4-6(kiro)", + wantProvider: "kiro", + wantModel: "claude-opus-4-6", + }, + { + name: "no separator returns empty provider and original ref", + ref: "gpt-4o", + wantProvider: "", + wantModel: "gpt-4o", + }, + { + name: "leading and trailing whitespace is trimmed", + ref: " openai / gpt-4o ", + wantProvider: "openai", + wantModel: "gpt-4o", + }, + { + name: "whitespace inside paren form is trimmed", + ref: " claude ( kiro ) ", + wantProvider: "kiro", + wantModel: "claude", + }, + { + name: "paren form preferred over slash when both present", + ref: "anthropic/claude(kiro)", + wantProvider: "kiro", + wantModel: "anthropic/claude", + }, + { + name: "empty provider inside parens falls back to slash form", + ref: "provider/model()", + wantProvider: "provider", + wantModel: "model()", + }, + { + name: "paren form matches even when model contains a slash", + ref: "a/b(provider)", + wantProvider: "provider", + wantModel: "a/b", + }, + { + name: "open paren at index zero is not treated as paren form", + ref: "(provider)", + wantProvider: "", + wantModel: "(provider)", + }, + { + name: "open paren without closing suffix uses slash form", + ref: "model(provider", + wantProvider: "", + wantModel: "model(provider", + }, + { + name: "empty string", + ref: "", + wantProvider: "", + wantModel: "", + }, + { + name: "only slash yields empty provider and model", + ref: "/", + wantProvider: "", + wantModel: "", + }, + { + name: "slash with empty provider", + ref: "/model", + wantProvider: "", + wantModel: "model", + }, + { + name: "slash with empty model", + ref: "provider/", + wantProvider: "provider", + wantModel: "", + }, + { + name: "first slash is used to split", + ref: "a/b/c", + wantProvider: "a", + wantModel: "b/c", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gotProvider, gotModel := modelref.Parse(tt.ref) + if gotProvider != tt.wantProvider { + t.Errorf("Parse(%q) provider = %q, want %q", tt.ref, gotProvider, tt.wantProvider) + } + if gotModel != tt.wantModel { + t.Errorf("Parse(%q) model = %q, want %q", tt.ref, gotModel, tt.wantModel) + } + }) + } +} diff --git a/internal/service/server/session/manager_test.go b/internal/service/server/session/manager_test.go new file mode 100644 index 00000000..5a2d2b3c --- /dev/null +++ b/internal/service/server/session/manager_test.go @@ -0,0 +1,220 @@ +package session_test + +import ( + "testing" + "time" + + "moonbridge/internal/extension/plugin" + sessionmgr "moonbridge/internal/service/server/session" +) + +// fakeConfig is a test ConfigAccessor with fixed TTL and max sessions. +type fakeConfig struct { + ttl time.Duration + maxSessions int +} + +func (c fakeConfig) SessionTTL() time.Duration { return c.ttl } +func (c fakeConfig) MaxSessions() int { return c.maxSessions } + +func newManager(ttl time.Duration, maxSessions int) *sessionmgr.InMemoryManager { + return sessionmgr.NewInMemoryManager(fakeConfig{ttl: ttl, maxSessions: maxSessions}, nil) +} + +func TestGetOrCreateCreatesAndReuses(t *testing.T) { + m := newManager(time.Hour, 0) + defer m.Stop() + + now := time.Now() + s1 := m.GetOrCreate("a", now) + if s1 == nil { + t.Fatal("GetOrCreate returned nil") + } + if s1.ID != "a" { + t.Errorf("session ID = %q, want a", s1.ID) + } + + s2 := m.GetOrCreate("a", now.Add(time.Minute)) + if s1 != s2 { + t.Error("GetOrCreate should return the same session instance for the same key") + } + + s3 := m.GetOrCreate("b", now) + if s3 == s1 { + t.Error("GetOrCreate should return a distinct session for a different key") + } +} + +func TestGetOrCreateInitializesExtensions(t *testing.T) { + m := newManager(time.Hour, 0) + defer m.Stop() + + // With a nil plugin registry, ExtensionData is initialized to nil (not left unset). + s := m.GetOrCreate("a", time.Now()) + if s.ExtensionData != nil { + t.Errorf("ExtensionData = %v, want nil with nil registry", s.ExtensionData) + } +} + +func TestListReturnsSnapshot(t *testing.T) { + m := newManager(time.Hour, 0) + defer m.Stop() + + now := time.Now() + m.GetOrCreate("a", now) + m.GetOrCreate("b", now) + + infos := m.List() + if len(infos) != 2 { + t.Fatalf("List returned %d sessions, want 2", len(infos)) + } + keys := map[string]bool{} + for _, info := range infos { + keys[info.Key] = true + if info.CreatedAt == "" { + t.Errorf("session %q has empty CreatedAt", info.Key) + } + if info.LastUsed == "" { + t.Errorf("session %q has empty LastUsed", info.Key) + } + } + if !keys["a"] || !keys["b"] { + t.Errorf("List keys = %v, want a and b", keys) + } +} + +func TestPruneRemovesExpiredSessions(t *testing.T) { + m := newManager(90*time.Minute, 0) + defer m.Stop() + + base := time.Now() + m.GetOrCreate("old", base) + m.GetOrCreate("fresh", base.Add(time.Hour)) + + // Prune at base+2h: "old" is stale (used at base), "fresh" used at base+1h. + m.Prune(base.Add(2 * time.Hour)) + + infos := m.List() + if len(infos) != 1 { + t.Fatalf("after prune got %d sessions, want 1", len(infos)) + } + if infos[0].Key != "fresh" { + t.Errorf("remaining session = %q, want fresh", infos[0].Key) + } +} + +func TestGetOrCreatePrunesBeforeLookup(t *testing.T) { + m := newManager(30*time.Minute, 0) + defer m.Stop() + + base := time.Now() + first := m.GetOrCreate("a", base) + + // Access with a much-later time: the stale entry is pruned and recreated. + second := m.GetOrCreate("a", base.Add(2*time.Hour)) + if first == second { + t.Error("expected a fresh session after the previous one expired") + } +} + +func TestMaxSessionsEvictsLRU(t *testing.T) { + m := newManager(time.Hour, 2) + defer m.Stop() + + base := time.Now() + m.GetOrCreate("a", base) + m.GetOrCreate("b", base.Add(time.Minute)) + // Adding a third session should evict the least-recently-used ("a"). + m.GetOrCreate("c", base.Add(2*time.Minute)) + + infos := m.List() + if len(infos) != 2 { + t.Fatalf("got %d sessions, want 2", len(infos)) + } + keys := map[string]bool{} + for _, info := range infos { + keys[info.Key] = true + } + if keys["a"] { + t.Error("expected LRU session 'a' to have been evicted") + } + if !keys["b"] || !keys["c"] { + t.Errorf("expected b and c to remain, got %v", keys) + } +} + +func TestMaxSessionsReuseUpdatesRecency(t *testing.T) { + m := newManager(time.Hour, 2) + defer m.Stop() + + base := time.Now() + m.GetOrCreate("a", base) + m.GetOrCreate("b", base.Add(time.Minute)) + // Touch "a" so it becomes most-recently-used. + m.GetOrCreate("a", base.Add(2*time.Minute)) + // Adding "c" should now evict "b" instead of "a". + m.GetOrCreate("c", base.Add(3*time.Minute)) + + keys := map[string]bool{} + for _, info := range m.List() { + keys[info.Key] = true + } + if keys["b"] { + t.Error("expected 'b' to be evicted as LRU") + } + if !keys["a"] || !keys["c"] { + t.Errorf("expected a and c to remain, got %v", keys) + } +} + +func TestGetOrCreateWithPluginRegistry(t *testing.T) { + reg := plugin.NewRegistry(nil) + m := sessionmgr.NewInMemoryManager(fakeConfig{ttl: time.Hour}, reg) + defer m.Stop() + + now := time.Now() + s1 := m.GetOrCreate("a", now) + if s1 == nil { + t.Fatal("GetOrCreate returned nil") + } + // Reusing the key exercises the ExtensionData backfill branch. + if s2 := m.GetOrCreate("a", now.Add(time.Minute)); s2 != s1 { + t.Error("expected the same session on reuse") + } +} + +func TestNewEphemeralWithPluginRegistry(t *testing.T) { + reg := plugin.NewRegistry(nil) + m := sessionmgr.NewInMemoryManager(fakeConfig{ttl: time.Hour}, reg) + defer m.Stop() + + if s := m.NewEphemeral(); s == nil { + t.Fatal("NewEphemeral returned nil") + } +} + +func TestNewEphemeralIsNotTracked(t *testing.T) { + m := newManager(time.Hour, 0) + defer m.Stop() + + s := m.NewEphemeral() + if s == nil { + t.Fatal("NewEphemeral returned nil") + } + if s.ID == "" { + t.Error("ephemeral session should have a generated ID") + } + if len(m.List()) != 0 { + t.Error("ephemeral session must not be tracked by the manager") + } +} + +func TestStopIsIdempotentlyCloseable(t *testing.T) { + m := newManager(time.Hour, 0) + m.Stop() + // A second Stop would panic on a double-close; ensure we only call it once + // but that the manager remains usable for reads after stopping. + if got := m.List(); len(got) != 0 { + t.Errorf("List after Stop = %v, want empty", got) + } +} diff --git a/internal/service/server/trace/writer_test.go b/internal/service/server/trace/writer_test.go new file mode 100644 index 00000000..dbbf8853 --- /dev/null +++ b/internal/service/server/trace/writer_test.go @@ -0,0 +1,172 @@ +package trace_test + +import ( + "bytes" + "os" + "path/filepath" + "testing" + + srvtrace "moonbridge/internal/service/server/trace" + mbtrace "moonbridge/internal/service/trace" +) + +func newTracer(t *testing.T, enabled bool) (*mbtrace.Tracer, string) { + t.Helper() + root := t.TempDir() + return mbtrace.New(mbtrace.Config{Enabled: enabled, Root: root, SessionID: "sess", Flat: true}), root +} + +// categoryFiles returns the JSON files written under root/. +func categoryFiles(t *testing.T, root, category string) []string { + t.Helper() + entries, err := os.ReadDir(filepath.Join(root, category)) + if os.IsNotExist(err) { + return nil + } + if err != nil { + t.Fatalf("ReadDir(%s): %v", category, err) + } + var names []string + for _, e := range entries { + names = append(names, e.Name()) + } + return names +} + +func TestWriteTraceDispatchesByCategory(t *testing.T) { + tests := []struct { + name string + record mbtrace.Record + wantDirs []string + }{ + { + name: "chat only", + record: mbtrace.Record{ChatRequest: map[string]any{"k": "v"}}, + wantDirs: []string{"Chat"}, + }, + { + name: "response via openai request", + record: mbtrace.Record{OpenAIRequest: map[string]any{"k": "v"}}, + wantDirs: []string{"Response"}, + }, + { + name: "response via upstream request", + record: mbtrace.Record{UpstreamRequest: map[string]any{"k": "v"}}, + wantDirs: []string{"Response"}, + }, + { + name: "anthropic only", + record: mbtrace.Record{AnthropicResponse: map[string]any{"k": "v"}}, + wantDirs: []string{"Anthropic"}, + }, + { + name: "all three categories", + record: mbtrace.Record{ + ChatRequest: map[string]any{"k": "v"}, + OpenAIRequest: map[string]any{"k": "v"}, + AnthropicRequest: map[string]any{"k": "v"}, + }, + wantDirs: []string{"Chat", "Response", "Anthropic"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tracer, root := newTracer(t, true) + var errBuf bytes.Buffer + w := srvtrace.NewFileWriter(tracer, &errBuf) + + w.WriteTrace(tt.record) + + for _, dir := range tt.wantDirs { + if files := categoryFiles(t, root, dir); len(files) == 0 { + t.Errorf("expected a trace file under %q, found none", dir) + } + } + for _, dir := range []string{"Chat", "Response", "Anthropic"} { + if !contains(tt.wantDirs, dir) { + if files := categoryFiles(t, root, dir); len(files) != 0 { + t.Errorf("did not expect trace files under %q, got %v", dir, files) + } + } + } + if errBuf.Len() != 0 { + t.Errorf("unexpected error output: %q", errBuf.String()) + } + }) + } +} + +func TestWriteTraceNoCategoryWritesNothing(t *testing.T) { + tracer, root := newTracer(t, true) + w := srvtrace.NewFileWriter(tracer, nil) + + // A record with no request/response payloads matches no category. + w.WriteTrace(mbtrace.Record{Model: "m"}) + + for _, dir := range []string{"Chat", "Response", "Anthropic"} { + if files := categoryFiles(t, root, dir); len(files) != 0 { + t.Errorf("expected no files under %q, got %v", dir, files) + } + } +} + +func TestWriteTraceDisabledTracer(t *testing.T) { + tracer, root := newTracer(t, false) + w := srvtrace.NewFileWriter(tracer, nil) + + w.WriteTrace(mbtrace.Record{ChatRequest: map[string]any{"k": "v"}}) + + if entries, err := os.ReadDir(root); err == nil && len(entries) != 0 { + t.Errorf("disabled tracer should not write, found %d entries", len(entries)) + } +} + +func TestWriteTraceNilTracer(t *testing.T) { + w := srvtrace.NewFileWriter(nil, nil) + // Must be a no-op and must not panic. + w.WriteTrace(mbtrace.Record{ChatRequest: map[string]any{"k": "v"}}) +} + +func TestWriteCategoryWritesFile(t *testing.T) { + tracer, root := newTracer(t, true) + w := srvtrace.NewFileWriter(tracer, nil) + + w.WriteCategory("Response", 7, mbtrace.Record{OpenAIRequest: map[string]any{"k": "v"}}) + + path := filepath.Join(root, "Response", "7.json") + if _, err := os.Stat(path); err != nil { + t.Fatalf("expected trace file at %s: %v", path, err) + } +} + +func TestWriteCategoryReportsError(t *testing.T) { + // Point the tracer root at a path whose parent is a regular file so that + // MkdirAll fails, exercising the error-reporting branch. + root := t.TempDir() + filePath := filepath.Join(root, "notadir") + if err := os.WriteFile(filePath, []byte("x"), 0o600); err != nil { + t.Fatal(err) + } + tracer := mbtrace.New(mbtrace.Config{Enabled: true, Root: filePath, SessionID: "sess", Flat: true}) + var errBuf bytes.Buffer + w := srvtrace.NewFileWriter(tracer, &errBuf) + + w.WriteCategory("Response", 1, mbtrace.Record{OpenAIRequest: map[string]any{"k": "v"}}) + + if errBuf.Len() == 0 { + t.Error("expected an error to be written to the error writer") + } +} + +// Ensure FileWriter satisfies the Writer interface. +var _ srvtrace.Writer = (*srvtrace.FileWriter)(nil) + +func contains(s []string, v string) bool { + for _, x := range s { + if x == v { + return true + } + } + return false +} diff --git a/internal/service/server/usage/tracker_test.go b/internal/service/server/usage/tracker_test.go new file mode 100644 index 00000000..e2456a7d --- /dev/null +++ b/internal/service/server/usage/tracker_test.go @@ -0,0 +1,94 @@ +package usage_test + +import ( + "testing" + + "moonbridge/internal/service/server/usage" + "moonbridge/internal/service/stats" +) + +func newStatsWithPricing(model string, p stats.ModelPricing) *stats.SessionStats { + s := stats.NewSessionStats() + s.SetPricing(map[string]stats.ModelPricing{model: p}) + return s +} + +func TestStatsTrackerRecordBilling(t *testing.T) { + s := newStatsWithPricing("alias", stats.ModelPricing{InputPrice: 1, OutputPrice: 2}) + tr := usage.NewStatsTracker(s) + + tr.RecordBilling("alias", "upstream-model", stats.BillingUsage{ + FreshInputTokens: 1_000_000, + OutputTokens: 500_000, + }) + + summary := s.Summary() + if summary.Requests != 1 { + t.Fatalf("Requests = %d, want 1", summary.Requests) + } + if summary.InputTokens != 1_000_000 { + t.Errorf("InputTokens = %d, want 1000000", summary.InputTokens) + } + if summary.OutputTokens != 500_000 { + t.Errorf("OutputTokens = %d, want 500000", summary.OutputTokens) + } + // 1M fresh input * 1/1M + 500k output * 2/1M = 1 + 1 = 2 + if got := summary.TotalCost; got != 2 { + t.Errorf("TotalCost = %v, want 2", got) + } + if summary.ActualModelNames["alias"] != "upstream-model" { + t.Errorf("ActualModelNames[alias] = %q, want upstream-model", summary.ActualModelNames["alias"]) + } +} + +func TestStatsTrackerRecordBillingNilStats(t *testing.T) { + tr := usage.NewStatsTracker(nil) + // Should be a no-op and must not panic. + tr.RecordBilling("alias", "upstream", stats.BillingUsage{FreshInputTokens: 10}) +} + +func TestStatsTrackerCostForRequest(t *testing.T) { + s := newStatsWithPricing("alias", stats.ModelPricing{ + InputPrice: 3, + OutputPrice: 6, + CacheWritePrice: 4, + CacheReadPrice: 1, + }) + tr := usage.NewStatsTracker(s) + + cost := tr.CostForRequest("alias", "upstream", "provider", stats.BillingUsage{ + FreshInputTokens: 1_000_000, + OutputTokens: 1_000_000, + CacheCreationInputTokens: 1_000_000, + CacheReadInputTokens: 1_000_000, + }) + + // 3 + 6 + 4 + 1 = 14 + if cost != 14 { + t.Errorf("CostForRequest = %v, want 14", cost) + } + + // CostForRequest must not record anything into the stats. + if reqs := s.Summary().Requests; reqs != 0 { + t.Errorf("CostForRequest should not record usage, Requests = %d", reqs) + } +} + +func TestStatsTrackerCostForRequestUnknownModel(t *testing.T) { + s := newStatsWithPricing("alias", stats.ModelPricing{InputPrice: 3}) + tr := usage.NewStatsTracker(s) + + if cost := tr.CostForRequest("unknown", "upstream", "provider", stats.BillingUsage{FreshInputTokens: 1_000_000}); cost != 0 { + t.Errorf("CostForRequest for unknown model = %v, want 0", cost) + } +} + +func TestStatsTrackerCostForRequestNilStats(t *testing.T) { + tr := usage.NewStatsTracker(nil) + if cost := tr.CostForRequest("alias", "upstream", "provider", stats.BillingUsage{FreshInputTokens: 1_000_000}); cost != 0 { + t.Errorf("CostForRequest with nil stats = %v, want 0", cost) + } +} + +// Ensure StatsTracker satisfies the Tracker interface. +var _ usage.Tracker = (*usage.StatsTracker)(nil)