package service import ( "context" "net" "testing" toki "git.toki-labs.com/toki/proto-socket/go" "google.golang.org/protobuf/proto" edgeevents "iop/apps/edge/internal/events" edgenode "iop/apps/edge/internal/node" "iop/packages/go/config" iop "iop/proto/gen/iop" ) func newUsageAttributionService(t *testing.T, providerType string) *Service { t.Helper() edgeConn, nodeConn := net.Pipe() t.Cleanup(func() { _ = edgeConn.Close() _ = nodeConn.Close() }) parserMap := toki.ParserMap{ toki.TypeNameOf(&iop.RunRequest{}): func(data []byte) (proto.Message, error) { message := &iop.RunRequest{} return message, proto.Unmarshal(data, message) }, toki.TypeNameOf(&iop.ProviderTunnelRequest{}): func(data []byte) (proto.Message, error) { message := &iop.ProviderTunnelRequest{} return message, proto.Unmarshal(data, message) }, } edgeClient := toki.NewTcpClient(edgeConn, 0, 0, parserMap) _ = toki.NewTcpClient(nodeConn, 0, 0, parserMap) registry := edgenode.NewRegistry() registry.Register(&edgenode.NodeEntry{ NodeID: "node-a", LifecycleState: edgenode.LifecycleConnected, Client: edgeClient, }) service := New(registry, edgeevents.NewBus()) store := edgenode.NewNodeStore() store.Add(&edgenode.NodeRecord{ ID: "node-a", Providers: []config.NodeProviderConf{{ ID: "provider-a", Type: providerType, Models: []string{"served-a"}, Health: "available", Capacity: 1, }}, }) service.SetNodeStore(store) service.SetModelCatalog([]config.ModelCatalogEntry{{ ID: "logical-model", UsageAttribution: config.UsageAttributionModelGroup, Providers: map[string]string{"provider-a": "served-a"}, }}) return service } func assertUsageAttributionDispatch( t *testing.T, dispatch RunDispatch, providerID, target, nodeID, attribution string, ) { t.Helper() if dispatch.ProviderID != providerID { t.Errorf("provider_id = %q, want %q", dispatch.ProviderID, providerID) } if dispatch.Target != target { t.Errorf("target = %q, want %q", dispatch.Target, target) } if dispatch.NodeID != nodeID { t.Errorf("node_id = %q, want %q", dispatch.NodeID, nodeID) } if dispatch.UsageAttribution != attribution { t.Errorf("usage_attribution = %q, want %q", dispatch.UsageAttribution, attribution) } } func TestRunDispatchActualProviderBinding(t *testing.T) { t.Run("direct normalized", func(t *testing.T) { service := newUsageAttributionService(t, "ollama") result, err := service.SubmitRun(context.Background(), SubmitRunRequest{ NodeRef: "node-a", ProviderID: "provider-direct", UsageAttribution: config.UsageAttributionProvider, Adapter: "cli", Target: "direct-target", Background: true, }) if err != nil { t.Fatalf("SubmitRun: %v", err) } defer result.Close() assertUsageAttributionDispatch( t, result.Dispatch(), "provider-direct", "direct-target", "node-a", config.UsageAttributionProvider, ) }) t.Run("provider pool normalized", func(t *testing.T) { service := newUsageAttributionService(t, "ollama") result, err := service.SubmitProviderPool(context.Background(), ProviderPoolDispatchRequest{ Run: SubmitRunRequest{ ModelGroupKey: "logical-model", UsageAttribution: config.UsageAttributionModelGroup, ProviderPool: true, Background: true, }, }) if err != nil { t.Fatalf("SubmitProviderPool: %v", err) } if result.Path != ProviderPoolPathNormalized || result.Run == nil { t.Fatalf("path = %q run=%v, want normalized run", result.Path, result.Run != nil) } defer result.Run.Close() defer service.queue.releaseRun(result.DispatchInfo.RunID, "test-cleanup") assertUsageAttributionDispatch( t, result.DispatchInfo, "provider-a", "served-a", "node-a", config.UsageAttributionModelGroup, ) assertUsageAttributionDispatch( t, result.Run.Dispatch(), "provider-a", "served-a", "node-a", config.UsageAttributionModelGroup, ) }) } func TestTunnelDispatchActualProviderBinding(t *testing.T) { t.Run("direct tunnel", func(t *testing.T) { service := newUsageAttributionService(t, "vllm") result, err := service.SubmitProviderTunnel(context.Background(), SubmitProviderTunnelRequest{ NodeRef: "node-a", ProviderID: "provider-direct", UsageAttribution: config.UsageAttributionProvider, Adapter: "openai_compat", Target: "direct-target", Method: "POST", Path: "/v1/chat/completions", }) if err != nil { t.Fatalf("SubmitProviderTunnel: %v", err) } defer result.Close() assertUsageAttributionDispatch( t, result.Dispatch(), "provider-direct", "direct-target", "node-a", config.UsageAttributionProvider, ) }) t.Run("provider pool tunnel", func(t *testing.T) { service := newUsageAttributionService(t, "vllm") result, err := service.SubmitProviderPool(context.Background(), ProviderPoolDispatchRequest{ Run: SubmitRunRequest{ ModelGroupKey: "logical-model", UsageAttribution: config.UsageAttributionModelGroup, ProviderPool: true, }, Tunnel: SubmitProviderTunnelRequest{ Method: "POST", Path: "/v1/chat/completions", }, }) if err != nil { t.Fatalf("SubmitProviderPool: %v", err) } if result.Path != ProviderPoolPathTunnel || result.Tunnel == nil { t.Fatalf("path = %q tunnel=%v, want provider tunnel", result.Path, result.Tunnel != nil) } defer result.Tunnel.Close() assertUsageAttributionDispatch( t, result.DispatchInfo, "provider-a", "served-a", "node-a", config.UsageAttributionModelGroup, ) assertUsageAttributionDispatch( t, result.Tunnel.Dispatch(), "provider-a", "served-a", "node-a", config.UsageAttributionModelGroup, ) }) } func TestManagedTunnelDispatchFreezesCredentialSlotAttribution(t *testing.T) { env := newProviderTunnelTestEnv(t) provider := &fakeCredentialLeaseProvider{valid: true} env.svc.SetCredentialLeaseProvider(provider) env.svc.SetCredentialLeaseLimit(1) dispatch := func(slot string, revision uint64) RunDispatch { t.Helper() result, err := env.svc.SubmitProviderPool(context.Background(), ProviderPoolDispatchRequest{ Run: SubmitRunRequest{ ModelGroupKey: "qwen3.6:35b", UsageAttribution: config.UsageAttributionProvider, ProviderPool: true, }, Tunnel: SubmitProviderTunnelRequest{ CredentialBinding: &CredentialBinding{ PrincipalRef: "principal-one", CredentialSlotRef: slot, RouteID: "route-one", ProfileID: "openai", CredentialRevision: revision, RouteRevision: 3, ProjectionGeneration: 5, }, Method: "POST", Path: "/v1/chat/completions", Operation: "chat_completions", BuildBody: func(string) ([]byte, error) { return []byte(`{"model":"served-qwen"}`), nil }, }, }) if err != nil { t.Fatalf("SubmitProviderPool slot %q: %v", slot, err) } got := result.DispatchInfo if got.CredentialSlotRef != slot || got.CredentialRevision != revision { t.Fatalf("dispatch slot attribution = %q/%d, want %q/%d", got.CredentialSlotRef, got.CredentialRevision, slot, revision) } if tunnelGot := result.Tunnel.Dispatch(); tunnelGot.CredentialSlotRef != slot || tunnelGot.CredentialRevision != revision { t.Fatalf("tunnel slot attribution = %q/%d, want %q/%d", tunnelGot.CredentialSlotRef, tunnelGot.CredentialRevision, slot, revision) } result.Tunnel.Close() waitForCondition(t, func() bool { return inflightRunCount(env.svc.queue) == 0 }, "closing managed tunnel did not release admission") return got } first := dispatch("slot-a", 11) retry := dispatch("slot-a", 11) second := dispatch("slot-b", 19) for name, got := range map[string]RunDispatch{"first": first, "retry": retry, "second": second} { if got.ProviderID != "prov-vllm-01" || got.Target != "served-qwen" || got.ModelGroupKey != "qwen3.6:35b" { t.Fatalf("%s dispatch changed provider/model binding: %+v", name, got) } } if retry.CredentialSlotRef != first.CredentialSlotRef || retry.CredentialRevision != first.CredentialRevision { t.Fatalf("retry changed frozen credential slot: first=%+v retry=%+v", first, retry) } if second.CredentialSlotRef == first.CredentialSlotRef { t.Fatalf("distinct slot dispatches collapsed: first=%+v second=%+v", first, second) } }