diff --git a/core/orch_test.go b/core/orch_test.go index 22edcf021f..ccb4de3d5b 100644 --- a/core/orch_test.go +++ b/core/orch_test.go @@ -1752,6 +1752,79 @@ func TestBYOCExternalCapsSenderPricing(t *testing.T) { assert.Equal(t, int64(10), getBYOCPrice(addr3), "falls back to default") } +func TestPriceInfoForCaps_BYOCUsesJobPrice(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + + n, _ := NewLivepeerNode(nil, "", nil) + n.AutoAdjustPrice = false + n.SetBasePrice("default", NewFixedPrice(big.NewRat(1, 1))) + n.Recipient = new(pm.MockRecipient) + orch := NewOrchestrator(n, nil) + + sender := ethcommon.HexToAddress("0x1000000000000000000000000000000000000000") + n.ExternalCapabilities.Capabilities["flux-schnell"] = &ExternalCapability{Name: "flux-schnell"} + n.SetPriceForExternalCapability("default", "flux-schnell", big.NewRat(42, 1)) + n.SetPriceForExternalCapability(sender.Hex(), "flux-schnell", big.NewRat(99, 1)) + // Built-in cap price must not be used for BYOC. + n.SetBasePriceForCap("default", Capability_BYOC, "flux-schnell", NewFixedPrice(big.NewRat(7, 1))) + + byocCaps := NewCapabilities([]Capability{Capability_BYOC}, nil) + byocCaps.SetPerCapabilityConstraints(PerCapabilityConstraints{ + Capability_BYOC: &CapabilityConstraints{ + Models: map[string]*ModelConstraint{ + "flux-schnell": {Warm: true, Capacity: 1}, + }, + }, + }) + netCaps := byocCaps.ToNetCapabilities() + + price, err := orch.PriceInfoForCaps(sender, "", netCaps) + require.Nil(err) + require.NotNil(price) + assert.Equal(int64(99), price.PricePerUnit) + assert.Equal(int64(1), price.PixelsPerUnit) + + other := ethcommon.HexToAddress("0x2000000000000000000000000000000000000000") + price, err = orch.PriceInfoForCaps(other, "", netCaps) + require.Nil(err) + require.NotNil(price) + assert.Equal(int64(42), price.PricePerUnit, "falls back to default job price") + + // Session-pinned price wins over job price. + n.Balances = NewAddressBalances(time.Minute) + n.Balances.Credit(sender, ManifestID("sess-1"), big.NewRat(0, 1)) + n.Balances.SetFixedPrice(sender, ManifestID("sess-1"), big.NewRat(5, 2)) + price, err = orch.PriceInfoForCaps(sender, ManifestID("sess-1"), netCaps) + require.Nil(err) + require.NotNil(price) + assert.Equal(int64(5), price.PricePerUnit) + assert.Equal(int64(2), price.PixelsPerUnit) + + // TicketParams must be minted from the same PriceInfoForCaps rate (not base/cap). + recipient := n.Recipient.(*pm.MockRecipient) + jobPrice := big.NewRat(99, 1) + recipient.On("TicketParams", sender, mock.MatchedBy(func(p *big.Rat) bool { + return p != nil && p.Cmp(jobPrice) == 0 + })).Return(&pm.TicketParams{ + Recipient: sender, + FaceValue: big.NewInt(100), + WinProb: big.NewInt(100), + RecipientRandHash: pm.RandHash(), + Seed: big.NewInt(1), + ExpirationBlock: big.NewInt(100), + PricePerPixel: jobPrice, + ExpirationParams: &pm.TicketExpirationParams{}, + }, nil).Once() + + price, err = orch.PriceInfoForCaps(sender, "", netCaps) + require.Nil(err) + params, err := orch.TicketParams(sender, price) + require.Nil(err) + require.NotNil(params) + recipient.AssertExpectations(t) +} + func TestBYOCExternalCapsPriceEdgeCases(t *testing.T) { addr := "0x1000000000000000000000000000000000000000" diff --git a/core/orchestrator.go b/core/orchestrator.go index 0d4d8417a2..da1e3c23cf 100644 --- a/core/orchestrator.go +++ b/core/orchestrator.go @@ -368,11 +368,37 @@ func (orch *orchestrator) PriceInfoForCaps(sender ethcommon.Address, manifestID return nil, nil } + if fixedPrice := orch.sessionFixedPrice(sender, manifestID); fixedPrice != nil { + return priceInfoFromRat(fixedPrice) + } + + // BYOC jobs are priced via GetPriceForJob / JobPriceInfo, not GetBasePriceForCap. + // When GetOrchestrator is called with BYOC caps, TicketParams must be minted at + // that same rate so Payment.ExpectedPrice matches recipientRandHash. + if caps != nil { + coreCaps := CapabilitiesFromNetCapabilities(caps) + if modelID := coreCaps.ModelIDForCapability(Capability_BYOC); modelID != "" { + return orch.JobPriceInfo(sender, modelID) + } + } + price, err := orch.priceInfo(sender, manifestID, caps) if err != nil { return nil, err } + return priceInfoFromRat(price) +} + +func (orch *orchestrator) sessionFixedPrice(sender ethcommon.Address, manifestID ManifestID) *big.Rat { + if manifestID == "" || orch.node.Balances == nil { + return nil + } + + return orch.node.Balances.FixedPrice(sender, manifestID) +} + +func priceInfoFromRat(price *big.Rat) (*net.PriceInfo, error) { if !price.Num().IsInt64() || !price.Denom().IsInt64() { fixedPrice, err := common.PriceToInt64(price) if err != nil { @@ -390,11 +416,8 @@ func (orch *orchestrator) PriceInfoForCaps(sender ethcommon.Address, manifestID // priceInfo returns price per pixel as a fixed point number wrapped in a big.Rat func (orch *orchestrator) priceInfo(sender ethcommon.Address, manifestID ManifestID, caps *net.Capabilities) (*big.Rat, error) { // If there is already a fixed price for the given session, use this price - if manifestID != "" { - fixedPrice := orch.node.Balances.FixedPrice(sender, manifestID) - if fixedPrice != nil { - return fixedPrice, nil - } + if fixedPrice := orch.sessionFixedPrice(sender, manifestID); fixedPrice != nil { + return fixedPrice, nil } transcodePrice := orch.node.GetBasePrice(sender.String()) diff --git a/pm/recipient_test.go b/pm/recipient_test.go index d84f5fab39..e695d17402 100644 --- a/pm/recipient_test.go +++ b/pm/recipient_test.go @@ -182,6 +182,41 @@ func TestReceiveTicket_InvalidSignature(t *testing.T) { assert.False(ok) } +// TestBYOCAlignedPrice_RecipientRandAccepted is the payment-alignment probe: +// TicketParams minted at the orch job price accept matching ExpectedPrice, and +// reject a divergent CapabilitiesPrices rate (the recipientRandHash failure mode). +func TestBYOCAlignedPrice_RecipientRandAccepted(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + + sender, b, _, gm, sm, tm, cfg, sig := newRecipientFixtureOrFatal(t) + sv := &stubSigVerifier{} + sv.SetVerifyResult(true) + v := NewValidator(sv, tm) + secret := [32]byte{3} + r := NewRecipientWithSecret(RandAddress(), b, v, gm, sm, tm, secret, cfg) + + jobPrice := big.NewRat(99, 1) // GetPriceForJob / PriceInfo + divergentCapPrice := big.NewRat(7, 1) // CapabilitiesPrices / base-cap price + + params, err := r.TicketParams(sender, jobPrice) + require.NoError(err) + require.Zero(params.PricePerPixel.Cmp(jobPrice)) + + ticket := newTicket(sender, params, 0) + _, _, err = r.ReceiveTicket(ticket, sig, params.Seed) + require.NoError(err, "payment at orch job price must validate recipientRand") + + // Signer/gateway substituting CapabilitiesPrices into ExpectedPrice. + badTicket := newTicket(sender, params, 1) + badTicket.PricePerPixel = divergentCapPrice + _, _, err = r.ReceiveTicket(badTicket, sig, params.Seed) + require.Error(err) + assert.Equal(errInvalidTicketRecipientRand.Error(), err.Error()) + _, fatal := err.(*FatalReceiveErr) + assert.True(fatal) +} + func TestReceiveTicket_InvalidSender(t *testing.T) { assert := assert.New(t) sender, b, v, gm, sm, tm, cfg, sig := newRecipientFixtureOrFatal(t) diff --git a/server/remote_discovery.go b/server/remote_discovery.go index d2c2b2f626..3b7ff674e2 100644 --- a/server/remote_discovery.go +++ b/server/remote_discovery.go @@ -392,15 +392,8 @@ func capabilityPrice(info *common.OrchNetworkCapabilities, capability core.Capab if info == nil { return nil } - // Check per-capability price if it exists - for _, capPrice := range info.CapabilitiesPrices { - if capPrice == nil || capPrice.PixelsPerUnit <= 0 || core.Capability(capPrice.Capability) != capability { - continue - } - price := new(big.Rat).SetFrac64(capPrice.PricePerUnit, capPrice.PixelsPerUnit) - if capPrice.Constraint == modelID { - return price - } + if capPrice := findCapPriceInfo(info.CapabilitiesPrices, capability, modelID, false); capPrice != nil { + return new(big.Rat).SetFrac64(capPrice.PricePerUnit, capPrice.PixelsPerUnit) } // Global fallback if no per-capability price is available. if info.PriceInfo == nil || info.PriceInfo.PixelsPerUnit <= 0 { diff --git a/server/remote_signer.go b/server/remote_signer.go index f34cdead52..4609a108b9 100644 --- a/server/remote_signer.go +++ b/server/remote_signer.go @@ -18,7 +18,6 @@ import ( "github.com/golang/glog" "github.com/golang/protobuf/proto" "github.com/livepeer/go-livepeer/ai/runner" - "github.com/livepeer/go-livepeer/byoc" "github.com/livepeer/go-livepeer/clog" "github.com/livepeer/go-livepeer/common" "github.com/livepeer/go-livepeer/core" @@ -74,7 +73,9 @@ func (ls *LivepeerServer) SignOrchestratorInfo(w http.ResponseWriter, r *http.Re _ = json.NewEncoder(w).Encode(results) } -// SignBYOCJobRequest signs a BYOC job using the V1 binary format (FlattenBYOCJob). +// SignBYOCJobRequest signs a BYOC job the same way gatewayJob.sign() and +// network orchestrator verifyJobCreds do today: eth-sign over +// request+parameters. (V1 FlattenBYOCJob is not yet deployed on network orchs.) type SignBYOCJobRequestInput struct { ID string `json:"id"` Capability string `json:"capability"` @@ -113,15 +114,7 @@ func (ls *LivepeerServer) SignBYOCJobRequest(w http.ResponseWriter, r *http.Requ return } - sigPayload := byoc.FlattenBYOCJob(&byoc.BYOCJobSigningInput{ - ID: req.ID, - Capability: req.Capability, - Request: req.Request, - Parameters: req.Parameters, - TimeoutSeconds: req.TimeoutSeconds, - }) - - sig, err := gw.Sign(sigPayload) + sig, err := gw.Sign([]byte(req.Request + req.Parameters)) if err != nil { clog.Errorf(ctx, "Failed to sign BYOC job request err=%q", err) respondJsonError(ctx, w, err, http.StatusInternalServerError) @@ -239,13 +232,14 @@ type RemotePaymentRequest struct { // Set if an ID is needed to tie into orch accounting for a session. Optional ManifestID string - // Number of pixels to generate a ticket for. Required if `type` is not set. + // Number of billable units to generate tickets for (e.g. compute-seconds + // for BYOC, or pixels for other callers). Required if `type` is not set. InPixels int64 `json:"inPixels"` // Job type to automatically calculate payments. Valid values: `lv2v`. Optional. Type string `json:"type"` - // Capabilities to include in the ticket. Optional; may be set for the lv2v job type. + // Capabilities to include in segment credentials / max-price policy. Optional. Capabilities []byte `json:"capabilities"` } @@ -355,6 +349,25 @@ func (ls *LivepeerServer) authLivePayment(r *http.Request, state *RemotePaymentS return *webhookResp.Status, &webhookResp, errors.New(webhookResp.Reason) } +// findCapPriceInfo returns the first CapabilitiesPrices entry matching +// capability+modelID. requirePositiveRate skips non-positive PricePerUnit and +// keeps scanning; otherwise a zero rate is allowed. +func findCapPriceInfo(prices []*net.PriceInfo, capability core.Capability, modelID string, requirePositiveRate bool) *net.PriceInfo { + for _, p := range prices { + if p == nil || core.Capability(p.Capability) != capability || p.Constraint != modelID { + continue + } + if p.PixelsPerUnit <= 0 { + continue + } + if requirePositiveRate && p.PricePerUnit <= 0 { + continue + } + return &net.PriceInfo{PricePerUnit: p.PricePerUnit, PixelsPerUnit: p.PixelsPerUnit} + } + return nil +} + // GenerateLivePayment handles remote generation of a payment for live streams. func (ls *LivepeerServer) GenerateLivePayment(w http.ResponseWriter, r *http.Request) { requestID := string(core.RandomManifestID()) @@ -389,14 +402,29 @@ func (ls *LivepeerServer) GenerateLivePayment(w http.ResponseWriter, r *http.Req respondJsonError(ctx, w, err, http.StatusBadRequest) return } - priceInfo := oInfo.PriceInfo - if priceInfo == nil || priceInfo.PricePerUnit == 0 || priceInfo.PixelsPerUnit == 0 { - err := fmt.Errorf("missing or zero priceInfo") + if oInfo.TicketParams == nil { + err := fmt.Errorf("missing ticketParams in OrchestratorInfo") respondJsonError(ctx, w, err, http.StatusBadRequest) return } - if oInfo.TicketParams == nil { - err := fmt.Errorf("missing ticketParams in OrchestratorInfo") + + var reqCaps *core.Capabilities + if len(req.Capabilities) > 0 { + var caps net.Capabilities + if err := proto.Unmarshal(req.Capabilities, &caps); err != nil { + clog.Errorf(ctx, "Failed to unmarshal capabilities err=%q", err) + respondJsonError(ctx, w, err, http.StatusBadRequest) + return + } + reqCaps = core.CapabilitiesFromNetCapabilities(&caps) + } + + // ExpectedPrice must match the rate baked into TicketParams.recipientRandHash. + // Never overwrite PriceInfo from CapabilitiesPrices — the orch is the sole + // rate source (via GetOrchestrator / JobPriceInfo). + priceInfo := oInfo.PriceInfo + if priceInfo == nil || priceInfo.PricePerUnit == 0 || priceInfo.PixelsPerUnit == 0 { + err := fmt.Errorf("missing or zero priceInfo") respondJsonError(ctx, w, err, http.StatusBadRequest) return } @@ -454,16 +482,8 @@ func (ls *LivepeerServer) GenerateLivePayment(w http.ResponseWriter, r *http.Req streamParams := &core.StreamParameters{ // Embedded within genSegCreds, may be used by orch for payment accounting - ManifestID: core.ManifestID(manifestID), - } - if len(req.Capabilities) > 0 { - var caps net.Capabilities - if err := proto.Unmarshal(req.Capabilities, &caps); err != nil { - clog.Errorf(ctx, "Failed to unmarshal capabilities err=%q", err) - respondJsonError(ctx, w, err, http.StatusBadRequest) - return - } - streamParams.Capabilities = core.CapabilitiesFromNetCapabilities(&caps) + ManifestID: core.ManifestID(manifestID), + Capabilities: reqCaps, } pmParams := pmTicketParams(oInfo.TicketParams) @@ -528,7 +548,8 @@ func (ls *LivepeerServer) GenerateLivePayment(w http.ResponseWriter, r *http.Req lastUpdate = now } billableSecs := now.Sub(lastUpdate).Seconds() - if req.Type == RemoteType_LiveVideoToVideo { + switch { + case req.Type == RemoteType_LiveVideoToVideo: info := defaultSegInfo if billableSecs <= 0 { // preload with 60 seconds of data for LV2V @@ -536,13 +557,13 @@ func (ls *LivepeerServer) GenerateLivePayment(w http.ResponseWriter, r *http.Req } pixelsPerSec := float64(info.Height) * float64(info.Width) * float64(info.FPS) pixels = int64(pixelsPerSec * billableSecs) // pixels to charge for - } else if req.Type != "" { + case req.Type != "": err = errors.New("invalid job type") respondJsonError(ctx, w, err, http.StatusBadRequest) return } if pixels <= 0 { - err = errors.New("missing pixels or job type") + err = errors.New("missing billable units or job type") respondJsonError(ctx, w, err, http.StatusBadRequest) return } diff --git a/server/remote_signer_test.go b/server/remote_signer_test.go index cb793b86b5..09e479d5ce 100644 --- a/server/remote_signer_test.go +++ b/server/remote_signer_test.go @@ -1741,3 +1741,147 @@ func TestGetOrchInfoSig_SendsConfiguredHeaders(t *testing.T) { require.Equal([]byte{0x12, 0x34}, []byte(resp.Address)) require.Equal([]byte{0xab, 0xcd}, []byte(resp.Signature)) } + +func byocCaps(t *testing.T, modelID string) *core.Capabilities { + t.Helper() + if modelID == "" { + return nil + } + caps := core.NewCapabilities([]core.Capability{core.Capability_BYOC}, nil) + caps.SetPerCapabilityConstraints(core.PerCapabilityConstraints{ + core.Capability_BYOC: &core.CapabilityConstraints{ + Models: map[string]*core.ModelConstraint{ + modelID: {}, + }, + }, + }) + return caps +} + +func byocCapsBlob(t *testing.T, modelID string) []byte { + t.Helper() + caps := byocCaps(t, modelID) + if caps == nil { + return nil + } + blob, err := proto.Marshal(caps.ToNetCapabilities()) + require.NoError(t, err) + return blob +} + +func TestGenerateLivePayment_ExplicitUnitsKeepOrchPrice(t *testing.T) { + require := require.New(t) + + ethClient := newTestEthClient(t) + node, _ := core.NewLivepeerNode(ethClient, "", nil) + node.Balances = core.NewAddressBalances(1 * time.Minute) + + ev := big.NewRat(10_000_000_000, 1) + var totalTickets uint32 + sender := newMockSender(mockSenderConfig{ + ev: ev, + createTicketBatchFn: func(args mock.Arguments, batch *pm.TicketBatch) { + size := args.Int(1) + *batch = *defaultTicketBatch() + var baseSig []byte + if len(batch.SenderParams) > 0 && batch.SenderParams[0] != nil { + baseSig = batch.SenderParams[0].Sig + } + batch.SenderParams = make([]*pm.TicketSenderParams, size) + for i := 0; i < size; i++ { + totalTickets++ + batch.SenderParams[i] = &pm.TicketSenderParams{SenderNonce: totalTickets, Sig: baseSig} + } + }, + }) + node.Sender = sender + ls := &LivepeerServer{LivepeerNode: node} + + autoPrice, err := core.NewAutoConvertedPrice("", big.NewRat(1_000, 1), nil) + require.NoError(err) + BroadcastCfg.SetMaxPrice(autoPrice) + defer BroadcastCfg.SetMaxPrice(nil) + + const orchPPU, divergentCapPPU, billableUnits = 10, 999, int64(60) + oInfo := &net.OrchestratorInfo{ + Address: ethClient.addr.Bytes(), + PriceInfo: &net.PriceInfo{PricePerUnit: orchPPU, PixelsPerUnit: 1}, + TicketParams: &net.TicketParams{ + Recipient: pm.RandAddress().Bytes(), + }, + AuthToken: stubAuthToken, + CapabilitiesPrices: []*net.PriceInfo{ + {Capability: uint32(core.Capability_BYOC), Constraint: "nano-banana", PricePerUnit: divergentCapPPU, PixelsPerUnit: 1}, + }, + } + + orchBlob, err := proto.Marshal(oInfo) + require.NoError(err) + reqBody, err := json.Marshal(RemotePaymentRequest{ + Orchestrator: orchBlob, + InPixels: billableUnits, + Capabilities: byocCapsBlob(t, "nano-banana"), + }) + require.NoError(err) + req := httptest.NewRequest(http.MethodPost, "/generate-live-payment", bytes.NewReader(reqBody)) + rr := httptest.NewRecorder() + ls.GenerateLivePayment(rr, req) + require.Equal(http.StatusOK, rr.Code, "body=%s", rr.Body.String()) + + var resp RemotePaymentResponse + require.NoError(json.NewDecoder(rr.Body).Decode(&resp)) + var state RemotePaymentState + require.NoError(json.Unmarshal(resp.State.State, &state)) + + require.EqualValues(orchPPU, state.InitialPricePerUnit, "must keep orch PriceInfo; ignore CapabilitiesPrices") + require.EqualValues(1, state.InitialPixelsPerUnit) + + fee := big.NewRat(orchPPU*billableUnits, 1) + smc := fee + if ev.Cmp(smc) > 0 { + smc = ev + } + q := new(big.Rat).Quo(smc, ev) + nt := new(big.Int).Quo(q.Num(), q.Denom()) + if new(big.Int).Rem(q.Num(), q.Denom()).Sign() != 0 { + nt.Add(nt, big.NewInt(1)) + } + wantBal := new(big.Rat).Sub(new(big.Rat).Mul(new(big.Rat).SetInt(nt), ev), fee) + gotBal := new(big.Rat) + _, ok := gotBal.SetString(state.Balance) + require.True(ok) + require.Zero(gotBal.Cmp(wantBal), "balance got=%s want=%s fee=%s", gotBal.RatString(), wantBal.RatString(), fee.RatString()) +} + +func TestGenerateLivePayment_RejectsByocTypeWithoutUnits(t *testing.T) { + require := require.New(t) + + ethClient := newTestEthClient(t) + node, _ := core.NewLivepeerNode(ethClient, "", nil) + node.Balances = core.NewAddressBalances(1 * time.Minute) + node.Sender = newMockSender(mockSenderConfig{ev: big.NewRat(1, 1)}) + ls := &LivepeerServer{LivepeerNode: node} + + oInfo := &net.OrchestratorInfo{ + Address: ethClient.addr.Bytes(), + PriceInfo: &net.PriceInfo{PricePerUnit: 10, PixelsPerUnit: 1}, + TicketParams: &net.TicketParams{ + Recipient: pm.RandAddress().Bytes(), + }, + AuthToken: stubAuthToken, + } + orchBlob, err := proto.Marshal(oInfo) + require.NoError(err) + + reqBody, err := json.Marshal(RemotePaymentRequest{ + Orchestrator: orchBlob, + Type: "byoc", + Capabilities: byocCapsBlob(t, "nano-banana"), + }) + require.NoError(err) + req := httptest.NewRequest(http.MethodPost, "/generate-live-payment", bytes.NewReader(reqBody)) + rr := httptest.NewRecorder() + ls.GenerateLivePayment(rr, req) + require.Equal(http.StatusBadRequest, rr.Code) + require.Contains(rr.Body.String(), "invalid job type") +} diff --git a/server/rpc_test.go b/server/rpc_test.go index bce71a09f8..ad73710dc2 100644 --- a/server/rpc_test.go +++ b/server/rpc_test.go @@ -1843,7 +1843,11 @@ func (o *mockOrchestrator) AuthToken(sessionID string, expiration int64) *net.Au return nil } func (r *mockOrchestrator) PriceInfoForCaps(sender ethcommon.Address, manifestID core.ManifestID, caps *net.Capabilities) (*net.PriceInfo, error) { - return &net.PriceInfo{PricePerUnit: 4, PixelsPerUnit: 1}, nil + args := r.Called(sender, manifestID, caps) + if args.Get(0) != nil { + return args.Get(0).(*net.PriceInfo), args.Error(1) + } + return nil, args.Error(1) } func (r *mockOrchestrator) TextToImage(ctx context.Context, requestID string, req worker.GenTextToImageJSONRequestBody) (interface{}, error) { return nil, nil @@ -2033,6 +2037,7 @@ func TestOrchestratorInfoWithCaps_NonNilEmptyCaps_DoesNotIncludeCapabilitiesPric orch.On("Nodes").Return() orch.On("Address").Return(addr) + orch.On("PriceInfoForCaps", addr, core.ManifestID(""), mock.Anything).Return(&net.PriceInfo{PricePerUnit: 4, PixelsPerUnit: 1}, nil) orch.On("TicketParams", addr, mock.Anything).Return(&net.TicketParams{Recipient: pm.RandBytes(32)}, nil) orch.On("AuthToken", mock.Anything, mock.Anything).Return(&net.AuthToken{Token: []byte("tok"), SessionId: "sess", Expiration: time.Now().Add(time.Hour).Unix()}) @@ -2044,3 +2049,40 @@ func TestOrchestratorInfoWithCaps_NonNilEmptyCaps_DoesNotIncludeCapabilitiesPric orch.AssertNotCalled(t, "GetCapabilitiesPrices", mock.Anything) orch.AssertNotCalled(t, "PriceInfo", mock.Anything) } + +func TestOrchestratorInfoWithCaps_BYOCPricePassedToTicketParams(t *testing.T) { + require := require.New(t) + + oldNodeStorage := drivers.NodeStorage + drivers.NodeStorage = drivers.NewMemoryDriver(nil) + defer func() { drivers.NodeStorage = oldNodeStorage }() + + orch := &mockOrchestrator{} + addr := ethcommon.HexToAddress("0x1") + byocPrice := &net.PriceInfo{PricePerUnit: 99, PixelsPerUnit: 1} + byocCaps := core.NewCapabilities([]core.Capability{core.Capability_BYOC}, nil) + byocCaps.SetPerCapabilityConstraints(core.PerCapabilityConstraints{ + core.Capability_BYOC: &core.CapabilityConstraints{ + Models: map[string]*core.ModelConstraint{ + "flux-schnell": {}, + }, + }, + }) + netCaps := byocCaps.ToNetCapabilities() + + orch.On("Nodes").Return() + orch.On("Address").Return(addr) + orch.On("PriceInfoForCaps", addr, core.ManifestID(""), netCaps).Return(byocPrice, nil) + orch.On("TicketParams", addr, byocPrice).Return(&net.TicketParams{Recipient: pm.RandBytes(32)}, nil) + orch.On("AuthToken", mock.Anything, mock.Anything).Return(&net.AuthToken{ + Token: []byte("tok"), + SessionId: "sess", + Expiration: time.Now().Add(time.Hour).Unix(), + }) + + info, err := orchestratorInfoWithCaps(orch, addr, "https://orch.example.com", "", netCaps) + require.NoError(err) + require.Equal(byocPrice, info.PriceInfo) + orch.AssertCalled(t, "TicketParams", addr, byocPrice) + orch.AssertNotCalled(t, "GetCapabilitiesPrices", mock.Anything) +}