@@ -5,13 +5,15 @@ import (
55 "fmt"
66 "time"
77
8+ "github.com/oklog/ulid/v2"
89 "github.com/samber/lo"
910
1011 "github.com/openmeterio/openmeter/openmeter/billing"
1112 "github.com/openmeterio/openmeter/openmeter/billing/charges/flatfee"
1213 "github.com/openmeterio/openmeter/openmeter/billing/charges/meta"
1314 metaadapter "github.com/openmeterio/openmeter/openmeter/billing/charges/meta/adapter"
1415 "github.com/openmeterio/openmeter/openmeter/billing/charges/models/chargemeta"
16+ "github.com/openmeterio/openmeter/openmeter/billing/charges/models/costbasis"
1517 "github.com/openmeterio/openmeter/openmeter/ent/db"
1618 dbchargeflatfee "github.com/openmeterio/openmeter/openmeter/ent/db/chargeflatfee"
1719 dbchargeflatfeeoverride "github.com/openmeterio/openmeter/openmeter/ent/db/chargeflatfeeoverride"
@@ -86,6 +88,10 @@ func (a *adapter) UpdateCharge(ctx context.Context, charge flatfee.ChargeBase) (
8688 return flatfee.ChargeBase {}, err
8789 }
8890
91+ if err := tx .loadCostBasisEdge (ctx , dbUpdatedChargeBase ); err != nil {
92+ return flatfee.ChargeBase {}, err
93+ }
94+
8995 if overrideLayer := charge .Intent .GetOverrideLayerMutableFields (); overrideLayer != nil {
9096 intentOverride , err := tx .updateIntentOverride (ctx , charge .GetChargeID (), overrideLayer , intent .Currency )
9197 if err != nil {
@@ -132,6 +138,9 @@ func (a *adapter) UpdateSubscriptionItemID(ctx context.Context, charge flatfee.C
132138 }
133139
134140 updatedChargeBase .Edges .IntentOverride = override
141+ if err := tx .loadCostBasisEdge (ctx , updatedChargeBase ); err != nil {
142+ return flatfee.Charge {}, err
143+ }
135144 mappedChargeBase , err := fromDBBaseWithCurrency (updatedChargeBase , charge .Intent .GetBaseIntent ().Currency )
136145 if err != nil {
137146 return flatfee.Charge {}, err
@@ -208,19 +217,69 @@ func (a *adapter) CreateCharges(ctx context.Context, in flatfee.CreateChargesInp
208217 }
209218
210219 return entutils .TransactingRepo (ctx , a , func (ctx context.Context , tx * adapter ) ([]flatfee.Charge , error ) {
211- creates , err := slicesx .MapWithErr (in .Intents , func (intent flatfee.IntentWithInitialStatus ) (* db.ChargeFlatFeeCreate , error ) {
212- return tx .buildCreateFlatFeeCharge (in .Namespace , intent )
220+ type preparedCreate struct {
221+ costBasis * db.ChargeFlatFeeCostBasisCreate
222+ charge * db.ChargeFlatFeeCreate
223+ }
224+
225+ preparedCreates := make ([]preparedCreate , 0 , len (in .Intents ))
226+ for _ , intent := range in .Intents {
227+ chargeCreate , err := tx .buildCreateFlatFeeCharge (in .Namespace , intent )
228+ if err != nil {
229+ return nil , err
230+ }
231+
232+ var costBasisCreate * db.ChargeFlatFeeCostBasisCreate
233+ if intent .Intent .CostBasis != nil {
234+ costBasisCreate , err = costbasis .Create (tx .db .ChargeFlatFeeCostBasis .Create (), costbasis.CreateInput {
235+ NamespacedID : models.NamespacedID {
236+ Namespace : in .Namespace ,
237+ ID : ulid .Make ().String (),
238+ },
239+ CurrencyID : intent .Intent .Currency .ID ,
240+ Intent : * intent .Intent .CostBasis ,
241+ State : intent .ResolvedCostBasis ,
242+ })
243+ if err != nil {
244+ return nil , fmt .Errorf ("building flat fee cost basis: %w" , err )
245+ }
246+ }
247+
248+ preparedCreates = append (preparedCreates , preparedCreate {
249+ costBasis : costBasisCreate ,
250+ charge : chargeCreate ,
251+ })
252+ }
253+
254+ costBasisCreates := lo .Filter (preparedCreates , func (create preparedCreate , _ int ) bool {
255+ return create .costBasis != nil
213256 })
214- if err != nil {
215- return nil , err
257+
258+ var createdCostBases []* db.ChargeFlatFeeCostBasis
259+ if len (costBasisCreates ) > 0 {
260+ var err error
261+ createdCostBases , err = tx .db .ChargeFlatFeeCostBasis .CreateBulk (
262+ lo .Map (costBasisCreates , func (create preparedCreate , _ int ) * db.ChargeFlatFeeCostBasisCreate {
263+ return create .costBasis
264+ })... ,
265+ ).Save (ctx )
266+ if err != nil {
267+ return nil , fmt .Errorf ("creating flat fee cost bases: %w" , err )
268+ }
269+
270+ lo .ForEach (costBasisCreates , func (create preparedCreate , idx int ) {
271+ create .charge .SetCostBasisID (createdCostBases [idx ].ID )
272+ })
216273 }
217274
218- entities , err := tx .db .ChargeFlatFee .CreateBulk (creates ... ).Save (ctx )
275+ chargeCreates := lo .Map (preparedCreates , func (create preparedCreate , _ int ) * db.ChargeFlatFeeCreate {
276+ return create .charge
277+ })
278+ entities , err := tx .db .ChargeFlatFee .CreateBulk (chargeCreates ... ).Save (ctx )
219279 if err != nil {
220280 return nil , metaadapter .MapChargeConstraintError (err )
221281 }
222282
223- // Let's reserve the charge IDs
224283 err = tx .metaAdapter .RegisterCharges (ctx , meta.RegisterChargesInput {
225284 Namespace : in .Namespace ,
226285 Type : meta .ChargeTypeFlatFee ,
@@ -235,7 +294,20 @@ func (a *adapter) CreateCharges(ctx context.Context, in flatfee.CreateChargesInp
235294 return nil , err
236295 }
237296
297+ costBasisByID := lo .SliceToMap (createdCostBases , func (entity * db.ChargeFlatFeeCostBasis ) (string , * db.ChargeFlatFeeCostBasis ) {
298+ return entity .ID , entity
299+ })
300+
238301 return lo .MapErr (entities , func (entity * db.ChargeFlatFee , idx int ) (flatfee.Charge , error ) {
302+ if entity .CostBasisID != nil {
303+ createdCostBasis , ok := costBasisByID [* entity .CostBasisID ]
304+ if ! ok {
305+ return flatfee.Charge {}, fmt .Errorf ("created flat fee cost basis %s not found" , * entity .CostBasisID )
306+ }
307+
308+ entity .Edges .CostBasis = createdCostBasis
309+ }
310+
239311 return FromDBWithCurrency (entity , in .Intents [idx ].Intent .Currency , meta .ExpandNone )
240312 })
241313 })
@@ -251,7 +323,8 @@ func (a *adapter) GetByIDs(ctx context.Context, input flatfee.GetByIDsInput) ([]
251323 Where (dbchargeflatfee .Namespace (input .Namespace )).
252324 Where (dbchargeflatfee .IDIn (input .IDs ... )).
253325 WithIntentOverride ().
254- WithCustomCurrency ()
326+ WithCustomCurrency ().
327+ WithCostBasis ()
255328
256329 if input .Expands .Has (meta .ExpandRealizations ) {
257330 query = expandRealizations (query )
@@ -294,7 +367,8 @@ func (a *adapter) GetByID(ctx context.Context, input flatfee.GetByIDInput) (flat
294367 Where (dbchargeflatfee .Namespace (input .ChargeID .Namespace )).
295368 Where (dbchargeflatfee .ID (input .ChargeID .ID )).
296369 WithIntentOverride ().
297- WithCustomCurrency ()
370+ WithCustomCurrency ().
371+ WithCostBasis ()
298372
299373 if input .Expands .Has (meta .ExpandRealizations ) {
300374 query = expandRealizations (query )
0 commit comments