diff --git a/charts/telemetry-api/values-prod.yaml b/charts/telemetry-api/values-prod.yaml index 18a4b411..c94d76a0 100644 --- a/charts/telemetry-api/values-prod.yaml +++ b/charts/telemetry-api/values-prod.yaml @@ -21,6 +21,7 @@ env: VEHICLE_NFT_ADDRESS: '0xbA5738a18d83D41847dfFbDC6101d37C69c9B0cF' MANUFACTURER_NFT_ADDRESS: '0x3b07e2A2ABdd0A9B8F7878bdE6487c502164B9dd' FETCH_API_GRPC_ENDPOINT: fetch-api-prod:8086 + CREDIT_TRACKER_ENDPOINT: credit-tracker-prod:8086 IDENTITY_API_URL: http://identity-api-prod:8080/query IDENTITY_API_REQUEST_TIMEOUT_SECONDS: 5 DEVICE_LAST_SEEN_BIN_HOURS: 3 diff --git a/charts/telemetry-api/values.yaml b/charts/telemetry-api/values.yaml index 34753d89..15fae444 100644 --- a/charts/telemetry-api/values.yaml +++ b/charts/telemetry-api/values.yaml @@ -37,6 +37,7 @@ env: VINVC_DATA_VERSION: VINVCv1.0 POMVC_DATA_VERSION: POMVCv1.0 FETCH_API_GRPC_ENDPOINT: fetch-api-dev:8086 + CREDIT_TRACKER_ENDPOINT: credit-tracker-dev:8086 IDENTITY_API_URL: http://identity-api-dev:8080/query IDENTITY_API_REQUEST_TIMEOUT_SECONDS: 5 DEVICE_LAST_SEEN_BIN_HOURS: 3 diff --git a/e2e/credit_tracker_test.go b/e2e/credit_tracker_test.go new file mode 100644 index 00000000..abb6520d --- /dev/null +++ b/e2e/credit_tracker_test.go @@ -0,0 +1,168 @@ +package e2e_test + +import ( + "context" + "fmt" + "net" + "sync" + "testing" + "time" + + ctgrpc "github.com/DIMO-Network/credit-tracker/pkg/grpc" + "github.com/stretchr/testify/require" + "google.golang.org/grpc" +) + +// mockCreditTrackerServer wraps the gRPC server and contains test configuration +type mockCreditTrackerServer struct { + grpcServer *grpc.Server + listener net.Listener + port int + mutex sync.Mutex + responses map[string]map[string]any // method -> request key -> response + ctgrpc.UnimplementedCreditTrackerServer + t *testing.T +} + +// setupCreditTrackerContainer creates and starts a gRPC server on a random available port +func setupCreditTrackerContainer(t *testing.T) *mockCreditTrackerServer { + // Find an available port + listener, err := net.Listen("tcp", ":0") + require.NoError(t, err) + + // Create the gRPC server + grpcServer := grpc.NewServer() + testServer := &mockCreditTrackerServer{ + grpcServer: grpcServer, + t: t, + listener: listener, + port: listener.Addr().(*net.TCPAddr).Port, + responses: make(map[string]map[string]any), + } + + ctgrpc.RegisterCreditTrackerServer(grpcServer, testServer) + + // Start the server + go func() { + if err := grpcServer.Serve(listener); err != nil { + t.Logf("server stopped: %v", err) + } + }() + + // Wait for server to be ready by attempting to connect + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + for { + select { + case <-ctx.Done(): + t.Fatal("timeout waiting for server to start") + default: + conn, err := net.Dial("tcp", testServer.URL()) + if err == nil { + _ = conn.Close() + return testServer + } + time.Sleep(10 * time.Millisecond) + } + } +} + +// Close gracefully stops the test server +func (ts *mockCreditTrackerServer) Close() { + ts.grpcServer.GracefulStop() + + if ts.listener != nil { + ts.listener.Close() //nolint:errcheck + } +} + +// URL returns the full address of the server +func (ts *mockCreditTrackerServer) URL() string { + return ts.listener.Addr().String() +} + +// SetResponse sets a response for a given method and request parameters +func (ts *mockCreditTrackerServer) SetResponse(method string, requestKey string, response any) { + ts.mutex.Lock() + defer ts.mutex.Unlock() + + if ts.responses[method] == nil { + ts.responses[method] = make(map[string]any) + } + ts.responses[method][requestKey] = response +} + +// getRequestKey generates a unique key for a request based on its parameters +func getRequestKey(req any) string { + switch r := req.(type) { + case *ctgrpc.CreditCheckRequest: + return fmt.Sprintf("%s:%s", r.DeveloperLicense, r.AssetDid) + case *ctgrpc.CreditDeductRequest: + return fmt.Sprintf("%s:%s:%d", r.DeveloperLicense, r.AssetDid, r.Amount) + case *ctgrpc.RefundCreditsRequest: + return fmt.Sprintf("%s:%s:%d:%s", r.DeveloperLicense, r.AssetDid, r.Amount, r.Reason) + default: + return "" + } +} + +// CheckCredits implements the gRPC CheckCredits method +func (s *mockCreditTrackerServer) CheckCredits(ctx context.Context, req *ctgrpc.CreditCheckRequest) (*ctgrpc.CreditCheckResponse, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + + requestKey := getRequestKey(req) + if responses, ok := s.responses["CheckCredits"]; ok { + if response, ok := responses[requestKey]; ok { + if resp, ok := response.(*ctgrpc.CreditCheckResponse); ok { + return resp, nil + } + } + } + + // Default response if no custom response is set + return &ctgrpc.CreditCheckResponse{ + RemainingCredits: 100, + }, nil +} + +// DeductCredits implements the gRPC DeductCredits method +func (s *mockCreditTrackerServer) DeductCredits(ctx context.Context, req *ctgrpc.CreditDeductRequest) (*ctgrpc.CreditDeductResponse, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + + requestKey := getRequestKey(req) + if responses, ok := s.responses["DeductCredits"]; ok { + if response, ok := responses[requestKey]; ok { + if resp, ok := response.(*ctgrpc.CreditDeductResponse); ok { + return resp, nil + } + } + } + + // Default response if no custom response is set + return &ctgrpc.CreditDeductResponse{ + RemainingCredits: 99, + }, nil +} + +// RefundCredits implements the gRPC RefundCredits method +func (s *mockCreditTrackerServer) RefundCredits(ctx context.Context, req *ctgrpc.RefundCreditsRequest) (*ctgrpc.RefundCreditsResponse, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + + requestKey := getRequestKey(req) + if responses, ok := s.responses["RefundCredits"]; ok { + if response, ok := responses[requestKey]; ok { + if resp, ok := response.(*ctgrpc.RefundCreditsResponse); ok { + return resp, nil + } + } + } + + // Default response if no custom response is set + return &ctgrpc.RefundCreditsResponse{ + RemainingCredits: 101, + }, nil +} diff --git a/e2e/permission_test.go b/e2e/permission_test.go index d5b29781..11fe52b9 100644 --- a/e2e/permission_test.go +++ b/e2e/permission_test.go @@ -29,7 +29,7 @@ func TestPermission(t *testing.T) { } }`, permissions: []int{}, - expectedErr: `[{"message":"unauthorized: token id does not match","path":["signalsLatest"]}]`, + expectedErr: "unauthorized: token id does not match", }, { name: "Token permissions", @@ -53,7 +53,7 @@ func TestPermission(t *testing.T) { } }`, permissions: []int{}, - expectedErr: `[{"message":"unauthorized: missing required privilege(s) VEHICLE_NON_LOCATION_DATA","path":["signalsLatest","speed"]}]`, + expectedErr: "unauthorized: missing required privilege(s) VEHICLE_NON_LOCATION_DATA", }, { name: "Non Location permissions", @@ -120,7 +120,7 @@ func TestPermission(t *testing.T) { } }`, permissions: []int{1}, - expectedErr: `[{"message":"unauthorized: requires at least one of the following privileges [VEHICLE_APPROXIMATE_LOCATION VEHICLE_ALL_TIME_LOCATION]" ,"path":["signalsLatest","currentLocationApproximateLatitude"]}]`, + expectedErr: "unauthorized: requires at least one of the following privileges [VEHICLE_APPROXIMATE_LOCATION VEHICLE_ALL_TIME_LOCATION]", }, } @@ -132,7 +132,7 @@ func TestPermission(t *testing.T) { err := telemetryClient.Post(tt.query, &result, WithToken(token)) if tt.expectedErr != "" { require.Error(t, err) - require.JSONEq(t, tt.expectedErr, err.Error()) + require.Contains(t, err.Error(), tt.expectedErr) return } require.NoError(t, err) diff --git a/e2e/setup_test.go b/e2e/setup_test.go index bafa0915..00a9e95e 100644 --- a/e2e/setup_test.go +++ b/e2e/setup_test.go @@ -19,6 +19,7 @@ type TestServices struct { Auth *mockAuthServer FetchServer *mockFetchServer CH *container.Container + CT *mockCreditTrackerServer Settings config.Settings } @@ -61,18 +62,21 @@ func GetTestServices(t *testing.T) *TestServices { auth := setupAuthServer(t, settings.VehicleNFTAddress, settings.ManufacturerNFTAddress) fetch := NewTestFetchAPI(t) ch := setupClickhouseContainer(t) - // Create test settings + ct := setupCreditTrackerContainer(t) + // Create test settings settings.FetchAPIGRPCEndpoint = fetch.URL() settings.Clickhouse = ch.Config() settings.IdentityAPIURL = identity.URL() settings.TokenExchangeJWTKeySetURL = auth.URL() + "/keys" + settings.CreditTrackerEndpoint = ct.URL() testServices = &TestServices{ Identity: identity, Auth: auth, FetchServer: fetch, CH: ch, + CT: ct, Settings: settings, } cleanup = func() { @@ -81,6 +85,7 @@ func GetTestServices(t *testing.T) *TestServices { auth.Close() fetch.Close() ch.Terminate(context.Background()) + ct.Close() }) } }) diff --git a/go.mod b/go.mod index a13594f8..6ce2af45 100644 --- a/go.mod +++ b/go.mod @@ -8,6 +8,7 @@ require ( github.com/DIMO-Network/attestation-api v0.0.22 github.com/DIMO-Network/clickhouse-infra v0.0.3 github.com/DIMO-Network/cloudevent v0.1.0 + github.com/DIMO-Network/credit-tracker v0.0.0-20250603213155-96e2e9965b01 github.com/DIMO-Network/fetch-api v0.0.12 github.com/DIMO-Network/model-garage v0.6.0 github.com/DIMO-Network/shared v1.0.3 @@ -25,7 +26,8 @@ require ( github.com/vektah/gqlparser/v2 v2.5.27 github.com/volatiletech/sqlboiler/v4 v4.19.1 go.uber.org/mock v0.5.2 - google.golang.org/grpc v1.72.1 + google.golang.org/genproto/googleapis/rpc v0.0.0-20250324211829-b45e905df463 + google.golang.org/grpc v1.72.2 google.golang.org/protobuf v1.36.6 ) @@ -141,7 +143,6 @@ require ( golang.org/x/text v0.25.0 // indirect golang.org/x/tools v0.33.0 // indirect golang.org/x/xerrors v0.0.0-20240716161551-93cc26a95ae9 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20250324211829-b45e905df463 // indirect gopkg.in/go-jose/go-jose.v2 v2.6.3 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index 137d4baa..3bc58193 100644 --- a/go.sum +++ b/go.sum @@ -18,6 +18,8 @@ github.com/DIMO-Network/clickhouse-infra v0.0.3 h1:B6/4IY9IxLcyydET14IjHUT+A5SDE github.com/DIMO-Network/clickhouse-infra v0.0.3/go.mod h1:NtpQ1btkPzebDvpYYygeqiiBmJ/q5oJb/T/JWzUVRlk= github.com/DIMO-Network/cloudevent v0.1.0 h1:ze0ngJQBXjSSyBnEAUO+YvClvkqM68kNko/czY3GfLo= github.com/DIMO-Network/cloudevent v0.1.0/go.mod h1:RS9Byb0ycb5b7OFe9y+xpF0nkR4pYYS6Om/ccs3N5Z4= +github.com/DIMO-Network/credit-tracker v0.0.0-20250603213155-96e2e9965b01 h1:KW+yKmX+LreOfSKt0lqEitKN20PaUcMkThIxYLSJOsM= +github.com/DIMO-Network/credit-tracker v0.0.0-20250603213155-96e2e9965b01/go.mod h1:Ze9AjcpcEXRo2s9N4dS5JhAGfZlExWUM4Qzo1ztJi2c= github.com/DIMO-Network/fetch-api v0.0.12 h1:pLUekaYNWHKmguoGAni+BtTWl1j1yOdpowN5T6bWZuo= github.com/DIMO-Network/fetch-api v0.0.12/go.mod h1:9FtpOZR6kChy9x8cXN1z/0002ATh7R6yDo40pABAyNw= github.com/DIMO-Network/model-garage v0.6.0 h1:VtRjcEljCTffeDQbq5k8LP1Y01FBt5OpgwXRrH5QiPY= @@ -492,8 +494,8 @@ google.golang.org/genproto/googleapis/api v0.0.0-20250218202821-56aae31c358a h1: google.golang.org/genproto/googleapis/api v0.0.0-20250218202821-56aae31c358a/go.mod h1:3kWAYMk1I75K4vykHtKt2ycnOgpA6974V7bREqbsenU= google.golang.org/genproto/googleapis/rpc v0.0.0-20250324211829-b45e905df463 h1:e0AIkUUhxyBKh6ssZNrAMeqhA7RKUj42346d1y02i2g= google.golang.org/genproto/googleapis/rpc v0.0.0-20250324211829-b45e905df463/go.mod h1:qQ0YXyHHx3XkvlzUtpXDkS29lDSafHMZBAZDc03LQ3A= -google.golang.org/grpc v1.72.1 h1:HR03wO6eyZ7lknl75XlxABNVLLFc2PAb6mHlYh756mA= -google.golang.org/grpc v1.72.1/go.mod h1:wH5Aktxcg25y1I3w7H69nHfXdOG3UiadoBtjh3izSDM= +google.golang.org/grpc v1.72.2 h1:TdbGzwb82ty4OusHWepvFWGLgIbNo1/SUynEN0ssqv8= +google.golang.org/grpc v1.72.2/go.mod h1:wH5Aktxcg25y1I3w7H69nHfXdOG3UiadoBtjh3izSDM= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM= diff --git a/internal/app/app.go b/internal/app/app.go index ce99f3f8..7f10a407 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -14,6 +14,7 @@ import ( "github.com/99designs/gqlgen/graphql/handler/transport" "github.com/DIMO-Network/telemetry-api/internal/auth" "github.com/DIMO-Network/telemetry-api/internal/config" + "github.com/DIMO-Network/telemetry-api/internal/dtcmiddleware" "github.com/DIMO-Network/telemetry-api/internal/graph" "github.com/DIMO-Network/telemetry-api/internal/limits" "github.com/DIMO-Network/telemetry-api/internal/metrics" @@ -21,6 +22,7 @@ import ( "github.com/DIMO-Network/telemetry-api/internal/repositories/attestation" "github.com/DIMO-Network/telemetry-api/internal/repositories/vc" "github.com/DIMO-Network/telemetry-api/internal/service/ch" + "github.com/DIMO-Network/telemetry-api/internal/service/credittracker" "github.com/DIMO-Network/telemetry-api/internal/service/fetchapi" "github.com/DIMO-Network/telemetry-api/internal/service/identity" "github.com/DIMO-Network/telemetry-api/pkg/errorhandler" @@ -57,6 +59,11 @@ func New(settings config.Settings) (*App, error) { return nil, fmt.Errorf("failed to create attestation repository: %w", err) } + ctClient, err := credittracker.NewClient(&settings) + if err != nil { + return nil, fmt.Errorf("failed to create credit tracker client: %w", err) + } + resolver := &graph.Resolver{ Repository: baseRepo, IdentityService: idService, @@ -72,7 +79,8 @@ func New(settings config.Settings) (*App, error) { cfg.Directives.IsSignal = noOp cfg.Directives.HasAggregation = noOp - server := newDefaultServer(graph.NewExecutableSchema(cfg)) + server := newServer(graph.NewExecutableSchema(cfg)) + server.Use(dtcmiddleware.NewDCT(ctClient)) authMiddleware, err := auth.NewJWTMiddleware(settings.TokenExchangeIssuer, settings.TokenExchangeJWTKeySetURL) if err != nil { @@ -124,7 +132,7 @@ func newVinVCServiceFromSettings(settings config.Settings) (*vc.Repository, erro return vc.New(fetchapiSvc, settings.VINVCDataVersion, settings.POMVCDataVersion, uint64(settings.ChainID), settings.VehicleNFTAddress), nil } -func newDefaultServer(es graphql.ExecutableSchema) *handler.Server { +func newServer(es graphql.ExecutableSchema) *handler.Server { srv := handler.New(es) srv.AddTransport(transport.Websocket{ @@ -162,7 +170,7 @@ func LoggerMiddleware(next http.Handler) http.Handler { // authLoggerMiddleware adds the authenticated user to the logger func authLoggerMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - validateClaims, ok := auth.GetValidatedClaims(r) + validateClaims, ok := auth.GetValidatedClaims(r.Context()) if !ok { next.ServeHTTP(w, r) return diff --git a/internal/auth/jwt.go b/internal/auth/jwt.go index 84fb0570..295d15d2 100644 --- a/internal/auth/jwt.go +++ b/internal/auth/jwt.go @@ -64,7 +64,7 @@ func AddClaimHandler(next http.Handler, vehicleAddr, mfrAddr common.Address) htt } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - claims, ok := GetValidatedClaims(r) + claims, ok := GetValidatedClaims(r.Context()) if !ok || claims.CustomClaims == nil { // unauthorized calls will not have a claims. next.ServeHTTP(w, r) @@ -87,8 +87,8 @@ func AddClaimHandler(next http.Handler, vehicleAddr, mfrAddr common.Address) htt } // GetValidatedClaims returns the validated claims from the request context. -func GetValidatedClaims(r *http.Request) (*validator.ValidatedClaims, bool) { - claim := r.Context().Value(jwtmiddleware.ContextKey{}) +func GetValidatedClaims(ctx context.Context) (*validator.ValidatedClaims, bool) { + claim := ctx.Value(jwtmiddleware.ContextKey{}) if claim == nil { return nil, false } diff --git a/internal/config/settings.go b/internal/config/settings.go index b5bd41e0..0c3ad77f 100644 --- a/internal/config/settings.go +++ b/internal/config/settings.go @@ -23,4 +23,5 @@ type Settings struct { DeviceLastSeenBinHrs int64 `yaml:"DEVICE_LAST_SEEN_BIN_HOURS"` ChainID int `yaml:"DIMO_REGISTRY_CHAIN_ID"` FetchAPIGRPCEndpoint string `yaml:"FETCH_API_GRPC_ENDPOINT"` + CreditTrackerEndpoint string `yaml:"CREDIT_TRACKER_ENDPOINT"` } diff --git a/internal/dtcmiddleware/dctmiddleware.go b/internal/dtcmiddleware/dctmiddleware.go new file mode 100644 index 00000000..81f6b579 --- /dev/null +++ b/internal/dtcmiddleware/dctmiddleware.go @@ -0,0 +1,168 @@ +// Package dtcmiddleware provides a middleware for the Developer Credit Tracker. +package dtcmiddleware + +import ( + "context" + "fmt" + "math/big" + "net/http" + + "github.com/99designs/gqlgen/graphql" + "github.com/DIMO-Network/credit-tracker/pkg/grpc" + "github.com/DIMO-Network/telemetry-api/internal/auth" + "github.com/DIMO-Network/telemetry-api/internal/service/credittracker" + "github.com/DIMO-Network/telemetry-api/pkg/errorhandler" + "github.com/rs/zerolog" + "github.com/vektah/gqlparser/v2/gqlerror" + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/status" +) + +var defaultCreditAmount = int64(1) + +// DCT provides a GraphQL middleware for the Developer Credit Tracker. +type DCT struct { + Tracker *credittracker.Client +} + +var _ interface { + graphql.HandlerExtension + graphql.ResponseInterceptor +} = DCT{} + +// ExtensionName returns the name of this extension. +func (DCT) ExtensionName() string { + return "DCT" +} + +// Validate validates the GraphQL schema. +func (DCT) Validate(schema graphql.ExecutableSchema) error { + return nil +} + +// NewDCT creates a new DCT middleware with default values. +func NewDCT(tracker *credittracker.Client) *DCT { + return &DCT{ + Tracker: tracker, + } +} + +// InterceptResponse intercepts GraphQL responses to handle errors from the credit tracker. +func (d DCT) InterceptResponse( + ctx context.Context, + next graphql.ResponseHandler, +) *graphql.Response { + if d.Tracker == nil { + return graphql.ErrorResponse(ctx, "DCT is not enabled") + } + + // Determine who to charge + developerID, tokenID, gqlError := d.getSubjectAndTokenID(ctx) + if gqlError != nil { + return &graphql.Response{ + Errors: gqlerror.List{gqlError}, + } + } + + // Determine how many credits to charge + credits, gqlError := d.calculateCredits(ctx) + if gqlError != nil { + return &graphql.Response{ + Errors: gqlerror.List{gqlError}, + } + } + + // Deduct the credits + err := d.Tracker.DeductCredits(ctx, developerID, tokenID, credits) + if err != nil { + zerolog.Ctx(ctx).Error().Err(err).Msg("Failed to deduct credits") + gqlError := processDCTErrorToGraphqlError(ctx, err) + return &graphql.Response{ + Errors: gqlerror.List{gqlError}, + } + } + + // Complete the request and get the response + response := next(ctx) + + // If it's our fault the request failed, refund the credits + if errorhandler.HasInternalError(&response.Errors) { + err := d.Tracker.RefundCredits(ctx, developerID, tokenID, credits) + if err != nil { + zerolog.Ctx(ctx).Error().Err(err).Msg("Failed to refund credits") + } + return response + } + + return response +} + +// processDCTError extracts and processes error details from a gRPC error +func processDCTErrorToGraphqlError(ctx context.Context, err error) *gqlerror.Error { + st, ok := status.FromError(err) + if !ok { + return graphql.DefaultErrorPresenter(ctx, err) + } + + for _, detail := range st.Details() { + if errorInfo, ok := detail.(*errdetails.ErrorInfo); ok { + if err := handleErrorDetails(ctx, errorInfo, err); err != nil { + return err + } + } + } + + return errorhandler.NewInternalErrorWithMsg(ctx, err, "Failed to process credit operation") +} + +// handleErrorDetails processes the error details from a gRPC status error +func handleErrorDetails(ctx context.Context, errorInfo *errdetails.ErrorInfo, originalError error) *gqlerror.Error { + switch errorInfo.Reason { + case grpc.ErrorReason_ERROR_REASON_INVALID_ASSET_DID.String(): + err := fmt.Errorf("invalid asset DID: %s", errorInfo.Metadata[grpc.MetadataKey_METADATA_KEY_ASSET_DID.String()]) + return errorhandler.NewInternalErrorWithMsg(ctx, err, "Failed to process credit operation") + case grpc.ErrorReason_ERROR_REASON_INSUFFICIENT_CREDITS.String(): + if txHash, ok := errorInfo.Metadata[grpc.MetadataKey_METADATA_KEY_TRANSACTION_HASH.String()]; ok { + return &gqlerror.Error{ + Message: fmt.Sprintf("insufficient credits, burn transaction initiated: %s", txHash), + Err: originalError, + Extensions: map[string]any{ + "reason": errorInfo.Reason, + "code": http.StatusPaymentRequired, + }, + } + } + return &gqlerror.Error{ + Message: fmt.Sprintf("insufficient credits for asset: %s", errorInfo.Metadata[grpc.MetadataKey_METADATA_KEY_ASSET_DID.String()]), + Extensions: map[string]any{ + "reason": errorInfo.Reason, + "code": http.StatusPaymentRequired, + }, + } + default: + return nil + } +} + +func (d DCT) getSubjectAndTokenID(ctx context.Context) (string, *big.Int, *gqlerror.Error) { + validateClaims, ok := auth.GetValidatedClaims(ctx) + if !ok || validateClaims.CustomClaims == nil { + return "", nil, errorhandler.NewUnauthorizedErrorWithMsg(ctx, fmt.Errorf("failed to get validated claims"), "Unauthorized") + } + telemClaim, ok := validateClaims.CustomClaims.(*auth.TelemetryClaim) + if !ok { + return "", nil, errorhandler.NewUnauthorizedErrorWithMsg(ctx, fmt.Errorf("failed to get cast exchange custom claims"), "Unauthorized") + } + tokenIDBig, ok := new(big.Int).SetString(telemClaim.TokenID, 10) + if !ok { + return "", nil, errorhandler.NewInternalErrorWithMsg(ctx, fmt.Errorf("failed to parse token ID"), "Failed to parse token ID") + } + + return validateClaims.RegisteredClaims.Subject, tokenIDBig, nil +} + +func (d DCT) calculateCredits(ctx context.Context) (int64, *gqlerror.Error) { + // TODO: We can add logic here to determine what the base cost for a given operations should be + return defaultCreditAmount, nil + +} diff --git a/internal/repositories/repositories.go b/internal/repositories/repositories.go index 5c9b8394..0dcc0b22 100644 --- a/internal/repositories/repositories.go +++ b/internal/repositories/repositories.go @@ -171,7 +171,7 @@ func (r *Repository) GetAvailableSignals(ctx context.Context, tokenID uint32, fi func handleDBError(ctx context.Context, err error) error { exceptionErr := &proto.Exception{} if errors.Is(err, context.DeadlineExceeded) || (errors.As(err, &exceptionErr) && exceptionErr.Code == ch.TimeoutErrCode) { - return errorhandler.NewInternalErrorWithMsg(ctx, err, "request exceeded or is estimated to exceed the maximum execution time") + return errorhandler.NewBadRequestErrorWithMsg(ctx, err, "request exceeded or is estimated to exceed the maximum execution time") } return errorhandler.NewInternalErrorWithMsg(ctx, err, "failed to query db") } diff --git a/internal/service/credittracker/credittracker.go b/internal/service/credittracker/credittracker.go new file mode 100644 index 00000000..47f3aaed --- /dev/null +++ b/internal/service/credittracker/credittracker.go @@ -0,0 +1,118 @@ +package credittracker + +import ( + "context" + "fmt" + "math/big" + "time" + + "github.com/DIMO-Network/cloudevent" + ctgrpc "github.com/DIMO-Network/credit-tracker/pkg/grpc" + "github.com/DIMO-Network/telemetry-api/internal/config" + "github.com/ethereum/go-ethereum/common" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" +) + +// Client implements the Client interface. +type Client struct { + conn *grpc.ClientConn + Endpoint string + RequestTimeout time.Duration + MaxRetries int + RetryTimeout time.Duration + ctClient ctgrpc.CreditTrackerClient + chainID uint64 + vehicleContractAddress common.Address +} + +// NewClient creates a new credit tracker client. +func NewClient(settings *config.Settings) (*Client, error) { + conn, err := grpc.NewClient(settings.CreditTrackerEndpoint, grpc.WithTransportCredentials(insecure.NewCredentials())) + if err != nil { + return nil, fmt.Errorf("failed to create credit tracker client: %w", err) + } + ctClient := ctgrpc.NewCreditTrackerClient(conn) + return &Client{ + conn: conn, + Endpoint: settings.CreditTrackerEndpoint, + RequestTimeout: 3 * time.Second, + MaxRetries: 3, + RetryTimeout: 100 * time.Millisecond, + ctClient: ctClient, + }, nil +} + +// DeductCredits deducts credits from the given developer license and token. +func (c *Client) DeductCredits(ctx context.Context, developerLicense string, tokenID *big.Int, amount int64) error { + trackerCtx, cancel := context.WithTimeout(ctx, c.RequestTimeout) + defer cancel() + + deductCredits := func() error { + _, err := c.ctClient.DeductCredits(trackerCtx, &ctgrpc.CreditDeductRequest{ + DeveloperLicense: developerLicense, + AssetDid: cloudevent.ERC721DID{ + ChainID: c.chainID, + ContractAddress: c.vehicleContractAddress, + TokenID: tokenID, + }.String(), + Amount: amount, + }) + if err != nil { + return fmt.Errorf("failed to send deduct request to credit tracker: %w", err) + } + return nil + } + err := c.runWithRetry(ctx, deductCredits) + if err != nil { + return err + } + return nil +} + +// RefundCredits refunds credits from the given developer license and token. +func (c *Client) RefundCredits(ctx context.Context, developerLicense string, tokenID *big.Int, amount int64) error { + trackerCtx, cancel := context.WithTimeout(ctx, c.RequestTimeout) + defer cancel() + + refundCredits := func() error { + _, err := c.ctClient.RefundCredits(trackerCtx, &ctgrpc.RefundCreditsRequest{ + DeveloperLicense: developerLicense, + AssetDid: cloudevent.ERC721DID{ + ChainID: c.chainID, + ContractAddress: c.vehicleContractAddress, + TokenID: tokenID, + }.String(), + Amount: amount, + }) + if err != nil { + return fmt.Errorf("failed to send refund request to credit tracker: %w", err) + } + return nil + } + err := c.runWithRetry(ctx, refundCredits) + if err != nil { + return err + } + return nil +} + +// Close closes the gRPC connection +func (c *Client) Close() error { + if c.conn != nil { + return c.conn.Close() + } + return nil +} + +func (c *Client) runWithRetry(ctx context.Context, f func() error) error { + var err error + for i := 0; i < c.MaxRetries; i++ { + if err = f(); err != nil { + time.Sleep(c.RetryTimeout) + continue + } + return nil + } + return err +} diff --git a/pkg/errorhandler/errorhandler.go b/pkg/errorhandler/errorhandler.go index 3f4c51f3..53ae2082 100644 --- a/pkg/errorhandler/errorhandler.go +++ b/pkg/errorhandler/errorhandler.go @@ -77,3 +77,38 @@ func NewUnauthorizedErrorWithMsg(ctx context.Context, err error, message string) func NewUnauthorizedError(ctx context.Context, err error) *gqlerror.Error { return NewUnauthorizedErrorWithMsg(ctx, err, err.Error()) } + +// HasInternalError checks if the gqlerror.List contains an internal server error. +func HasInternalError(gqlErrs *gqlerror.List) bool { + for _, err := range *gqlErrs { + if IsInternalError(err) { + return true + } + } + return false +} + +// IsInternalError checks if the gqlerror.Error is an internal server error. +func IsInternalError(gqlErr *gqlerror.Error) bool { + return gqlErr.Extensions["code"] == http.StatusInternalServerError +} + +// HasBadRequestError checks if the gqlerror.List contains a bad request error. +func HasBadRequestError(gqlErrs *gqlerror.List) bool { + for _, err := range *gqlErrs { + if IsBadRequestError(err) { + return true + } + } + return false +} + +// IsBadRequestError checks if the gqlerror.Error is a bad request error. +func IsBadRequestError(gqlErr *gqlerror.Error) bool { + return gqlErr.Extensions["code"] == http.StatusBadRequest +} + +// IsUnauthorizedError checks if the gqlerror.Error is an unauthorized error. +func IsUnauthorizedError(gqlErr *gqlerror.Error) bool { + return gqlErr.Extensions["code"] == http.StatusUnauthorized +}