package openai import ( "bytes" "context" "io" "sync" "testing" "time" edgeservice "iop/apps/edge/internal/service" "iop/packages/go/streamgate" ) type dispatcherEventSource struct{} func (dispatcherEventSource) NextEvent(context.Context) (streamgate.NormalizedEvent, error) { return streamgate.NormalizedEvent{}, io.EOF } type dispatcherRunHandle struct { dispatch edgeservice.RunDispatch close func() once sync.Once } func (h *dispatcherRunHandle) Dispatch() edgeservice.RunDispatch { return h.dispatch } func (h *dispatcherRunHandle) Stream() edgeservice.RunStream { return edgeservice.RunStream{} } func (h *dispatcherRunHandle) WaitTimeout() time.Duration { return time.Second } func (h *dispatcherRunHandle) Close() { if h != nil && h.close != nil { h.once.Do(h.close) } } type dispatcherTunnelHandle struct { dispatch edgeservice.RunDispatch close func() once sync.Once } func (h *dispatcherTunnelHandle) Dispatch() edgeservice.RunDispatch { return h.dispatch } func (h *dispatcherTunnelHandle) Stream() edgeservice.ProviderTunnelStream { return edgeservice.ProviderTunnelStream{} } func (h *dispatcherTunnelHandle) WaitTimeout() time.Duration { return time.Second } func (h *dispatcherTunnelHandle) SetHeaders(map[string]string) {} func (h *dispatcherTunnelHandle) Close() { if h != nil && h.close != nil { h.once.Do(h.close) } } type dispatcherServiceSpy struct { poolPath string runCalls int tunnelCalls int poolCalls int cancelCalls int closeCalls int lastHeaders map[string]string } func (s *dispatcherServiceSpy) dispatch(path string) edgeservice.RunDispatch { return edgeservice.RunDispatch{ RunID: "attempt-" + path, NodeID: "node.actual", ModelGroupKey: "alias", Adapter: "adapter.actual", Target: "model.actual", SessionID: "session.actual", ProviderID: "provider.actual", ProviderType: "openai", ExecutionPath: path, } } func (s *dispatcherServiceSpy) SubmitRun(context.Context, edgeservice.SubmitRunRequest) (edgeservice.RunResult, error) { s.runCalls++ return &dispatcherRunHandle{dispatch: s.dispatch("normalized"), close: func() { s.closeCalls++ }}, nil } func (s *dispatcherServiceSpy) SubmitProviderTunnel(_ context.Context, request edgeservice.SubmitProviderTunnelRequest) (edgeservice.ProviderTunnelResult, error) { s.tunnelCalls++ s.lastHeaders = request.Headers return &dispatcherTunnelHandle{dispatch: s.dispatch("provider_tunnel"), close: func() { s.closeCalls++ }}, nil } func (s *dispatcherServiceSpy) SubmitProviderPool(_ context.Context, request edgeservice.ProviderPoolDispatchRequest) (*edgeservice.ProviderPoolDispatchResult, error) { s.poolCalls++ if s.poolPath == "provider_tunnel" { tunnel := request.Tunnel var err error if request.PrepareTunnel != nil { tunnel, err = request.PrepareTunnel(tunnel) if err != nil { return nil, err } } s.lastHeaders = tunnel.Headers handle := &dispatcherTunnelHandle{dispatch: s.dispatch("provider_tunnel"), close: func() { s.closeCalls++ }} return &edgeservice.ProviderPoolDispatchResult{ Path: edgeservice.ProviderPoolPathTunnel, Tunnel: handle, DispatchInfo: handle.Dispatch(), }, nil } handle := &dispatcherRunHandle{dispatch: s.dispatch("normalized"), close: func() { s.closeCalls++ }} return &edgeservice.ProviderPoolDispatchResult{ Path: edgeservice.ProviderPoolPathNormalized, Run: handle, DispatchInfo: handle.Dispatch(), }, nil } func (s *dispatcherServiceSpy) CancelRun(context.Context, edgeservice.CancelRunRequest) (edgeservice.CommandResult, error) { s.cancelCalls++ return edgeservice.CommandResult{}, nil } func (s *dispatcherServiceSpy) OllamaAPI(context.Context, edgeservice.OllamaAPIRequest) (edgeservice.OllamaAPIView, error) { return edgeservice.OllamaAPIView{}, nil } func rebuiltRequestForDispatcher(t *testing.T, rebuilder *openAIRequestRebuilder, ref streamgate.RecoveryRequestSnapshotRef, id string) streamgate.RebuiltRequest { t.Helper() directive, err := streamgate.NewRecoveryDirectiveExact(ref.SnapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveExact: %v", err) } plan := mustOpenAIRecoveryPlan(t, id, streamgate.RecoveryStrategyExactReplay, directive) draft, err := rebuilder.RebuildRequest(context.Background(), ref, plan) if err != nil { t.Fatalf("RebuildRequest: %v", err) } _, request, err := plan.FinalizeRebuiltRequest(draft) if err != nil { t.Fatalf("FinalizeRebuiltRequest: %v", err) } return request } func newDispatcherFixture(t *testing.T, service *dispatcherServiceSpy, builder openAIAttemptAdmissionBuilder) (*openAIRequestRebuilder, streamgate.RecoveryRequestSnapshotRef, *openAIAttemptDispatcher) { t.Helper() body := []byte(`{"model":"alias","messages":[]}`) _, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 4096) dispatcher, err := newOpenAIAttemptDispatcher(service, rebuilder.RebuiltStore(), builder, func(openAIAttemptTransport) (streamgate.NormalizedEventSource, error) { return dispatcherEventSource{}, nil }) if err != nil { t.Fatalf("newOpenAIAttemptDispatcher: %v", err) } return rebuilder, ref, dispatcher } func TestOpenAIAttemptDispatcherExistingAdmissionSurfaces(t *testing.T) { tests := []struct { name string kind openAIAdmissionKind }{ {name: "run", kind: openAIAdmissionRun}, {name: "tunnel", kind: openAIAdmissionTunnel}, {name: "pool", kind: openAIAdmissionPool}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { service := &dispatcherServiceSpy{poolPath: "normalized"} rebuilder, ref, dispatcher := newDispatcherFixture(t, service, func(_ context.Context, _ streamgate.RebuiltRequest, body []byte) (openAIAttemptAdmission, error) { admission := openAIAttemptAdmission{kind: tc.kind} admission.run = edgeservice.SubmitRunRequest{ModelGroupKey: "alias", Target: "model"} admission.tunnel = edgeservice.SubmitProviderTunnelRequest{Path: openAIRebuildEndpointChat, Body: body} admission.pool = edgeservice.ProviderPoolDispatchRequest{ Run: edgeservice.SubmitRunRequest{ModelGroupKey: "alias", ProviderPool: true}, Tunnel: edgeservice.SubmitProviderTunnelRequest{Path: openAIRebuildEndpointChat, Body: body}, } return admission, nil }) request := rebuiltRequestForDispatcher(t, rebuilder, ref, "plan.surface."+tc.name) binding, err := dispatcher.DispatchAttempt(context.Background(), request) if err != nil { t.Fatalf("DispatchAttempt: %v", err) } if binding.Model() != "model.actual" || binding.Provider() != "provider.actual" { t.Fatalf("binding did not use actual admission: model=%s provider=%s", binding.Model(), binding.Provider()) } if err := binding.Controller().AbortAttempt(context.Background()); err != nil { t.Fatalf("AbortAttempt: %v", err) } if service.runCalls+service.tunnelCalls+service.poolCalls != 1 { t.Fatalf("admission calls = run:%d tunnel:%d pool:%d", service.runCalls, service.tunnelCalls, service.poolCalls) } }) } } func TestOpenAIAttemptDispatcherPoolPathSwitchAndFreshAuth(t *testing.T) { service := &dispatcherServiceSpy{poolPath: "normalized"} token := "token-one" authCalls := 0 rebuilder, ref, dispatcher := newDispatcherFixture(t, service, func(_ context.Context, _ streamgate.RebuiltRequest, body []byte) (openAIAttemptAdmission, error) { return openAIAttemptAdmission{ kind: openAIAdmissionPool, pool: edgeservice.ProviderPoolDispatchRequest{ Run: edgeservice.SubmitRunRequest{ModelGroupKey: "alias", ProviderPool: true}, Tunnel: edgeservice.SubmitProviderTunnelRequest{Path: openAIRebuildEndpointChat, Body: body}, }, authorize: func(context.Context) (map[string]string, error) { authCalls++ return map[string]string{"Authorization": token}, nil }, }, nil }) first := rebuiltRequestForDispatcher(t, rebuilder, ref, "plan.switch.one") firstBinding, err := dispatcher.DispatchAttempt(context.Background(), first) if err != nil { t.Fatalf("first DispatchAttempt: %v", err) } if firstBinding.ExecutionPath() != "normalized" || authCalls != 0 { t.Fatalf("normalized binding/auth = %s/%d", firstBinding.ExecutionPath(), authCalls) } _ = firstBinding.Controller().AbortAttempt(context.Background()) service.poolPath = "provider_tunnel" token = "token-two" second := rebuiltRequestForDispatcher(t, rebuilder, ref, "plan.switch.two") secondBinding, err := dispatcher.DispatchAttempt(context.Background(), second) if err != nil { t.Fatalf("second DispatchAttempt: %v", err) } if secondBinding.ExecutionPath() != "provider_tunnel" || authCalls != 1 { t.Fatalf("tunnel binding/auth = %s/%d", secondBinding.ExecutionPath(), authCalls) } if service.lastHeaders["Authorization"] != "token-two" { t.Fatalf("stale auth header = %#v", service.lastHeaders) } if err := secondBinding.Controller().AbortAttempt(context.Background()); err != nil { t.Fatalf("AbortAttempt: %v", err) } if err := secondBinding.Controller().AbortAttempt(context.Background()); err != nil { t.Fatalf("idempotent AbortAttempt: %v", err) } if service.cancelCalls != 2 || service.closeCalls != 2 { t.Fatalf("cancel/close calls = %d/%d, want 2/2", service.cancelCalls, service.closeCalls) } } func TestOpenAIAttemptDispatcherDoesNotStoreAuthInSnapshot(t *testing.T) { service := &dispatcherServiceSpy{} secret := "dispatch-only-secret" rebuilder, ref, dispatcher := newDispatcherFixture(t, service, func(_ context.Context, _ streamgate.RebuiltRequest, body []byte) (openAIAttemptAdmission, error) { return openAIAttemptAdmission{ kind: openAIAdmissionTunnel, tunnel: edgeservice.SubmitProviderTunnelRequest{Path: openAIRebuildEndpointChat, Body: body}, authorize: func(context.Context) (map[string]string, error) { return map[string]string{"Authorization": secret}, nil }, }, nil }) request := rebuiltRequestForDispatcher(t, rebuilder, ref, "plan.auth") binding, err := dispatcher.DispatchAttempt(context.Background(), request) if err != nil { t.Fatalf("DispatchAttempt: %v", err) } canonical, _ := rebuilder.ingress.canonicalBody() semantic, _ := rebuilder.ingress.semanticBody() if bytes.Contains(canonical, []byte(secret)) || bytes.Contains(semantic, []byte(secret)) { t.Fatal("auth secret entered snapshot") } _ = binding.Controller().AbortAttempt(context.Background()) } func TestOpenAIAttemptControllerRecordsUsageOnceAcrossAbortAndClose(t *testing.T) { service := &dispatcherServiceSpy{} recorder := &openAIUsageRecorder{ request: usageRequestLabels{routeModel: "alias", endpoint: usageEndpointChatCompletions}, attempts: make(map[string]usageDispatchBinding), } attemptUsage := &openAIAttemptUsage{} attemptUsage.observe(usageObservation{inputTokens: 4, outputTokens: 2, providerReported: true}) controller := &openAIAttemptController{ service: service, dispatch: edgeservice.RunDispatch{ RunID: "attempt-controller", NodeID: "node.actual", ProviderID: "provider.actual", Target: "model.actual", }, closeTransport: func() { service.closeCalls++ }, usageRecorder: recorder, usageBinding: usageDispatchBinding{ attemptID: "attempt-controller", usageAttribution: "provider", providerID: "provider.actual", servedModel: "model.actual", responseMode: responseModeNormalized, }, usage: attemptUsage, } if err := controller.AbortAttempt(context.Background()); err != nil { t.Fatalf("AbortAttempt: %v", err) } if err := controller.AbortAttempt(context.Background()); err != nil { t.Fatalf("second AbortAttempt: %v", err) } if err := controller.CloseAttempt(context.Background()); err != nil { t.Fatalf("CloseAttempt after abort: %v", err) } recorder.mu.Lock() attempts := len(recorder.attempts) providerUsage := recorder.providerUsage recorder.mu.Unlock() if attempts != 1 || !providerUsage { t.Fatalf("recorded attempts/provider usage = %d/%v, want 1/true", attempts, providerUsage) } if service.cancelCalls != 1 || service.closeCalls != 1 { t.Fatalf("cancel/close calls = %d/%d, want 1/1", service.cancelCalls, service.closeCalls) } } func TestOpenAIAttemptControllerUnobservedAbortEmitsNoProviderUsage(t *testing.T) { service := &dispatcherServiceSpy{} recorder := &openAIUsageRecorder{ request: usageRequestLabels{routeModel: "alias", endpoint: usageEndpointChatCompletions}, attempts: make(map[string]usageDispatchBinding), } controller := &openAIAttemptController{ service: service, dispatch: edgeservice.RunDispatch{ RunID: "attempt-unobserved", NodeID: "node.actual", ProviderID: "provider.actual", Target: "model.actual", }, usageRecorder: recorder, usageBinding: usageDispatchBinding{ attemptID: "attempt-unobserved", usageAttribution: "provider", providerID: "provider.actual", servedModel: "model.actual", responseMode: responseModeNormalized, }, usage: &openAIAttemptUsage{}, } if err := controller.AbortAttempt(context.Background()); err != nil { t.Fatalf("AbortAttempt: %v", err) } recorder.mu.Lock() attempts := len(recorder.attempts) providerUsage := recorder.providerUsage recorder.mu.Unlock() if attempts != 1 || providerUsage { t.Fatalf("recorded attempts/provider usage = %d/%v, want 1/false", attempts, providerUsage) } } func TestActualOpenAIProviderDoesNotFallbackToAdapterOrNode(t *testing.T) { dispatch := edgeservice.RunDispatch{ RunID: "attempt-no-provider", NodeID: "node.must-not-be-provider", Adapter: "adapter.must-not-be-provider", Target: "served-model", } if got := actualOpenAIProvider(dispatch); got != "" { t.Fatalf("actual provider fallback = %q, want empty", got) } recorder := &openAIUsageRecorder{ request: usageRequestLabels{routeModel: "alias", endpoint: usageEndpointChatCompletions}, attempts: make(map[string]usageDispatchBinding), } recorder.RecordAttempt(newUsageDispatchBinding(dispatch, responseModeNormalized), usageObservation{ inputTokens: 9, providerReported: true, }) recorder.mu.Lock() providerUsage := recorder.providerUsage recorder.mu.Unlock() if providerUsage { t.Fatal("missing strict provider id must not create provider-attributed usage") } }