Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion monkeyai/backend/internal/app/app.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down
13 changes: 13 additions & 0 deletions monkeyai/backend/internal/app/billing.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
6 changes: 5 additions & 1 deletion monkeyai/backend/internal/imagegen/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Expand Down
13 changes: 13 additions & 0 deletions monkeyai/backend/internal/imagegen/postgres_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
39 changes: 24 additions & 15 deletions monkeyai/backend/internal/imagegen/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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")
Expand Down
36 changes: 35 additions & 1 deletion monkeyai/backend/internal/imagegen/service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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 {
Expand Down
1 change: 1 addition & 0 deletions monkeyai/backend/internal/model/model.go
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,7 @@ type AgentModel struct {
type Target struct {
ID string
UserID string
OwnershipType string
UpstreamModelID string
Protocol Protocol
Kind Kind
Expand Down
24 changes: 24 additions & 0 deletions monkeyai/backend/internal/model/postgres_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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)
Expand Down
11 changes: 11 additions & 0 deletions monkeyai/backend/internal/model/query.sql
Original file line number Diff line number Diff line change
Expand Up @@ -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';
1 change: 1 addition & 0 deletions monkeyai/backend/internal/model/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
41 changes: 41 additions & 0 deletions monkeyai/backend/internal/model/sqlc/query.sql.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

30 changes: 30 additions & 0 deletions monkeyai/backend/internal/model/usage.go
Original file line number Diff line number Diff line change
@@ -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
}
34 changes: 32 additions & 2 deletions monkeyai/backend/internal/proxy/billing.go
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand All @@ -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
Expand Down
Loading
Loading