요청 종료 계수와 실제 provider 시도 사용량을 분리하고, 직접·pool·retry 경로의 attribution을 보존한다. 관련 계약·스펙과 완료된 task archive 정리도 함께 반영한다.
357 lines
14 KiB
Go
357 lines
14 KiB
Go
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")
|
|
}
|
|
}
|