From d8d031690cfd48ac40ccbc13e89eec4fee9e7401 Mon Sep 17 00:00:00 2001 From: yokowu <18836617@qq.com> Date: Thu, 24 Sep 2026 14:15:12 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=EF=BC=9A=E8=87=AA=E9=85=8D?= =?UTF-8?q?=E6=A8=A1=E5=9E=8B=E7=BB=95=E8=BF=87=E8=AE=A1=E8=B4=B9=E5=B9=B6?= =?UTF-8?q?=E4=BF=9D=E7=95=99=E7=94=A8=E9=87=8F=E7=BB=9F=E8=AE=A1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- monkeyai/backend/internal/app/app.go | 3 +- monkeyai/backend/internal/app/billing.go | 13 ++ .../backend/internal/imagegen/postgres.go | 6 +- .../internal/imagegen/postgres_test.go | 13 ++ monkeyai/backend/internal/imagegen/service.go | 39 +++--- .../backend/internal/imagegen/service_test.go | 36 +++++- monkeyai/backend/internal/model/model.go | 1 + .../backend/internal/model/postgres_test.go | 24 ++++ monkeyai/backend/internal/model/query.sql | 11 ++ monkeyai/backend/internal/model/service.go | 1 + .../backend/internal/model/sqlc/query.sql.go | 41 ++++++ monkeyai/backend/internal/model/usage.go | 30 +++++ monkeyai/backend/internal/proxy/billing.go | 34 ++++- .../backend/internal/proxy/billing_test.go | 120 +++++++++++++++++- monkeyai/backend/internal/proxy/proxy.go | 11 +- monkeyai/backend/internal/proxy/usage.go | 2 +- 16 files changed, 361 insertions(+), 24 deletions(-) create mode 100644 monkeyai/backend/internal/model/usage.go diff --git a/monkeyai/backend/internal/app/app.go b/monkeyai/backend/internal/app/app.go index 97740e5b8..37eb12221 100644 --- a/monkeyai/backend/internal/app/app.go +++ b/monkeyai/backend/internal/app/app.go @@ -189,7 +189,7 @@ func newApplicationHandler(ctx context.Context, logger *slog.Logger, pool *pgxpo charges.RegisterAgent(agent) router := chi.NewRouter() - modelProxy := proxy.NewProxy(modelResolver{service: models}, logger).WithBilling(modelBilling{service: charges}) + modelProxy := proxy.NewProxy(modelResolver{service: models}, logger).WithBilling(modelBilling{service: charges}).WithUsageRecorder(modelUsageRecorder{models: modelRepo}) modelProxy.Register(router) imageproxy.NewProxy(modelResolver{service: models}, keys, imageService, imageService, imageService). WithInputs(imageUploader{inputs: imageInputs}).WithOutputs(imageOutputs).Register(router) @@ -225,6 +225,7 @@ func (r modelResolver) Resolve(ctx context.Context, credential, requestedModel s } return proxy.Target{ ModelID: target.ID, + OwnershipType: target.OwnershipType, UpstreamModel: target.UpstreamModelID, Protocol: string(target.Protocol), UserID: target.UserID, diff --git a/monkeyai/backend/internal/app/billing.go b/monkeyai/backend/internal/app/billing.go index d16cd168c..0f0c33ec9 100644 --- a/monkeyai/backend/internal/app/billing.go +++ b/monkeyai/backend/internal/app/billing.go @@ -25,6 +25,19 @@ type applicationHandler struct { endpoints *endpoint.Service } type modelBilling struct{ service *billing.Service } +type modelUsageRecorder struct{ models *model.Postgres } + +func (r modelUsageRecorder) Record(ctx context.Context, c proxy.Call) error { + if c.InputTokens > math.MaxInt64 || c.CachedInputTokens > math.MaxInt64 || c.OutputTokens > math.MaxInt64 || c.CachedInputTokens > c.InputTokens { + return errors.New("模型用量无效") + } + return r.models.RecordCall(ctx, model.Call{ + ModelID: c.ModelID, UserID: c.UserID, RequestID: c.RequestID, + Status: c.Result, ErrorCode: c.ErrorCode, + InputTokens: int64(c.InputTokens), CachedInputTokens: int64(c.CachedInputTokens), + OutputTokens: int64(c.OutputTokens), StartedAt: c.StartedAt, CompletedAt: c.CompletedAt, + }) +} type reconciliationModels interface { Get(context.Context, string) (model.Model, error) diff --git a/monkeyai/backend/internal/imagegen/postgres.go b/monkeyai/backend/internal/imagegen/postgres.go index 516b15a0f..715f4c9a9 100644 --- a/monkeyai/backend/internal/imagegen/postgres.go +++ b/monkeyai/backend/internal/imagegen/postgres.go @@ -123,8 +123,12 @@ func (p *Postgres) Get(ctx context.Context, userID, id string) (Job, error) { } func (p *Postgres) Reserve(ctx context.Context, job Job, transactionID string) error { + var billingID *string + if transactionID != "" { + billingID = &transactionID + } count, err := sqlc.New(p.pool).SetJobReservation(ctx, sqlc.SetJobReservationParams{ - ID: job.ID, UserID: job.UserID, BillingTransactionID: &transactionID, + ID: job.ID, UserID: job.UserID, BillingTransactionID: billingID, }) return one(count, err) } diff --git a/monkeyai/backend/internal/imagegen/postgres_test.go b/monkeyai/backend/internal/imagegen/postgres_test.go index f73086ed3..e22ebdcb0 100644 --- a/monkeyai/backend/internal/imagegen/postgres_test.go +++ b/monkeyai/backend/internal/imagegen/postgres_test.go @@ -108,6 +108,19 @@ VALUES($1,'user',$2,'image-model','Image Model','image_generation','image','open if _, err := repo.Get(ctx, otherID, created.ID); !errors.Is(err, resource.NotFound) { t.Fatalf("跨用户任务查询未阻止: %v", err) } + freeJob := job + freeJob.RequestHash, freeJob.IdempotencyKey = "hash-free", "request-free" + freeJob, _, err = repo.Create(ctx, freeJob) + if err != nil { + t.Fatal(err) + } + if err := repo.Reserve(ctx, freeJob, ""); err != nil { + t.Fatal(err) + } + freeJob, err = repo.Get(ctx, userID, freeJob.ID) + if err != nil || freeJob.Status != "reserved" || freeJob.BillingTransactionID != nil { + t.Fatalf("自配生图任务应无计费交易: %+v, %v", freeJob, err) + } var imageBytes bytes.Buffer img := image.NewRGBA(image.Rect(0, 0, 2, 3)) diff --git a/monkeyai/backend/internal/imagegen/service.go b/monkeyai/backend/internal/imagegen/service.go index 168d1f0b1..97c4b4bfd 100644 --- a/monkeyai/backend/internal/imagegen/service.go +++ b/monkeyai/backend/internal/imagegen/service.go @@ -218,9 +218,12 @@ func (s *Service) submit(ctx context.Context, target proxy.Target, operation, re (mask != nil && (!cap.SupportsMask || operation != "edit")) { return imageproxy.Task{}, resource.Invalid("参考图或编辑参数不受支持") } - unit, err := Quote(item, operation, quality, aspect) - if err != nil { - return imageproxy.Task{}, resource.Invalid(err.Error()) + var unit billing.Amount + if item.OwnershipType != "user" { + unit, err = Quote(item, operation, quality, aspect) + if err != nil { + return imageproxy.Task{}, resource.Invalid(err.Error()) + } } images := make([]Input, 0, len(references)) fileIDs := make([]string, 0, len(references)+1) @@ -269,22 +272,28 @@ func (s *Service) submit(ctx context.Context, target proxy.Target, operation, re s.failUnsubmitted(ctx, job, "invalid_reference") return imageproxy.Task{}, err } - reservation, err := s.billing.Begin(ctx, billing.Request{ - UserID: target.UserID, ResourceID: item.ID, Category: "image", ImageCount: int64(imageCount), - ImageUnitPrice: unit, IdempotencyKey: job.ID, RequestHash: job.RequestHash, - }) - if err != nil { - s.failUnsubmitted(ctx, job, "billing_failed") - return imageproxy.Task{}, err + var reservationID string + if item.OwnershipType != "user" { + reservation, err := s.billing.Begin(ctx, billing.Request{ + UserID: target.UserID, ResourceID: item.ID, Category: "image", ImageCount: int64(imageCount), + ImageUnitPrice: unit, IdempotencyKey: job.ID, RequestHash: job.RequestHash, + }) + if err != nil { + s.failUnsubmitted(ctx, job, "billing_failed") + return imageproxy.Task{}, err + } + reservationID = reservation.ID + job.BillingTransactionID = &reservationID } - job.BillingTransactionID = &reservation.ID - if err := s.jobs.Reserve(ctx, job, reservation.ID); err != nil { + if err := s.jobs.Reserve(ctx, job, reservationID); err != nil { s.failUnsubmitted(ctx, job, "reservation_failed") return imageproxy.Task{}, err } - if err := s.billing.Start(ctx, reservation.ID); err != nil { - s.failUnsubmitted(ctx, job, "billing_start_failed") - return imageproxy.Task{}, err + if reservationID != "" { + if err := s.billing.Start(ctx, reservationID); err != nil { + s.failUnsubmitted(ctx, job, "billing_start_failed") + return imageproxy.Task{}, err + } } if err := s.jobs.Submitted(ctx, job.ID); err != nil { s.failUnsubmitted(ctx, job, "submit_failed") diff --git a/monkeyai/backend/internal/imagegen/service_test.go b/monkeyai/backend/internal/imagegen/service_test.go index f1e212399..75d6842ff 100644 --- a/monkeyai/backend/internal/imagegen/service_test.go +++ b/monkeyai/backend/internal/imagegen/service_test.go @@ -51,7 +51,9 @@ func (*testJobs) LinkInputs(context.Context, string, []string) error { return ni func (s *testJobs) Reserve(_ context.Context, _ Job, tx string) error { s.Lock() defer s.Unlock() - s.job.BillingTransactionID = &tx + if tx != "" { + s.job.BillingTransactionID = &tx + } s.job.Status = "reserved" return nil } @@ -174,6 +176,38 @@ func TestCapabilitiesUsesProviderDefaultWithoutModel(t *testing.T) { } } +func TestUserImageModelSkipsBillingAndKeepsJobUsage(t *testing.T) { + var imageBytes bytes.Buffer + if err := png.Encode(&imageBytes, image.NewRGBA(image.Rect(0, 0, 2, 2))); err != nil { + t.Fatal(err) + } + svc, jobs, _ := testImageService(generatorFunc(func(_ context.Context, _ proxy.Target, _ ProviderRequest) (ProviderResult, error) { + return ProviderResult{Status: "succeeded", Images: []Image{{Data: imageBytes.Bytes(), Width: 2, Height: 2}}}, nil + })) + svc.models = modelReaderFunc(func(_ context.Context, id string) (model.Model, error) { + return model.Model{ID: id, ModelID: "upstream", OwnershipType: "user", Kind: model.KindImage, + Protocol: model.ProtocolImage, Provider: model.ProviderOpenAIImages, + ImageConfig: &model.ImageConfig{Qualities: []string{"1K"}, AspectRatios: []string{"1:1"}, DefaultQuality: "1K", DefaultAspectRatio: "1:1"}, + ImagePricing: &model.ImagePricing{BaseCreditsPerImage: "0"}}, nil + }) + svc.billing = nil + target := proxy.Target{ModelID: "model-1", UpstreamModel: "upstream", UserID: "owner", Protocol: "image_generation"} + result, err := svc.Generate(t.Context(), target, imageproxy.GenerateRequest{Model: "image@model-1", Prompt: "猫"}) + if err != nil || result.Status != "pending" { + t.Fatalf("自配生图受理失败: %+v, %v", result, err) + } + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + if err := svc.Wait(ctx); err != nil { + t.Fatal(err) + } + jobs.Lock() + defer jobs.Unlock() + if jobs.job.Status != "succeeded" || jobs.job.GeneratedImages != 1 || jobs.job.BillingTransactionID != nil { + t.Fatalf("自配生图应保留任务统计但不创建计费交易: %+v", jobs.job) + } +} + func TestGenerateReservesAndSettlesActualImages(t *testing.T) { var imageBytes bytes.Buffer if err := png.Encode(&imageBytes, image.NewRGBA(image.Rect(0, 0, 2, 2))); err != nil { diff --git a/monkeyai/backend/internal/model/model.go b/monkeyai/backend/internal/model/model.go index 1d632e0cf..d6119da0b 100644 --- a/monkeyai/backend/internal/model/model.go +++ b/monkeyai/backend/internal/model/model.go @@ -134,6 +134,7 @@ type AgentModel struct { type Target struct { ID string UserID string + OwnershipType string UpstreamModelID string Protocol Protocol Kind Kind diff --git a/monkeyai/backend/internal/model/postgres_test.go b/monkeyai/backend/internal/model/postgres_test.go index ca20aef40..dd2c40c6f 100644 --- a/monkeyai/backend/internal/model/postgres_test.go +++ b/monkeyai/backend/internal/model/postgres_test.go @@ -7,6 +7,7 @@ import ( "path/filepath" "strings" "testing" + "time" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/resource" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/rootgroup" @@ -133,8 +134,31 @@ func TestResolveModelName(t *testing.T) { if err != nil || target.ID != tc.want || target.UserID != tc.user || target.UpstreamModelID == "" || target.APIKey != "test-key" { t.Fatalf("模型解析错误: %+v, %v", target, err) } + if tc.name == "自有模型" && target.OwnershipType != "user" || tc.name == "模型名" && target.OwnershipType != "system" { + t.Fatalf("模型归属类型未传递: %+v", target) + } }) } + t.Run("自配模型独立记录用量", func(t *testing.T) { + call := Call{ModelID: private.ID, UserID: owner, RequestID: "own_response", Status: "succeeded", + InputTokens: 11, CachedInputTokens: 5, OutputTokens: 7, + StartedAt: time.Now().Add(-time.Second), CompletedAt: time.Now()} + if err := repo.RecordCall(ctx, call); err != nil { + t.Fatal(err) + } + var count, input, cached, output int + if err := pool.QueryRow(ctx, `SELECT count(*),sum(input_tokens),sum(cached_input_tokens),sum(output_tokens) FROM model_calls WHERE model_id=$1`, private.ID).Scan(&count, &input, &cached, &output); err != nil || count != 1 || input != 11 || cached != 5 || output != 7 { + t.Fatalf("独立用量记录不正确: %d %d %d %d, %v", count, input, cached, output, err) + } + var transactions int + if err := pool.QueryRow(ctx, `SELECT count(*) FROM billing_transactions WHERE resource_id=$1`, private.ID).Scan(&transactions); err != nil || transactions != 0 { + t.Fatalf("自配模型不应有计费交易: %d, %v", transactions, err) + } + call.ModelID = shared.ID + if !errors.Is(repo.RecordCall(ctx, call), ErrNotFound) { + t.Fatal("系统模型不应由独立用量写入器记录") + } + }) models, err := NewService(repo).AgentModels(ctx, user, false) if err != nil || len(models) != 2 { t.Fatalf("下发模型目录错误: %+v, %v", models, err) diff --git a/monkeyai/backend/internal/model/query.sql b/monkeyai/backend/internal/model/query.sql index c99c1b8e2..69df98976 100644 --- a/monkeyai/backend/internal/model/query.sql +++ b/monkeyai/backend/internal/model/query.sql @@ -290,3 +290,14 @@ ORDER BY 1, 3, 2; + +-- name: RecordModelCall :execresult +INSERT INTO model_calls (user_id, model_id, request_id, status, input_tokens, cached_input_tokens, + output_tokens, cache_hit, error_code, started_at, completed_at) +SELECT sqlc.arg(user_id)::uuid, m.id, NULLIF(sqlc.arg(request_id)::text, ''), + sqlc.arg(status)::text, sqlc.arg(input_tokens)::bigint, sqlc.arg(cached_input_tokens)::bigint, + sqlc.arg(output_tokens)::bigint, sqlc.arg(cached_input_tokens)::bigint > 0, + NULLIF(sqlc.arg(error_code)::text, ''), sqlc.arg(started_at)::timestamptz, + sqlc.arg(completed_at)::timestamptz +FROM models m +WHERE m.id = sqlc.arg(model_id)::uuid AND m.ownership_type = 'user'; diff --git a/monkeyai/backend/internal/model/service.go b/monkeyai/backend/internal/model/service.go index 2512f1e9b..f5f4cea0a 100644 --- a/monkeyai/backend/internal/model/service.go +++ b/monkeyai/backend/internal/model/service.go @@ -222,6 +222,7 @@ func (s *Service) Resolve(ctx context.Context, credential, requestedModel string return Target{ ID: item.ID, UserID: userID, + OwnershipType: item.OwnershipType, UpstreamModelID: item.ModelID, Protocol: item.Protocol, Kind: item.Kind, diff --git a/monkeyai/backend/internal/model/sqlc/query.sql.go b/monkeyai/backend/internal/model/sqlc/query.sql.go index 518a353b9..388a8bcef 100644 --- a/monkeyai/backend/internal/model/sqlc/query.sql.go +++ b/monkeyai/backend/internal/model/sqlc/query.sql.go @@ -7,6 +7,7 @@ package sqlc import ( "context" + "time" "github.com/jackc/pgx/v5/pgconn" ) @@ -584,6 +585,46 @@ func (q *Queries) LockOwned(ctx context.Context, arg LockOwnedParams) (string, e return id, err } +const recordModelCall = `-- name: RecordModelCall :execresult +INSERT INTO model_calls (user_id, model_id, request_id, status, input_tokens, cached_input_tokens, + output_tokens, cache_hit, error_code, started_at, completed_at) +SELECT $1::uuid, m.id, NULLIF($2::text, ''), + $3::text, $4::bigint, $5::bigint, + $6::bigint, $5::bigint > 0, + NULLIF($7::text, ''), $8::timestamptz, + $9::timestamptz +FROM models m +WHERE m.id = $10::uuid AND m.ownership_type = 'user' +` + +type RecordModelCallParams struct { + UserID string + RequestID string + Status string + InputTokens int64 + CachedInputTokens int64 + OutputTokens int64 + ErrorCode string + StartedAt time.Time + CompletedAt time.Time + ModelID string +} + +func (q *Queries) RecordModelCall(ctx context.Context, arg RecordModelCallParams) (pgconn.CommandTag, error) { + return q.db.Exec(ctx, recordModelCall, + arg.UserID, + arg.RequestID, + arg.Status, + arg.InputTokens, + arg.CachedInputTokens, + arg.OutputTokens, + arg.ErrorCode, + arg.StartedAt, + arg.CompletedAt, + arg.ModelID, + ) +} + const resolveModel = `-- name: ResolveModel :one WITH RECURSIVE user_groups(group_id) AS ( SELECT id diff --git a/monkeyai/backend/internal/model/usage.go b/monkeyai/backend/internal/model/usage.go new file mode 100644 index 000000000..7fb7b8410 --- /dev/null +++ b/monkeyai/backend/internal/model/usage.go @@ -0,0 +1,30 @@ +package model + +import ( + "context" + "time" + + "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/model/sqlc" +) + +type Call struct { + ModelID, UserID, RequestID, Status, ErrorCode string + InputTokens, CachedInputTokens, OutputTokens int64 + StartedAt, CompletedAt time.Time +} + +func (p *Postgres) RecordCall(ctx context.Context, call Call) error { + result, err := sqlc.New(p.pool).RecordModelCall(ctx, sqlc.RecordModelCallParams{ + ModelID: call.ModelID, UserID: call.UserID, RequestID: call.RequestID, + Status: call.Status, ErrorCode: call.ErrorCode, + InputTokens: call.InputTokens, CachedInputTokens: call.CachedInputTokens, + OutputTokens: call.OutputTokens, StartedAt: call.StartedAt, CompletedAt: call.CompletedAt, + }) + if err != nil { + return err + } + if result.RowsAffected() != 1 { + return ErrNotFound + } + return nil +} diff --git a/monkeyai/backend/internal/proxy/billing.go b/monkeyai/backend/internal/proxy/billing.go index ac5acc30c..513c47abd 100644 --- a/monkeyai/backend/internal/proxy/billing.go +++ b/monkeyai/backend/internal/proxy/billing.go @@ -25,7 +25,37 @@ type Billing interface { func (p *Proxy) WithBilling(b Billing) *Proxy { p.billing = b; return p } func (p *Proxy) finish(ctx context.Context, pc *proxyContext, call Call) { - if p.billing == nil || pc == nil || pc.reservation.ID == "" { + if pc == nil { + return + } + if pc.target.OwnershipType == "user" { + if p.recorder == nil { + return + } + pc.settled.Do(func() { + c, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + call.ModelID, call.UserID = pc.target.ModelID, pc.target.UserID + call.SessionID = pc.target.SessionID + if call.StartedAt.IsZero() { + call.StartedAt = pc.startedAt + } + if call.CompletedAt.IsZero() { + call.CompletedAt = time.Now() + } + if call.Result == "" { + call.Result = "succeeded" + } + if !call.Known && call.ErrorCode == "" { + call.ErrorCode = "usage_unknown" + } + if err := p.recorder.Record(c, call); err != nil { + p.logger.ErrorContext(c, "记录自配模型用量失败", "model_id", call.ModelID, "error", err) + } + }) + return + } + if p.billing == nil || pc.reservation.ID == "" { return } pc.settled.Do(func() { @@ -36,7 +66,7 @@ func (p *Proxy) finish(ctx context.Context, pc *proxyContext, call Call) { } }) } -func billRequest(body []byte, path string, limit int64, stream bool) ([]byte, error) { +func prepareRequest(body []byte, path string, limit int64, stream bool) ([]byte, error) { var data map[string]json.RawMessage if err := json.Unmarshal(body, &data); err != nil { return nil, err diff --git a/monkeyai/backend/internal/proxy/billing_test.go b/monkeyai/backend/internal/proxy/billing_test.go index 61c0b7490..0955c4608 100644 --- a/monkeyai/backend/internal/proxy/billing_test.go +++ b/monkeyai/backend/internal/proxy/billing_test.go @@ -1,13 +1,131 @@ package proxy import ( + "context" "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" "testing" + "time" ) +type billingMustNotRun struct{} + +func (billingMustNotRun) Begin(context.Context, Target, BillingRequest) (Reservation, error) { + panic("自配模型不应预留积分") +} +func (billingMustNotRun) Start(context.Context, string) error { + panic("自配模型不应启动计费") +} +func (billingMustNotRun) Finish(context.Context, string, Call) error { + panic("自配模型不应结算") +} + +func TestUserModelBypassesBillingAndRecordsUsage(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + if !strings.Contains(string(body), `"n":2`) { + t.Errorf("自配模型请求被计费代理修改: %s", body) + } + _, _ = io.WriteString(w, `{"id":"own_request","usage":{"prompt_tokens":11,"completion_tokens":7}}`) + })) + t.Cleanup(upstream.Close) + recorder := &usageRecorderStub{calls: make(chan Call, 1)} + p := NewProxy(ResolverFunc(func(context.Context, string, string) (Target, error) { + target := testTarget(upstream.URL + "/v1") + target.OwnershipType = "user" + return target, nil + }), discardLogger()).WithBilling(billingMustNotRun{}).WithUsageRecorder(recorder) + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"model":"gpt-5","n":2}`)) + req.Header.Set("Authorization", "Bearer own-key") + w := httptest.NewRecorder() + p.ServeHTTP(w, req) + if w.Code != http.StatusOK || w.Header().Get("X-Billing-Transaction-ID") != "" { + t.Fatalf("自配模型响应错误: %d %s", w.Code, w.Body.String()) + } + select { + case call := <-recorder.calls: + if call.ModelID != "model-1" || call.InputTokens != 11 || call.OutputTokens != 7 || call.Result != "succeeded" { + t.Fatalf("自配模型用量未记录: %+v", call) + } + case <-time.After(time.Second): + t.Fatal("自配模型用量未记录") + } + p.captures.Wait() + select { + case <-recorder.calls: + t.Fatal("自配模型用量重复记录") + default: + } +} + +func TestUserModelRecordsFailedCallWithoutBilling(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusBadGateway) + })) + t.Cleanup(upstream.Close) + recorder := &usageRecorderStub{calls: make(chan Call, 1)} + p := NewProxy(ResolverFunc(func(context.Context, string, string) (Target, error) { + target := testTarget(upstream.URL + "/v1") + target.OwnershipType = "user" + return target, nil + }), discardLogger()).WithBilling(billingMustNotRun{}).WithUsageRecorder(recorder) + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"model":"gpt-5"}`)) + req.Header.Set("Authorization", "Bearer own-key") + w := httptest.NewRecorder() + p.ServeHTTP(w, req) + if w.Code != http.StatusBadGateway { + t.Fatalf("上游错误响应丢失: %d", w.Code) + } + select { + case call := <-recorder.calls: + if call.Result != "failed" || call.ErrorCode != "upstream_http_error" || call.InputTokens != 0 { + t.Fatalf("自配模型失败记录不正确: %+v", call) + } + case <-time.After(time.Second): + t.Fatal("未记录自配模型失败") + } +} + +func TestUserModelStreamStillRequestsUsage(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + if !strings.Contains(string(body), `"include_usage":true`) || !strings.Contains(string(body), `"max_tokens":4096`) { + t.Errorf("自配模型流式请求未保留参数或用量采集: %s", body) + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = io.WriteString(w, "data: {\"id\":\"own_stream\",\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2}}\n\ndata: [DONE]\n\n") + })) + t.Cleanup(upstream.Close) + recorder := &usageRecorderStub{calls: make(chan Call, 1)} + p := NewProxy(ResolverFunc(func(context.Context, string, string) (Target, error) { + target := testTarget(upstream.URL + "/v1") + target.OwnershipType = "user" + return target, nil + }), discardLogger()).WithBilling(billingMustNotRun{}).WithUsageRecorder(recorder) + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"model":"gpt-5","stream":true,"max_tokens":4096}`)) + req.Header.Set("Authorization", "Bearer own-key") + w := httptest.NewRecorder() + p.ServeHTTP(w, req) + p.captures.Wait() + if w.Code != http.StatusOK { + t.Fatalf("流式代理失败: %d", w.Code) + } + select { + case call := <-recorder.calls: + if call.InputTokens != 3 || call.OutputTokens != 2 { + t.Fatalf("流式用量未记录: %+v", call) + } + default: + t.Fatal("流式用量未记录") + } +} + func TestEnforcedRequestLimits(t *testing.T) { for _, path := range []string{"/v1/chat/completions", "/v1/responses", "/v1/messages"} { - out, err := billRequest([]byte(`{"max_tokens":99999,"max_completion_tokens":99999}`), path, 100, true) + out, err := prepareRequest([]byte(`{"max_tokens":99999,"max_completion_tokens":99999}`), path, 100, true) if err != nil { t.Fatal(err) } diff --git a/monkeyai/backend/internal/proxy/proxy.go b/monkeyai/backend/internal/proxy/proxy.go index e5f23c3f2..e52000acb 100644 --- a/monkeyai/backend/internal/proxy/proxy.go +++ b/monkeyai/backend/internal/proxy/proxy.go @@ -32,6 +32,7 @@ var endpoints = []struct { type Target struct { ModelID string + OwnershipType string UpstreamModel string Protocol string UserID string @@ -168,13 +169,13 @@ func (p *Proxy) ServeHTTP(w http.ResponseWriter, r *http.Request) { } } var reservation Reservation - if p.billing != nil { + if p.billing != nil && target.OwnershipType != "user" { reservation, err = p.billing.Begin(r.Context(), target, BillingRequest{Path: r.URL.Path, Body: body, IdempotencyKey: r.Header.Get("Idempotency-Key"), SessionID: r.Header.Get("X-Session-ID")}) if err != nil { billingError(w, err) return } - body, err = billRequest(body, r.URL.Path, reservation.OutputLimit, meta.Stream) + body, err = prepareRequest(body, r.URL.Path, reservation.OutputLimit, meta.Stream) if err != nil { pc := &proxyContext{reservation: reservation, stream: meta.Stream} p.finish(r.Context(), pc, Call{Known: true, Stream: meta.Stream, Result: "failed", ErrorCode: "invalid_request"}) @@ -186,6 +187,12 @@ func (p *Proxy) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } w.Header().Set("X-Billing-Transaction-ID", reservation.ID) + } else if target.OwnershipType == "user" && meta.Stream && r.URL.Path == "/v1/chat/completions" { + body, err = prepareRequest(body, r.URL.Path, 0, true) + if err != nil { + billingError(w, err) + return + } } r.Body = io.NopCloser(bytes.NewReader(body)) r.ContentLength = int64(len(body)) diff --git a/monkeyai/backend/internal/proxy/usage.go b/monkeyai/backend/internal/proxy/usage.go index d1c797bad..f80ddf153 100644 --- a/monkeyai/backend/internal/proxy/usage.go +++ b/monkeyai/backend/internal/proxy/usage.go @@ -155,7 +155,7 @@ func (p *Proxy) recordUsage(ctx context.Context, proxyCtx *proxyContext, result ) } p.finish(ctx, proxyCtx, call) - if p.recorder == nil || !result.hasTokens() { + if target.OwnershipType == "user" || p.billing != nil || p.recorder == nil || !result.hasTokens() { return } if err := p.recorder.Record(context.WithoutCancel(ctx), call); err != nil {