From 5497ed202d3a85c9245c14c2f4556e6b50618a98 Mon Sep 17 00:00:00 2001 From: Johnny Luo Date: Fri, 6 Jun 2025 17:42:07 +1000 Subject: [PATCH 1/3] feat: update plugin policy handling to use UUID and add payroll plugin integration --- api/plugin.go | 50 +++++++++++++++++++++++++---------- cmd/payroll/worker/config.go | 1 + cmd/payroll/worker/main.go | 9 ++++++- go.mod | 4 +-- go.sum | 10 +++---- plugin/dca/dca.go | 20 ++++++++------ plugin/payroll/payroll.go | 3 +++ plugin/payroll/policy.go | 4 +-- plugin/payroll/transaction.go | 31 ++++++++++++++++++---- service/policy.go | 8 +++--- storage/db.go | 4 +-- storage/postgres/policy.go | 5 ++-- 12 files changed, 103 insertions(+), 46 deletions(-) diff --git a/api/plugin.go b/api/plugin.go index 5a4199c..d03bae3 100644 --- a/api/plugin.go +++ b/api/plugin.go @@ -59,8 +59,18 @@ func (s *Server) SignPluginMessages(c echo.Context) error { if policy.PluginID.String() != req.PluginID { return fmt.Errorf("policy plugin ID mismatch") } - - if err := s.plugin.ValidateProposedTransactions(policy, []vtypes.PluginKeysignRequest{req}); err != nil { + policyUpdate := vtypes.PluginPolicyCreateUpdate{ + ID: policy.ID, + PublicKey: policy.PublicKey, + PluginID: policy.PluginID, + PluginVersion: policy.PluginVersion, + PolicyVersion: policy.PolicyVersion, + Signature: policy.Signature, + Recipe: policy.Recipe, + BillingRecipe: "", + Active: policy.Active, + } + if err := s.plugin.ValidateProposedTransactions(policyUpdate, []vtypes.PluginKeysignRequest{req}); err != nil { return fmt.Errorf("failed to validate transaction proposal: %w", err) } @@ -146,8 +156,14 @@ func (s *Server) GetPluginPolicyById(c echo.Context) error { if policyID == "" { return c.JSON(http.StatusBadRequest, NewErrorResponse("invalid policy ID")) } - - policy, err := s.policyService.GetPluginPolicy(c.Request().Context(), policyID) + uPolicyID, err := uuid.Parse(policyID) + if err != nil { + s.logger.WithError(err). + WithField("policy_id", policyID). + Error("failed to parse policy ID") + return c.JSON(http.StatusBadRequest, NewErrorResponse("invalid policy ID")) + } + policy, err := s.policyService.GetPluginPolicy(c.Request().Context(), uPolicyID) if err != nil { s.logger.WithError(err). WithField("policy_id", policyID). @@ -183,7 +199,7 @@ func (s *Server) GetAllPluginPolicies(c echo.Context) error { } func (s *Server) CreatePluginPolicy(c echo.Context) error { - var policy vtypes.PluginPolicy + var policy vtypes.PluginPolicyCreateUpdate if err := c.Bind(&policy); err != nil { return fmt.Errorf("fail to parse request, err: %w", err) } @@ -199,12 +215,12 @@ func (s *Server) CreatePluginPolicy(c echo.Context) error { policy.ID = uuid.New() } - if !s.verifyPolicySignature(policy, false) { + if !s.verifyPolicySignature(policy.ToPluginPolicy(), false) { s.logger.Error("invalid policy signature") return c.JSON(http.StatusForbidden, NewErrorResponse("Invalid policy signature")) } - newPolicy, err := s.policyService.CreatePolicy(c.Request().Context(), policy) + newPolicy, err := s.policyService.CreatePolicy(c.Request().Context(), policy.ToPluginPolicy()) if err != nil { s.logger.WithError(err).Error("Failed to create plugin policy") return c.JSON(http.StatusInternalServerError, NewErrorResponse("failed to create policy")) @@ -214,7 +230,7 @@ func (s *Server) CreatePluginPolicy(c echo.Context) error { } func (s *Server) UpdatePluginPolicyById(c echo.Context) error { - var policy vtypes.PluginPolicy + var policy vtypes.PluginPolicyCreateUpdate if err := c.Bind(&policy); err != nil { return fmt.Errorf("fail to parse request, err: %w", err) } @@ -227,12 +243,12 @@ func (s *Server) UpdatePluginPolicyById(c echo.Context) error { return c.JSON(http.StatusBadRequest, NewErrorResponse("failed to validate policy")) } - if !s.verifyPolicySignature(policy, true) { + if !s.verifyPolicySignature(policy.ToPluginPolicy(), true) { s.logger.Error("invalid policy signature") return c.JSON(http.StatusForbidden, NewErrorResponse("Invalid policy signature")) } - updatedPolicy, err := s.policyService.UpdatePolicy(c.Request().Context(), policy) + updatedPolicy, err := s.policyService.UpdatePolicy(c.Request().Context(), policy.ToPluginPolicy()) if err != nil { s.logger.WithError(err).Error("Failed to update plugin policy") return c.JSON(http.StatusInternalServerError, NewErrorResponse("failed to update policy")) @@ -254,8 +270,14 @@ func (s *Server) DeletePluginPolicyById(c echo.Context) error { if policyID == "" { return c.JSON(http.StatusBadRequest, NewErrorResponse("invalid policy ID")) } - - policy, err := s.policyService.GetPluginPolicy(c.Request().Context(), policyID) + uPolicyID, err := uuid.Parse(policyID) + if err != nil { + s.logger.WithError(err). + WithField("policy_id", policyID). + Error("Failed to parse policy ID") + return c.JSON(http.StatusBadRequest, NewErrorResponse("invalid policy ID")) + } + policy, err := s.policyService.GetPluginPolicy(c.Request().Context(), uPolicyID) if err != nil { s.logger.WithError(err). WithField("policy_id", policyID). @@ -270,7 +292,7 @@ func (s *Server) DeletePluginPolicyById(c echo.Context) error { return c.JSON(http.StatusForbidden, NewErrorResponse("Invalid policy signature")) } - if err := s.policyService.DeletePolicy(c.Request().Context(), policyID, reqBody.Signature); err != nil { + if err := s.policyService.DeletePolicy(c.Request().Context(), uPolicyID, reqBody.Signature); err != nil { s.logger.WithError(err). WithField("policy_id", policyID). Error("Failed to delete plugin policy") @@ -304,7 +326,7 @@ func (s *Server) GetPolicySchema(c echo.Context) error { func (s *Server) GetRecipeSpecification(c echo.Context) error { recipeSpec := s.plugin.GetRecipeSpecification() - return c.JSON(http.StatusOK, recipeSpec) + return c.JSON(http.StatusOK, &recipeSpec) } func (s *Server) GetPluginPolicyTransactionHistory(c echo.Context) error { diff --git a/cmd/payroll/worker/config.go b/cmd/payroll/worker/config.go index b2e2d3d..ef4ab01 100644 --- a/cmd/payroll/worker/config.go +++ b/cmd/payroll/worker/config.go @@ -21,6 +21,7 @@ type PayrollWorkerConfig struct { Database struct { DSN string `mapstructure:"dsn" json:"dsn,omitempty"` } `mapstructure:"database" json:"database,omitempty"` + BaseConfigPath string `mapstructure:"base_file_path" json:"base_file_path,omitempty"` } func GetConfigure() (*PayrollWorkerConfig, error) { diff --git a/cmd/payroll/worker/main.go b/cmd/payroll/worker/main.go index 61eed6a..a26533a 100644 --- a/cmd/payroll/worker/main.go +++ b/cmd/payroll/worker/main.go @@ -10,6 +10,7 @@ import ( "github.com/vultisig/plugin/internal/scheduler" "github.com/vultisig/plugin/internal/tasks" + "github.com/vultisig/plugin/plugin/payroll" "github.com/vultisig/plugin/storage/postgres" ) @@ -56,14 +57,20 @@ func main() { if err != nil { panic(fmt.Errorf("failed to create postgres backend: %w", err)) } + p, err := payroll.NewPayrollPlugin(postgressDB, cfg.BaseConfigPath) + if err != nil { + panic(fmt.Errorf("failed to create payroll plugin: %w", err)) + } schedulerSvc, err := scheduler.NewSchedulerService(postgressDB, client, redisOptions) if err != nil { panic(fmt.Errorf("failed to create scheduler service: %w", err)) } + schedulerSvc.Start() defer schedulerSvc.Stop() + mux := asynq.NewServeMux() - // mux.HandleFunc(tasks.TypePluginTransaction, vaultService.HandlePluginTransaction) + mux.HandleFunc(tasks.TypePluginTransaction, p.HandleSchedulerTrigger) mux.HandleFunc(tasks.TypeKeySignDKLS, vaultService.HandleKeySignDKLS) mux.HandleFunc(tasks.TypeReshareDKLS, vaultService.HandleReshareDKLS) if err := srv.Run(mux); err != nil { diff --git a/go.mod b/go.mod index f6a6e39..d6ccea8 100644 --- a/go.mod +++ b/go.mod @@ -19,8 +19,8 @@ require ( github.com/spf13/viper v1.20.1 github.com/vultisig/commondata v0.0.0-20250430024109-a2492623ef05 github.com/vultisig/mobile-tss-lib v0.0.0-20250316003201-2e7e570a4a74 - github.com/vultisig/recipes v0.0.0-20250603213257-b4a6b0afe2a0 - github.com/vultisig/verifier v0.0.0-20250531103513-c1f5cb38b103 + github.com/vultisig/recipes v0.0.0-20250604212709-58772375f814 + github.com/vultisig/verifier v0.0.0-20250606071917-2f9ec25b5689 github.com/vultisig/vultiserver v0.0.0-20250515110921-82d56d3d9cc9 ) diff --git a/go.sum b/go.sum index c571da7..7089d4a 100644 --- a/go.sum +++ b/go.sum @@ -738,12 +738,10 @@ github.com/vultisig/go-wrappers v0.0.0-20250403041248-86911e8aa33f h1:124Xlloih1 github.com/vultisig/go-wrappers v0.0.0-20250403041248-86911e8aa33f/go.mod h1:UfGCxUQW08kiwxyNBiHwXe+ePPuBmHVVS+BS51aU/Jg= github.com/vultisig/mobile-tss-lib v0.0.0-20250316003201-2e7e570a4a74 h1:goqwk4nQ/NEVIb3OPP9SUx7/u9ZfsUIcd5fIN/e4DVU= github.com/vultisig/mobile-tss-lib v0.0.0-20250316003201-2e7e570a4a74/go.mod h1:nOykk4nOy1L3yXtLSlYvVsgizBnCQ3tR2N5uwGPdvaM= -github.com/vultisig/recipes v0.0.0-20250531132511-82a0eb621885 h1:hXju0II1G1CydXchhjTla1JxK8WceDUGJhQhJHJSUIE= -github.com/vultisig/recipes v0.0.0-20250531132511-82a0eb621885/go.mod h1:BXXJ25U75xaexJLoiiaLEZ6TQrcy8+8UGJGd4c+hNSE= -github.com/vultisig/recipes v0.0.0-20250603213257-b4a6b0afe2a0 h1:JauwztTOr7EqVhhXIQw+fWX6+VKDc5T6zLVBFIvGuBk= -github.com/vultisig/recipes v0.0.0-20250603213257-b4a6b0afe2a0/go.mod h1:BXXJ25U75xaexJLoiiaLEZ6TQrcy8+8UGJGd4c+hNSE= -github.com/vultisig/verifier v0.0.0-20250531103513-c1f5cb38b103 h1:1Qgt3uwzG11UqKoizd5Ft65Ij/8c2Kxpg3VB6N3/Fb0= -github.com/vultisig/verifier v0.0.0-20250531103513-c1f5cb38b103/go.mod h1:8b9k8CJtF1KwjMoZuWFYmz0Q+zdikrx3q0ExbZqsLgk= +github.com/vultisig/recipes v0.0.0-20250604212709-58772375f814 h1:uY4x91tXkpDwuUOLBxUrG0ASvJzlvNKnbGqOMrBs0U8= +github.com/vultisig/recipes v0.0.0-20250604212709-58772375f814/go.mod h1:BXXJ25U75xaexJLoiiaLEZ6TQrcy8+8UGJGd4c+hNSE= +github.com/vultisig/verifier v0.0.0-20250606071917-2f9ec25b5689 h1:hD5rRrAImxdqN7yyI5BqqtZWye9Eoera9bbvimf0dhU= +github.com/vultisig/verifier v0.0.0-20250606071917-2f9ec25b5689/go.mod h1:Llu11hCj/HxVdzWIvnpwAJ9s4QnooaCCxZr42ExZS4g= github.com/vultisig/vultiserver v0.0.0-20250515110921-82d56d3d9cc9 h1:pZhGN8q8+gPB1JJjVDC1hDg8qn6Tbj0XBJymgTQ8qQg= github.com/vultisig/vultiserver v0.0.0-20250515110921-82d56d3d9cc9/go.mod h1:HwP2IgW6Mcu/gX8paFuKvfibrGE9UmPgkOFTub6dskM= github.com/xordataexchange/crypt v0.0.3-0.20170626215501-b2862e3d0a77/go.mod h1:aYKd//L2LvnjZzWKhF00oedf4jCCReLcmhLdhm1A27Q= diff --git a/plugin/dca/dca.go b/plugin/dca/dca.go index 012ac2e..54cd8fd 100644 --- a/plugin/dca/dca.go +++ b/plugin/dca/dca.go @@ -21,14 +21,16 @@ import ( "github.com/vultisig/mobile-tss-lib/tss" "github.com/vultisig/verifier/address" vcommon "github.com/vultisig/verifier/common" + "github.com/vultisig/verifier/plugin" vtypes "github.com/vultisig/verifier/types" + rtypes "github.com/vultisig/recipes/types" + "github.com/vultisig/plugin/common" "github.com/vultisig/plugin/internal/sigutil" "github.com/vultisig/plugin/internal/types" "github.com/vultisig/plugin/pkg/uniswap" "github.com/vultisig/plugin/storage" - rtypes "github.com/vultisig/recipes/types" ) const ( @@ -46,6 +48,8 @@ var ( ErrCompletedPolicy = errors.New("policy completed all swaps") ) +var _ plugin.Plugin = (*DCAPlugin)(nil) + type DCAPlugin struct { uniswapClient *uniswap.Client rpcClient *ethclient.Client @@ -90,7 +94,7 @@ func (p *DCAPlugin) SigningComplete( ctx context.Context, signature tss.KeysignResponse, signRequest vtypes.PluginKeysignRequest, - policy vtypes.PluginPolicy, + policy vtypes.PluginPolicyCreateUpdate, ) error { var dcaPolicy DCAPolicy // TODO: convert recipe to DCAPolicy @@ -131,7 +135,7 @@ func (p *DCAPlugin) SigningComplete( return nil } -func (p *DCAPlugin) ValidatePluginPolicy(policyDoc vtypes.PluginPolicy) error { +func (p *DCAPlugin) ValidatePluginPolicy(policyDoc vtypes.PluginPolicyCreateUpdate) error { if policyDoc.PluginID != vtypes.PluginVultisigDCA_0000 { return fmt.Errorf("policy does not match plugin type, expected: %s, got: %s", pluginType, policyDoc.PluginID) } @@ -266,7 +270,7 @@ func validateInterval(intervalStr string, frequency string) error { return nil } -func (p *DCAPlugin) ProposeTransactions(policy vtypes.PluginPolicy) ([]vtypes.PluginKeysignRequest, error) { +func (p *DCAPlugin) ProposeTransactions(policy vtypes.PluginPolicyCreateUpdate) ([]vtypes.PluginKeysignRequest, error) { p.logger.Info("DCA: PROPOSE TRANSACTIONS") var txs []vtypes.PluginKeysignRequest @@ -297,7 +301,7 @@ func (p *DCAPlugin) ProposeTransactions(policy vtypes.PluginPolicy) ([]vtypes.Pl } if completedSwaps >= totalOrders.Int64() { - if err := p.completePolicy(context.Background(), policy); err != nil { + if err := p.completePolicy(context.Background(), policy.ToPluginPolicy()); err != nil { return txs, fmt.Errorf("fail to complete policy: %w", err) } return txs, ErrCompletedPolicy @@ -343,9 +347,9 @@ func (p *DCAPlugin) ProposeTransactions(policy vtypes.PluginPolicy) ([]vtypes.Pl common.PluginPartyID, common.VerifierPartyID}, PluginID: policy.PluginID.String(), + PolicyID: policy.ID, }, Transaction: hex.EncodeToString(data.RlpTxBytes), - PolicyID: policy.ID.String(), TransactionType: data.Type, } txs = append(txs, signRequest) @@ -354,7 +358,7 @@ func (p *DCAPlugin) ProposeTransactions(policy vtypes.PluginPolicy) ([]vtypes.Pl return txs, nil } -func (p *DCAPlugin) ValidateProposedTransactions(policy vtypes.PluginPolicy, txs []vtypes.PluginKeysignRequest) error { +func (p *DCAPlugin) ValidateProposedTransactions(policy vtypes.PluginPolicyCreateUpdate, txs []vtypes.PluginKeysignRequest) error { p.logger.Info("DCA: VALIDATE TRANSACTION PROPOSAL") if len(txs) == 0 { @@ -403,7 +407,7 @@ func (p *DCAPlugin) ValidateProposedTransactions(policy vtypes.PluginPolicy, txs } // TODO: Change this to make the policy to status COMPLETED if: completed swaps == total orders. if completedSwaps >= totalOrders.Int64() { - if err := p.completePolicy(context.Background(), policy); err != nil { + if err := p.completePolicy(context.Background(), policy.ToPluginPolicy()); err != nil { return fmt.Errorf("fail to complete policy: %w", err) } p.logger.Info("DCA: COMPLETED SWAPS: ", totalOrders.Int64()) diff --git a/plugin/payroll/payroll.go b/plugin/payroll/payroll.go index f36d6ae..af768ca 100644 --- a/plugin/payroll/payroll.go +++ b/plugin/payroll/payroll.go @@ -5,10 +5,13 @@ import ( "github.com/ethereum/go-ethereum/ethclient" "github.com/sirupsen/logrus" + "github.com/vultisig/verifier/plugin" "github.com/vultisig/plugin/storage" ) +var _ plugin.Plugin = (*PayrollPlugin)(nil) + type PayrollPlugin struct { db storage.DatabaseStorage nonceManager *NonceManager diff --git a/plugin/payroll/policy.go b/plugin/payroll/policy.go index fca7b3b..70bc2c9 100644 --- a/plugin/payroll/policy.go +++ b/plugin/payroll/policy.go @@ -37,7 +37,7 @@ type Schedule struct { EndTime string `json:"end_time,omitempty"` } -func (p *PayrollPlugin) ValidateProposedTransactions(policy vtypes.PluginPolicy, txs []vtypes.PluginKeysignRequest) error { +func (p *PayrollPlugin) ValidateProposedTransactions(policy vtypes.PluginPolicyCreateUpdate, txs []vtypes.PluginKeysignRequest) error { err := p.ValidatePluginPolicy(policy) if err != nil { return fmt.Errorf("failed to validate plugin policy: %v", err) @@ -196,7 +196,7 @@ func (p *PayrollPlugin) checkRule(rule *rtypes.Rule) error { } return nil } -func (p *PayrollPlugin) ValidatePluginPolicy(policyDoc vtypes.PluginPolicy) error { +func (p *PayrollPlugin) ValidatePluginPolicy(policyDoc vtypes.PluginPolicyCreateUpdate) error { if policyDoc.PluginID != vtypes.PluginVultisigPayroll_0000 { return fmt.Errorf("policy does not match plugin type, expected: %s, got: %s", vtypes.PluginVultisigPayroll_0000, policyDoc.PluginID) } diff --git a/plugin/payroll/transaction.go b/plugin/payroll/transaction.go index 508ace7..615063d 100644 --- a/plugin/payroll/transaction.go +++ b/plugin/payroll/transaction.go @@ -3,12 +3,14 @@ package payroll import ( "context" "encoding/hex" + "encoding/json" "fmt" "math/big" "strconv" "strings" "github.com/google/uuid" + "github.com/vultisig/vultiserver/contexthelper" "github.com/ethereum/go-ethereum" "github.com/ethereum/go-ethereum/accounts/abi" @@ -22,6 +24,8 @@ import ( "github.com/vultisig/verifier/address" vcommon "github.com/vultisig/verifier/common" vtypes "github.com/vultisig/verifier/types" + + "github.com/vultisig/plugin/internal/types" ) // TODO: remove once the plugin installation is implemented @@ -29,7 +33,26 @@ const ( hexEncryptionKey = "hexencryptionkey" ) -func (p *PayrollPlugin) ProposeTransactions(policy vtypes.PluginPolicy) ([]vtypes.PluginKeysignRequest, error) { +func (p *PayrollPlugin) HandleSchedulerTrigger(ctx context.Context, t *asynq.Task) error { + if err := contexthelper.CheckCancellation(ctx); err != nil { + p.logger.WithError(err).Warn("Context cancelled, skipping scheduler trigger") + return err + } + var trigger types.TimeTrigger + if err := json.Unmarshal(t.Payload(), &trigger); err != nil { + p.logger.WithError(err).Error("Failed to unmarshal trigger payload") + return fmt.Errorf("failed to unmarshal trigger payload: %s, %w", err, asynq.SkipRetry) + } + pluginPolicy, err := p.db.GetPluginPolicy(ctx, trigger.PolicyID) + if err != nil { + p.logger.WithError(err).Error("Failed to get plugin policy from database") + return fmt.Errorf("failed to get plugin policy: %s, %w", err, asynq.SkipRetry) + } + // propose transaction and get it signed + _ = pluginPolicy + return nil +} +func (p *PayrollPlugin) ProposeTransactions(policy vtypes.PluginPolicyCreateUpdate) ([]vtypes.PluginKeysignRequest, error) { var txs []vtypes.PluginKeysignRequest err := p.ValidatePluginPolicy(policy) if err != nil { @@ -73,8 +96,6 @@ func (p *PayrollPlugin) ProposeTransactions(policy vtypes.PluginPolicy) ([]vtype PluginID: policy.PluginID.String(), }, Transaction: hex.EncodeToString(rawTx), - - PolicyID: policy.ID.String(), } txs = append(txs, signRequest) } @@ -211,8 +232,8 @@ func (p *PayrollPlugin) generatePayrollTransaction(amountString, recipientString return txHash, rawTx, nil } -func (p *PayrollPlugin) SigningComplete(ctx context.Context, signature tss.KeysignResponse, signRequest vtypes.PluginKeysignRequest, policy vtypes.PluginPolicy) error { - R, S, V, originalTx, chainID, _, err := p.convertData(signature, signRequest, policy) +func (p *PayrollPlugin) SigningComplete(ctx context.Context, signature tss.KeysignResponse, signRequest vtypes.PluginKeysignRequest, policy vtypes.PluginPolicyCreateUpdate) error { + R, S, V, originalTx, chainID, _, err := p.convertData(signature, signRequest, policy.ToPluginPolicy()) if err != nil { return fmt.Errorf("failed to convert R and S: %v", err) } diff --git a/service/policy.go b/service/policy.go index 67eed23..7094b3d 100644 --- a/service/policy.go +++ b/service/policy.go @@ -17,9 +17,9 @@ import ( type Policy interface { CreatePolicy(ctx context.Context, policy vtypes.PluginPolicy) (*vtypes.PluginPolicy, error) UpdatePolicy(ctx context.Context, policy vtypes.PluginPolicy) (*vtypes.PluginPolicy, error) - DeletePolicy(ctx context.Context, policyID, signature string) error + DeletePolicy(ctx context.Context, policyID uuid.UUID, signature string) error GetPluginPolicies(ctx context.Context, pluginID vtypes.PluginID, publicKey string) ([]vtypes.PluginPolicy, error) - GetPluginPolicy(ctx context.Context, policyID string) (vtypes.PluginPolicy, error) + GetPluginPolicy(ctx context.Context, policyID uuid.UUID) (vtypes.PluginPolicy, error) GetPluginPolicyTransactionHistory(ctx context.Context, policyID string) ([]types.TransactionHistory, error) } @@ -107,7 +107,7 @@ func (s *PolicyService) UpdatePolicy(ctx context.Context, policy vtypes.PluginPo return updatedPolicy, nil } -func (s *PolicyService) DeletePolicy(ctx context.Context, policyID, signature string) error { +func (s *PolicyService) DeletePolicy(ctx context.Context, policyID uuid.UUID, signature string) error { tx, err := s.db.Pool().Begin(ctx) if err != nil { @@ -135,7 +135,7 @@ func (s *PolicyService) GetPluginPolicies(ctx context.Context, pluginID vtypes.P return policies, nil } -func (s *PolicyService) GetPluginPolicy(ctx context.Context, policyID string) (vtypes.PluginPolicy, error) { +func (s *PolicyService) GetPluginPolicy(ctx context.Context, policyID uuid.UUID) (vtypes.PluginPolicy, error) { policy, err := s.db.GetPluginPolicy(ctx, policyID) if err != nil { return vtypes.PluginPolicy{}, fmt.Errorf("failed to get policy: %w", err) diff --git a/storage/db.go b/storage/db.go index 1a1d293..5b52de7 100644 --- a/storage/db.go +++ b/storage/db.go @@ -14,9 +14,9 @@ import ( type DatabaseStorage interface { Close() error - GetPluginPolicy(ctx context.Context, id string) (vtypes.PluginPolicy, error) + GetPluginPolicy(ctx context.Context, id uuid.UUID) (vtypes.PluginPolicy, error) GetAllPluginPolicies(ctx context.Context, publicKey string, pluginID vtypes.PluginID) ([]vtypes.PluginPolicy, error) - DeletePluginPolicyTx(ctx context.Context, dbTx pgx.Tx, id string) error + DeletePluginPolicyTx(ctx context.Context, dbTx pgx.Tx, id uuid.UUID) error InsertPluginPolicyTx(ctx context.Context, dbTx pgx.Tx, policy vtypes.PluginPolicy) (*vtypes.PluginPolicy, error) UpdatePluginPolicyTx(ctx context.Context, dbTx pgx.Tx, policy vtypes.PluginPolicy) (*vtypes.PluginPolicy, error) diff --git a/storage/postgres/policy.go b/storage/postgres/policy.go index d850f9d..36fb6b6 100644 --- a/storage/postgres/policy.go +++ b/storage/postgres/policy.go @@ -5,12 +5,13 @@ import ( "errors" "fmt" + "github.com/google/uuid" "github.com/jackc/pgx/v5" vtypes "github.com/vultisig/verifier/types" ) -func (p *PostgresBackend) GetPluginPolicy(ctx context.Context, id string) (vtypes.PluginPolicy, error) { +func (p *PostgresBackend) GetPluginPolicy(ctx context.Context, id uuid.UUID) (vtypes.PluginPolicy, error) { if p.pool == nil { return vtypes.PluginPolicy{}, fmt.Errorf("database pool is nil") } @@ -151,7 +152,7 @@ func (p *PostgresBackend) UpdatePluginPolicyTx(ctx context.Context, dbTx pgx.Tx, return &updatedPolicy, nil } -func (p *PostgresBackend) DeletePluginPolicyTx(ctx context.Context, dbTx pgx.Tx, id string) error { +func (p *PostgresBackend) DeletePluginPolicyTx(ctx context.Context, dbTx pgx.Tx, id uuid.UUID) error { _, err := dbTx.Exec(ctx, ` DELETE FROM transaction_history WHERE policy_id = $1 From 04d3c759c90c6e3401f1751c9408e8a5b18b7eb3 Mon Sep 17 00:00:00 2001 From: Johnny Luo Date: Fri, 6 Jun 2025 19:37:46 +1000 Subject: [PATCH 2/3] fix: simplify policy update validation by using conversion method --- api/plugin.go | 14 ++------------ 1 file changed, 2 insertions(+), 12 deletions(-) diff --git a/api/plugin.go b/api/plugin.go index d03bae3..a3b12ca 100644 --- a/api/plugin.go +++ b/api/plugin.go @@ -59,18 +59,8 @@ func (s *Server) SignPluginMessages(c echo.Context) error { if policy.PluginID.String() != req.PluginID { return fmt.Errorf("policy plugin ID mismatch") } - policyUpdate := vtypes.PluginPolicyCreateUpdate{ - ID: policy.ID, - PublicKey: policy.PublicKey, - PluginID: policy.PluginID, - PluginVersion: policy.PluginVersion, - PolicyVersion: policy.PolicyVersion, - Signature: policy.Signature, - Recipe: policy.Recipe, - BillingRecipe: "", - Active: policy.Active, - } - if err := s.plugin.ValidateProposedTransactions(policyUpdate, []vtypes.PluginKeysignRequest{req}); err != nil { + + if err := s.plugin.ValidateProposedTransactions(policy.ToPluginPolicyCreateUpdate(), []vtypes.PluginKeysignRequest{req}); err != nil { return fmt.Errorf("failed to validate transaction proposal: %w", err) } From da767f354870400ab6d74cd3a7d1c9c0ecdbefbd Mon Sep 17 00:00:00 2001 From: Johnny Luo Date: Fri, 6 Jun 2025 19:38:21 +1000 Subject: [PATCH 3/3] fix: remove unnecessary parameter from verifyPolicySignature method --- api/plugin.go | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/api/plugin.go b/api/plugin.go index a3b12ca..fb78bbd 100644 --- a/api/plugin.go +++ b/api/plugin.go @@ -205,7 +205,7 @@ func (s *Server) CreatePluginPolicy(c echo.Context) error { policy.ID = uuid.New() } - if !s.verifyPolicySignature(policy.ToPluginPolicy(), false) { + if !s.verifyPolicySignature(policy.ToPluginPolicy()) { s.logger.Error("invalid policy signature") return c.JSON(http.StatusForbidden, NewErrorResponse("Invalid policy signature")) } @@ -233,7 +233,7 @@ func (s *Server) UpdatePluginPolicyById(c echo.Context) error { return c.JSON(http.StatusBadRequest, NewErrorResponse("failed to validate policy")) } - if !s.verifyPolicySignature(policy.ToPluginPolicy(), true) { + if !s.verifyPolicySignature(policy.ToPluginPolicy()) { s.logger.Error("invalid policy signature") return c.JSON(http.StatusForbidden, NewErrorResponse("Invalid policy signature")) } @@ -278,7 +278,7 @@ func (s *Server) DeletePluginPolicyById(c echo.Context) error { // This is because we have different signature stored in the database. policy.Signature = reqBody.Signature - if !s.verifyPolicySignature(policy, true) { + if !s.verifyPolicySignature(policy) { return c.JSON(http.StatusForbidden, NewErrorResponse("Invalid policy signature")) } @@ -337,7 +337,7 @@ func (s *Server) GetPluginPolicyTransactionHistory(c echo.Context) error { return c.JSON(http.StatusOK, policyHistory) } -func (s *Server) verifyPolicySignature(policy vtypes.PluginPolicy, update bool) bool { +func (s *Server) verifyPolicySignature(policy vtypes.PluginPolicy) bool { msgBytes, err := policyToMessageHex(policy) if err != nil { s.logger.WithError(err).Error("Failed to convert policy to message hex")