From f8d92f5905b419908323ebaa53c18af3feea0c12 Mon Sep 17 00:00:00 2001 From: Loris Leiva Date: Fri, 25 Sep 2026 12:20:51 +0100 Subject: [PATCH] Adapt instruction visitors to Codama v2 --- packages/errors/src/codes.ts | 2 + packages/errors/src/context.ts | 10 +- packages/errors/src/messages.ts | 5 +- packages/visitors/README.md | 31 ++- ...tInstructionAccountDefaultValuesVisitor.ts | 218 +++++++++++------- .../setInstructionDiscriminatorsVisitor.ts | 179 +++++++++++--- .../visitors/src/setNumberWrappersVisitor.ts | 195 ++++++++++++++-- ...ructionAccountDefaultValuesVisitor.test.ts | 167 ++++++++++++++ ...etInstructionDiscriminatorsVisitor.test.ts | 200 ++++++++++++++++ .../test/setNumberWrappersVisitor.test.ts | 171 ++++++++++++++ 10 files changed, 1039 insertions(+), 139 deletions(-) create mode 100644 packages/visitors/test/setInstructionAccountDefaultValuesVisitor.test.ts create mode 100644 packages/visitors/test/setInstructionDiscriminatorsVisitor.test.ts create mode 100644 packages/visitors/test/setNumberWrappersVisitor.test.ts diff --git a/packages/errors/src/codes.ts b/packages/errors/src/codes.ts index 67a3887ab..3ef36e65b 100644 --- a/packages/errors/src/codes.ts +++ b/packages/errors/src/codes.ts @@ -61,6 +61,7 @@ export const CODAMA_ERROR__VISITORS__UNRECOGNIZED_UPDATE_KEYS = 1200015; export const CODAMA_ERROR__VISITORS__INSTRUCTION_DATA_FIELD_NOT_FOUND = 1200016; export const CODAMA_ERROR__VISITORS__INSTRUCTION_ACCOUNT_NOT_FOUND = 1200017; export const CODAMA_ERROR__VISITORS__DEFINED_TYPE_MEMBER_NOT_FOUND = 1200018; +export const CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR = 1200019; // Anchor-related errors. // Reserve error codes in the range [2100000-2100999]. @@ -166,6 +167,7 @@ export type CodamaErrorCode = | typeof CODAMA_ERROR__VISITORS__CANNOT_FLATTEN_STRUCT_WITH_CONFLICTING_ATTRIBUTES | typeof CODAMA_ERROR__VISITORS__CANNOT_FLATTEN_STRUCT_WITH_PLUGINS | typeof CODAMA_ERROR__VISITORS__CANNOT_REMOVE_LAST_PATH_IN_NODE_STACK + | typeof CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR | typeof CODAMA_ERROR__VISITORS__CANNOT_USE_OPTIONAL_ACCOUNT_AS_PDA_SEED_VALUE | typeof CODAMA_ERROR__VISITORS__CYCLIC_DEPENDENCY_DETECTED_WHEN_RESOLVING_INSTRUCTION_DEFAULT_VALUES | typeof CODAMA_ERROR__VISITORS__DEFINED_TYPE_MEMBER_NOT_FOUND diff --git a/packages/errors/src/context.ts b/packages/errors/src/context.ts index 87d6316ac..570ebe09a 100644 --- a/packages/errors/src/context.ts +++ b/packages/errors/src/context.ts @@ -73,6 +73,7 @@ import { CODAMA_ERROR__VISITORS__CANNOT_FLATTEN_STRUCT_WITH_CONFLICTING_ATTRIBUTES, CODAMA_ERROR__VISITORS__CANNOT_FLATTEN_STRUCT_WITH_PLUGINS, CODAMA_ERROR__VISITORS__CANNOT_REMOVE_LAST_PATH_IN_NODE_STACK, + CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR, CODAMA_ERROR__VISITORS__CANNOT_USE_OPTIONAL_ACCOUNT_AS_PDA_SEED_VALUE, CODAMA_ERROR__VISITORS__CYCLIC_DEPENDENCY_DETECTED_WHEN_RESOLVING_INSTRUCTION_DEFAULT_VALUES, CODAMA_ERROR__VISITORS__DEFINED_TYPE_MEMBER_NOT_FOUND, @@ -284,6 +285,11 @@ export type CodamaErrorContext = DefaultUnspecifiedErrorContextToUndefined<{ [CODAMA_ERROR__VISITORS__CANNOT_REMOVE_LAST_PATH_IN_NODE_STACK]: { path: readonly Node[]; }; + [CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR]: { + instruction: InstructionNode; + instructionName: IdentifierString; + reason: string; + }; [CODAMA_ERROR__VISITORS__CANNOT_USE_OPTIONAL_ACCOUNT_AS_PDA_SEED_VALUE]: { instruction: InstructionNode; instructionAccount: InstructionAccountNode; @@ -334,7 +340,9 @@ export type CodamaErrorContext = DefaultUnspecifiedErrorContextToUndefined<{ parentName: IdentifierString | PathString; }; [CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER]: { - wrapper: string; + kind: string; + reason: string; + wrapper: object; }; [CODAMA_ERROR__VISITORS__INVALID_PDA_SEED_VALUES]: { instruction: InstructionNode; diff --git a/packages/errors/src/messages.ts b/packages/errors/src/messages.ts index d633fddc1..c3752a294 100644 --- a/packages/errors/src/messages.ts +++ b/packages/errors/src/messages.ts @@ -53,6 +53,7 @@ import { CODAMA_ERROR__VISITORS__CANNOT_FLATTEN_STRUCT_WITH_CONFLICTING_ATTRIBUTES, CODAMA_ERROR__VISITORS__CANNOT_FLATTEN_STRUCT_WITH_PLUGINS, CODAMA_ERROR__VISITORS__CANNOT_REMOVE_LAST_PATH_IN_NODE_STACK, + CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR, CODAMA_ERROR__VISITORS__CANNOT_USE_OPTIONAL_ACCOUNT_AS_PDA_SEED_VALUE, CODAMA_ERROR__VISITORS__CYCLIC_DEPENDENCY_DETECTED_WHEN_RESOLVING_INSTRUCTION_DEFAULT_VALUES, CODAMA_ERROR__VISITORS__DEFINED_TYPE_MEMBER_NOT_FOUND, @@ -147,6 +148,8 @@ export const CodamaErrorMessages: Readonly<{ [CODAMA_ERROR__VISITORS__CANNOT_FLATTEN_STRUCT_WITH_PLUGINS]: 'Cannot flatten the struct of field [$fieldName] since it carries plugins that would be lost.', [CODAMA_ERROR__VISITORS__CANNOT_REMOVE_LAST_PATH_IN_NODE_STACK]: 'Cannot remove the last path in the node stack.', + [CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR]: + 'Cannot set the discriminator of instruction [$instructionName]: $reason.', [CODAMA_ERROR__VISITORS__CANNOT_USE_OPTIONAL_ACCOUNT_AS_PDA_SEED_VALUE]: 'Cannot use optional account [$seedValueName] as the [$seedName] PDA seed for the [$instructionAccountName] account of the [$instructionName] instruction.', [CODAMA_ERROR__VISITORS__CYCLIC_DEPENDENCY_DETECTED_WHEN_RESOLVING_INSTRUCTION_DEFAULT_VALUES]: @@ -162,7 +165,7 @@ export const CodamaErrorMessages: Readonly<{ 'Could not find an enum data field named [$fieldName] for instruction [$instructionName].', [CODAMA_ERROR__VISITORS__INVALID_INSTRUCTION_DEFAULT_VALUE_DEPENDENCY]: 'Dependency [$dependencyName] of kind [$dependencyKind] is not a valid dependency of [$parentName] of kind [$parentKind] in the [$instructionName] instruction.', - [CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER]: 'Invalid number wrapper kind [$wrapper].', + [CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER]: 'Invalid number wrapper [$kind]: $reason.', [CODAMA_ERROR__VISITORS__INVALID_PDA_SEED_VALUES]: 'Invalid seed values for PDA [$pdaName] in instruction [$instructionName].', [CODAMA_ERROR__VISITORS__INVALID_PROVIDED_VALUE]: diff --git a/packages/visitors/README.md b/packages/visitors/README.md index 6f4049d4b..a0d32769d 100644 --- a/packages/visitors/README.md +++ b/packages/visitors/README.md @@ -205,11 +205,14 @@ codama.update(setFixedAccountSizesVisitor()); ### `setInstructionAccountDefaultValuesVisitor` -This visitor helps set the default values of instruction accounts in bulk. It accepts an array of "rule" objects that must contain the default value to set and the name of the instruction account to set it on. The account name may also be a regular expression to match more complex patterns. +This visitor helps set the default values of instruction accounts in bulk, including the accounts of sub-instructions. It accepts an array of "rule" objects that must contain the default value to set and the identifier of the instruction account to set it on (matched exactly). The account identifier may also be a regular expression to match more complex patterns. Rules restricted to an `instruction` take precedence over the others, and `ignoreIfOptional` leaves optional or already defaulted accounts untouched. + +Missing seeds of `PdaValueNode` default values are filled using the `fillDefaultPdaSeedValuesVisitor`; a rule whose seeds cannot all be filled is skipped for that account. The `getCommonInstructionAccountDefaultRules` function returns rules for common accounts such as payers, authorities, well-known programs and sysvars, matching both their camelCase and snake_case identifiers. ```ts codama.update( setInstructionAccountDefaultValuesVisitor([ + ...getCommonInstructionAccountDefaultRules(), { // Set this public key as default value to any account named 'counterProgram'. account: 'counterProgram', @@ -226,28 +229,42 @@ codama.update( ### `setInstructionDiscriminatorsVisitor` -This visitor adds a new instruction argument to each of the provided instruction names. The new argument is added before any existing argument and marked as a discriminator of the instruction. This is useful if your Codama IDL is missing discriminators in the instruction data. +This visitor adds a discriminator to the data of each of the provided instructions. This is useful if your Codama IDL is missing discriminators in the instruction data. + +- When the instruction data is an inline `StructTypeNode` without transforms (or absent), the discriminator is added as its first field (named `discriminator` by default) and a `FieldDiscriminatorNode` pointing to it is added. +- Otherwise, such as when the data links to a defined type or carries transforms, the discriminator is added as a `HiddenPrefixTransformNode` on the data, leaving the defined type untouched, and a `ConstantDiscriminatorNode` is added. + +The discriminator type defaults to `u8` and must have a fixed size, since the offsets of existing discriminators are shifted by it. ```ts codama.update( setInstructionDiscriminatorsVisitor({ - mint: { name: 'discriminator', type: numberTypeNode('u8'), value: numberValueNode(0) }, - transfer: { name: 'discriminator', type: numberTypeNode('u8'), value: numberValueNode(1) }, - burn: { name: 'discriminator', type: numberTypeNode('u8'), value: numberValueNode(2) }, + mint: { value: integerValueNode('0') }, + transfer: { value: integerValueNode('1') }, + burn: { identifier: 'kind', type: integerTypeNode('u32'), value: integerValueNode('2') }, }), ); ``` ### `setNumberWrappersVisitor` -This visitor helps wrap `NumberTypeNodes` matching a given name with a specific number wrapper. +This visitor gives semantic meaning to the numbers matching the provided `NodeSelectors`, using the following wrappers: + +- `FixedPoint` and `SolAmount` wrap an integer in a `FixedPointTypeNode`, `SolAmount` being a fixed point of scale 9 in `SOL`. +- `DateTime` and `Duration` wrap an integer in a `DateTimeTypeNode` or a `DurationTypeNode`. +- `Unit` sets the `unit` of an integer or a float. +- `AmountDisplay` and `UnitDisplay` set the `display` of an integer to an `AmountNumberDisplayNode` or a `UnitNumberDisplayNode`. `UnitDisplay` also applies to floats. + +Wrappers carry the transforms of the number they wrap. Integers used as sizes or prefixes, and numbers within the type of a constant, are left untouched. ```ts codama.update( setNumberWrappersVisitor({ lamports: { kind: 'SolAmount' }, timestamp: { kind: 'DateTime' }, - percent: { decimals: 2, kind: 'Amount', unit: '%' }, + 'mint.supply': { kind: 'FixedPoint', scale: 6, unit: 'USDC' }, + 'transfer.amount': { decimals: injectedValueNode({ key: 'decimals' }), kind: 'AmountDisplay' }, + percent: { kind: 'Unit', unit: '%' }, }), ); ``` diff --git a/packages/visitors/src/setInstructionAccountDefaultValuesVisitor.ts b/packages/visitors/src/setInstructionAccountDefaultValuesVisitor.ts index b7309e2a3..f111fe11a 100644 --- a/packages/visitors/src/setInstructionAccountDefaultValuesVisitor.ts +++ b/packages/visitors/src/setInstructionAccountDefaultValuesVisitor.ts @@ -1,7 +1,10 @@ -import { camelCase } from '@codama/fragments/casing'; +import { CODAMA_ERROR__VISITORS__INVALID_PDA_SEED_VALUES, isCodamaError } from '@codama/errors'; +import { snakeCase } from '@codama/fragments/casing'; import { + assertIsNode, identityValueNode, InstructionAccountNode, + instructionAccountNode, InstructionInputValueNode, InstructionNode, instructionNode, @@ -10,192 +13,247 @@ import { publicKeyValueNode, } from '@codama/nodes'; import { - extendVisitor, + bottomUpTransformerVisitor, LinkableDictionary, - NodeStack, - nonNullableIdentityVisitor, + NodePath, pipe, recordLinkablesOnFirstVisitVisitor, - recordNodeStackVisitor, visit, } from '@codama/visitors-core'; import { fillDefaultPdaSeedValuesVisitor } from './fillDefaultPdaSeedValuesVisitor'; export type InstructionAccountDefaultRule = { - /** The name of the instruction account or a pattern to match on it. */ + /** The identifier of the instruction account (matched exactly) or a pattern to match on it. */ account: RegExp | string; /** The default value to assign to it. */ defaultValue: InstructionInputValueNode; - /** @defaultValue `false`. */ + /** + * Whether to leave the account untouched when it is optional or + * already has a default value. + * @defaultValue `false`. + */ ignoreIfOptional?: boolean; - /** @defaultValue Defaults to searching accounts on all instructions. */ + /** + * The identifier of the instruction to restrict the rule to (matched exactly). + * @defaultValue Defaults to searching accounts on all instructions. + */ instruction?: string; }; +/** + * Match any of the given camelCase identifiers, as is or in snake_case, + * since identifiers keep the casing of the program they come from. + */ +function anyIdentifierOf(...identifiers: string[]): RegExp { + const alternatives = [...new Set(identifiers.flatMap(identifier => [identifier, snakeCase(identifier)]))]; + return new RegExp(`^(${alternatives.join('|')})$`); +} + +/** + * Default value rules for commonly used accounts (payers, authorities, + * well-known programs and sysvars), matching their camelCase and + * snake_case identifiers. + */ export const getCommonInstructionAccountDefaultRules = (): InstructionAccountDefaultRule[] => [ { - account: /^(payer|feePayer)$/, + account: anyIdentifierOf('payer', 'feePayer'), defaultValue: payerValueNode(), ignoreIfOptional: true, }, { - account: /^(authority)$/, + account: anyIdentifierOf('authority'), defaultValue: identityValueNode(), ignoreIfOptional: true, }, { - account: /^(programId)$/, + account: anyIdentifierOf('programId'), defaultValue: programIdValueNode(), ignoreIfOptional: true, }, { - account: /^(systemProgram|splSystemProgram)$/, - defaultValue: publicKeyValueNode('11111111111111111111111111111111', 'splSystem'), + account: anyIdentifierOf('systemProgram', 'splSystemProgram'), + defaultValue: publicKeyValueNode('11111111111111111111111111111111', { identifier: 'splSystem' }), ignoreIfOptional: true, }, { - account: /^(tokenProgram|splTokenProgram)$/, - defaultValue: publicKeyValueNode('TokenkegQfeZyiNwAJbNbGKPFXCWuBvf9Ss623VQ5DA', 'splToken'), + account: anyIdentifierOf('tokenProgram', 'splTokenProgram'), + defaultValue: publicKeyValueNode('TokenkegQfeZyiNwAJbNbGKPFXCWuBvf9Ss623VQ5DA', { identifier: 'splToken' }), ignoreIfOptional: true, }, { - account: /^(ataProgram|splAtaProgram)$/, - defaultValue: publicKeyValueNode('ATokenGPvbdGVxr1b2hvZbsiqW5xWH25efTNsLJA8knL', 'splAssociatedToken'), + account: anyIdentifierOf('ataProgram', 'splAtaProgram'), + defaultValue: publicKeyValueNode('ATokenGPvbdGVxr1b2hvZbsiqW5xWH25efTNsLJA8knL', { + identifier: 'splAssociatedToken', + }), ignoreIfOptional: true, }, { - account: /^(tokenMetadataProgram|mplTokenMetadataProgram)$/, - defaultValue: publicKeyValueNode('metaqbxxUerdq28cj1RbAWkYQm3ybzjb6a8bt518x1s', 'mplTokenMetadata'), + account: anyIdentifierOf('tokenMetadataProgram', 'mplTokenMetadataProgram'), + defaultValue: publicKeyValueNode('metaqbxxUerdq28cj1RbAWkYQm3ybzjb6a8bt518x1s', { + identifier: 'mplTokenMetadata', + }), ignoreIfOptional: true, }, { - account: /^(tokenAuth|mplTokenAuth|authorization|mplAuthorization|auth|mplAuth)RulesProgram$/, - defaultValue: publicKeyValueNode('auth9SigNpDKz4sJJ1DfCTuZrZNSAgh9sFD3rboVmgg', 'mplTokenAuthRules'), + account: anyIdentifierOf( + 'tokenAuthRulesProgram', + 'mplTokenAuthRulesProgram', + 'authorizationRulesProgram', + 'mplAuthorizationRulesProgram', + 'authRulesProgram', + 'mplAuthRulesProgram', + ), + defaultValue: publicKeyValueNode('auth9SigNpDKz4sJJ1DfCTuZrZNSAgh9sFD3rboVmgg', { + identifier: 'mplTokenAuthRules', + }), ignoreIfOptional: true, }, { - account: /^(candyMachineProgram|mplCandyMachineProgram)$/, - defaultValue: publicKeyValueNode('CndyV3LdqHUfDLmE5naZjVN8rBZz4tqhdefbAnjHG3JR', 'mplCandyMachine'), + account: anyIdentifierOf('candyMachineProgram', 'mplCandyMachineProgram'), + defaultValue: publicKeyValueNode('CndyV3LdqHUfDLmE5naZjVN8rBZz4tqhdefbAnjHG3JR', { + identifier: 'mplCandyMachine', + }), ignoreIfOptional: true, }, { - account: /^(candyGuardProgram|mplCandyGuardProgram)$/, - defaultValue: publicKeyValueNode('Guard1JwRhJkVH6XZhzoYxeBVQe872VH6QggF4BWmS9g', 'mplCandyGuard'), + account: anyIdentifierOf('candyGuardProgram', 'mplCandyGuardProgram'), + defaultValue: publicKeyValueNode('Guard1JwRhJkVH6XZhzoYxeBVQe872VH6QggF4BWmS9g', { + identifier: 'mplCandyGuard', + }), ignoreIfOptional: true, }, { - account: /^(clockSysvar|sysvarClock)$/, + account: anyIdentifierOf('clockSysvar', 'sysvarClock'), defaultValue: publicKeyValueNode('SysvarC1ock11111111111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(epochScheduleSysvar|sysvarEpochSchedule)$/, + account: anyIdentifierOf('epochScheduleSysvar', 'sysvarEpochSchedule'), defaultValue: publicKeyValueNode('SysvarEpochSchedu1e111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(instructions?Sysvar|sysvarInstructions?)(Account)?$/, + account: anyIdentifierOf( + ...['instructionSysvar', 'instructionsSysvar', 'sysvarInstruction', 'sysvarInstructions'].flatMap( + identifier => [identifier, `${identifier}Account`], + ), + ), defaultValue: publicKeyValueNode('Sysvar1nstructions1111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(recentBlockhashesSysvar|sysvarRecentBlockhashes)$/, + account: anyIdentifierOf('recentBlockhashesSysvar', 'sysvarRecentBlockhashes'), defaultValue: publicKeyValueNode('SysvarRecentB1ockHashes11111111111111111111'), ignoreIfOptional: true, }, { - account: /^(rent|rentSysvar|sysvarRent)$/, + account: anyIdentifierOf('rent', 'rentSysvar', 'sysvarRent'), defaultValue: publicKeyValueNode('SysvarRent111111111111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(rewardsSysvar|sysvarRewards)$/, + account: anyIdentifierOf('rewardsSysvar', 'sysvarRewards'), defaultValue: publicKeyValueNode('SysvarRewards111111111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(slotHashesSysvar|sysvarSlotHashes)$/, + account: anyIdentifierOf('slotHashesSysvar', 'sysvarSlotHashes'), defaultValue: publicKeyValueNode('SysvarS1otHashes111111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(slotHistorySysvar|sysvarSlotHistory)$/, + account: anyIdentifierOf('slotHistorySysvar', 'sysvarSlotHistory'), defaultValue: publicKeyValueNode('SysvarS1otHistory11111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(stakeHistorySysvar|sysvarStakeHistory)$/, + account: anyIdentifierOf('stakeHistorySysvar', 'sysvarStakeHistory'), defaultValue: publicKeyValueNode('SysvarStakeHistory1111111111111111111111111'), ignoreIfOptional: true, }, { - account: /^(mplCoreProgram)$/, - defaultValue: publicKeyValueNode('CoREENxT6tW1HoK8ypY1SxRMZTcVPm7R94rH4PZNhX7d', 'mplCore'), + account: anyIdentifierOf('mplCoreProgram'), + defaultValue: publicKeyValueNode('CoREENxT6tW1HoK8ypY1SxRMZTcVPm7R94rH4PZNhX7d', { identifier: 'mplCore' }), ignoreIfOptional: true, }, ]; +/** + * Set the default values of instruction accounts, including those of + * sub-instructions, in bulk using the given rules. + * + * Rules restricted to an instruction take precedence over the others; + * otherwise, the first matching rule wins. Missing seeds of PDA default + * values are filled from the instruction's accounts and data (see + * `fillDefaultPdaSeedValuesVisitor`): a rule whose PDA seeds cannot all be + * filled is skipped for that account. + * + * @example + * ```ts + * setInstructionAccountDefaultValuesVisitor([ + * ...getCommonInstructionAccountDefaultRules(), + * { account: 'counterProgram', defaultValue: publicKeyValueNode('MyCounterProgram11111111111111111111111111') }, + * { account: /^(associatedToken|ata)$/, defaultValue: pdaValueNode('associatedToken') }, + * ]); + * ``` + */ export function setInstructionAccountDefaultValuesVisitor(rules: InstructionAccountDefaultRule[]) { const linkables = new LinkableDictionary(); - const stack = new NodeStack(); - // Place the rules with instructions first. - const sortedRules = rules.sort((a, b) => { - const ia = 'instruction' in a; - const ib = 'instruction' in b; - if ((ia && ib) || (!a && !ib)) return 0; - return ia ? -1 : 1; - }); + // Place the rules with instructions first, without mutating the given rules. + const sortedRules = [ + ...rules.filter(rule => rule.instruction !== undefined), + ...rules.filter(rule => rule.instruction === undefined), + ]; - function matchRule( + const matchRule = ( instruction: InstructionNode, account: InstructionAccountNode, - ): InstructionAccountDefaultRule | undefined { - return sortedRules.find(rule => { - if ('instruction' in rule && rule.instruction && camelCase(rule.instruction) !== instruction.identifier) { - return false; - } + ): InstructionAccountDefaultRule | undefined => + sortedRules.find(rule => { + if (rule.instruction !== undefined && rule.instruction !== instruction.identifier) return false; return typeof rule.account === 'string' - ? camelCase(rule.account) === account.identifier + ? rule.account === account.identifier : rule.account.test(account.identifier); }); - } + + const applyRule = ( + account: InstructionAccountNode, + rule: InstructionAccountDefaultRule, + instructionPath: NodePath, + ): InstructionAccountNode => { + if ((rule.ignoreIfOptional ?? false) && (account.isOptional || !!account.defaultValue)) return account; + try { + const defaultValue = visit( + rule.defaultValue, + fillDefaultPdaSeedValuesVisitor(instructionPath, linkables, true), + ); + return instructionAccountNode({ ...account, defaultValue }); + } catch (error) { + // The rule does not apply when its PDA seeds cannot all be filled. + if (isCodamaError(error, CODAMA_ERROR__VISITORS__INVALID_PDA_SEED_VALUES)) return account; + throw error; + } + }; return pipe( - nonNullableIdentityVisitor({ keys: ['rootNode', 'programNode', 'instructionNode'] }), - v => - extendVisitor(v, { - visitInstruction(node) { + bottomUpTransformerVisitor([ + { + select: '[instructionNode]', + transform: (node, stack) => { + assertIsNode(node, 'instructionNode'); const instructionPath = stack.getPath('instructionNode'); - const instructionAccounts = (node.accounts ?? []).map((account): InstructionAccountNode => { - const rule = matchRule(node, account); - if (!rule) return account; - - if ((rule.ignoreIfOptional ?? false) && (account.isOptional || !!account.defaultValue)) { - return account; - } - - try { - return { - ...account, - defaultValue: visit( - rule.defaultValue, - fillDefaultPdaSeedValuesVisitor(instructionPath, linkables, true), - ), - }; - } catch { - return account; - } - }); - return instructionNode({ ...node, - accounts: instructionAccounts, + accounts: (node.accounts ?? []).map(account => { + const rule = matchRule(node, account); + return rule ? applyRule(account, rule, instructionPath) : account; + }), }); }, - }), - v => recordNodeStackVisitor(v, stack), + }, + ]), v => recordLinkablesOnFirstVisitVisitor(v, linkables), ); } diff --git a/packages/visitors/src/setInstructionDiscriminatorsVisitor.ts b/packages/visitors/src/setInstructionDiscriminatorsVisitor.ts index 98eda168b..7f65b1dc9 100644 --- a/packages/visitors/src/setInstructionDiscriminatorsVisitor.ts +++ b/packages/visitors/src/setInstructionDiscriminatorsVisitor.ts @@ -1,49 +1,170 @@ +import { CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR, CodamaError } from '@codama/errors'; import { + addTypeNodeTransforms, assertIsNode, + constantDiscriminatorNode, + constantValueNode, + DiscriminatorNode, fieldDiscriminatorNode, - instructionArgumentNode, + hiddenPrefixTransformNode, + InstructionNode, instructionNode, - numberTypeNode, + integerTypeNode, + isNode, + sizeDiscriminatorNode, + structFieldTypeNode, + structTypeNode, + TextNode, TypeNode, ValueNode, } from '@codama/nodes'; -import { BottomUpNodeTransformerWithSelector, bottomUpTransformerVisitor } from '@codama/visitors-core'; - -type Discriminator = { - /** @defaultValue `[]` */ - docs?: string[]; - /** @defaultValue `"discriminator"` */ - name?: string; - /** @defaultValue `"omitted"` */ +import { + BottomUpNodeTransformerWithSelector, + bottomUpTransformerVisitor, + getByteSizeVisitor, + LinkableDictionary, + pipe, + recordLinkablesOnFirstVisitVisitor, + visit, +} from '@codama/visitors-core'; + +import { assertValidUpdateKeys } from './updateHelpers'; + +export type InstructionDiscriminator = { + /** Only used when the discriminator is added as a data field. */ + docs?: TextNode | string; + /** + * The identifier of the discriminator field, when added as a data field. + * @defaultValue `"discriminator"` + */ + identifier?: string; + /** + * The default value strategy of the discriminator field. Only `omitted` + * is supported when the discriminator is added as a hidden prefix. + * @defaultValue `"omitted"` + */ strategy?: 'omitted' | 'optional'; - /** @defaultValue `numberTypeNode('u8')` */ + /** + * The type of the discriminator, which must have a fixed size. + * @defaultValue `integerTypeNode('u8')` + */ type?: TypeNode; + /** The value of the discriminator. */ value: ValueNode; }; -export function setInstructionDiscriminatorsVisitor(map: Record) { - return bottomUpTransformerVisitor( - Object.entries(map).map(([selector, discriminator]): BottomUpNodeTransformerWithSelector => ({ +const DISCRIMINATOR_KEYS = ['docs', 'identifier', 'strategy', 'type', 'value']; + +/** + * Prepend a discriminator to the data of the selected instructions. + * + * - When the instruction data is an inline struct without transforms (or + * absent), the discriminator is added as its first field, with the given value as + * default value, and a `fieldDiscriminatorNode` pointing to it is added. + * - Otherwise (e.g. a `definedTypeLinkNode` or a struct with transforms, + * such as a size prefix), the discriminator is added as + * a `hiddenPrefixTransformNode` wrapping the data, leaving any shared + * defined type untouched, and a `constantDiscriminatorNode` is added. + * + * Since the discriminator is written before the rest of the data, the + * offsets of existing field and constant discriminators and the size of + * existing size discriminators are shifted by its size. + * + * @throws {CODAMA_ERROR__VISITORS__UNRECOGNIZED_UPDATE_KEYS} if a + * discriminator contains an unrecognised key (e.g. `name` instead of + * `identifier`). + * @throws {CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR} if + * the data already has a field with the discriminator's identifier, if the + * `optional` strategy is used with a hidden prefix, or if the discriminator + * type does not have a fixed size. + * + * @example + * ```ts + * setInstructionDiscriminatorsVisitor({ + * mint: { value: integerValueNode('0') }, + * transfer: { identifier: 'kind', type: integerTypeNode('u32'), value: integerValueNode('1') }, + * }); + * ``` + */ +export function setInstructionDiscriminatorsVisitor(map: Record) { + const linkables = new LinkableDictionary(); + + const transformers = Object.entries(map).map(([selector, discriminator]): BottomUpNodeTransformerWithSelector => { + assertValidUpdateKeys(selector, discriminator, DISCRIMINATOR_KEYS); + return { select: ['[instructionNode]', selector], - transform: node => { + transform: (node, stack) => { assertIsNode(node, 'instructionNode'); - const discriminatorArgument = instructionArgumentNode({ - defaultValue: discriminator.value, - defaultValueStrategy: discriminator.strategy ?? 'omitted', - docs: discriminator.docs ?? [], - identifier: discriminator.name ?? 'discriminator', - type: discriminator.type ?? numberTypeNode('u8'), - }); + const type = discriminator.type ?? integerTypeNode('u8'); + const size = visit(type, getByteSizeVisitor(linkables, { stack: stack.clone() })); + if (size === null) throw cannotSet(node, 'the discriminator type must have a fixed size'); + const discriminators = shiftDiscriminators(node.discriminators ?? [], size); + + // A field is only at byte 0 of a struct that has no transforms. + const isPlainStruct = isNode(node.data, 'structTypeNode') && (node.data.transforms ?? []).length === 0; + if (node.data === undefined || isPlainStruct) { + const identifier = discriminator.identifier ?? 'discriminator'; + const fields = isNode(node.data, 'structTypeNode') ? (node.data.fields ?? []) : []; + if (fields.some(field => field.identifier === identifier)) { + throw cannotSet(node, `the data already has a field named \`${identifier}\``); + } + const field = structFieldTypeNode({ + defaultValue: discriminator.value, + defaultValueStrategy: discriminator.strategy ?? 'omitted', + docs: discriminator.docs, + identifier, + type, + }); + return instructionNode({ + ...node, + data: structTypeNode([field, ...fields], { ...node.data }), + discriminators: [fieldDiscriminatorNode(identifier), ...discriminators], + }); + } + if (discriminator.strategy === 'optional') { + throw cannotSet( + node, + 'the `optional` strategy is not supported when the discriminator is added as a hidden prefix', + ); + } + const constant = constantValueNode(type, discriminator.value); return instructionNode({ ...node, - arguments: [discriminatorArgument, ...(node.arguments ?? [])], - discriminators: [ - fieldDiscriminatorNode(discriminator.name ?? 'discriminator'), - ...(node.discriminators ?? []), - ], + data: addTypeNodeTransforms(node.data, [hiddenPrefixTransformNode([constant])]), + discriminators: [constantDiscriminatorNode(constant, { offset: 0 }), ...discriminators], }); }, - })), - ); + }; + }); + + return pipe(bottomUpTransformerVisitor(transformers), v => recordLinkablesOnFirstVisitVisitor(v, linkables)); +} + +/** Account for bytes prepended to the data in existing discriminators. */ +function shiftDiscriminators(discriminators: DiscriminatorNode[], size: number): DiscriminatorNode[] { + return discriminators.map(discriminator => { + switch (discriminator.kind) { + case 'fieldDiscriminatorNode': + return fieldDiscriminatorNode(discriminator.path, { + ...discriminator, + offset: discriminator.offset + size, + }); + case 'constantDiscriminatorNode': + return constantDiscriminatorNode(discriminator.constant, { + ...discriminator, + offset: discriminator.offset + size, + }); + case 'sizeDiscriminatorNode': + return sizeDiscriminatorNode(discriminator.size + size, { ...discriminator }); + } + }); +} + +function cannotSet(instruction: InstructionNode, reason: string): CodamaError { + return new CodamaError(CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR, { + instruction, + instructionName: instruction.identifier, + reason, + }); } diff --git a/packages/visitors/src/setNumberWrappersVisitor.ts b/packages/visitors/src/setNumberWrappersVisitor.ts index f85992ccf..3b2751137 100644 --- a/packages/visitors/src/setNumberWrappersVisitor.ts +++ b/packages/visitors/src/setNumberWrappersVisitor.ts @@ -1,31 +1,184 @@ import { CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER, CodamaError } from '@codama/errors'; -import { amountTypeNode, assertIsNestedTypeNode, dateTimeTypeNode, solAmountTypeNode } from '@codama/nodes'; -import { BottomUpNodeTransformerWithSelector, bottomUpTransformerVisitor } from '@codama/visitors-core'; +import { + amountNumberDisplayNode, + dateTimeTypeNode, + durationTypeNode, + fixedPointTypeNode, + FloatTypeNode, + floatTypeNode, + InjectableIntegerValueNode, + InjectableStringValueNode, + IntegerTypeNode, + integerTypeNode, + isNode, + Node, + NodeKind, + TypeNode, + unitNumberDisplayNode, +} from '@codama/nodes'; +import { BottomUpNodeTransformerWithSelector, bottomUpTransformerVisitor, NodePath } from '@codama/visitors-core'; +/** A semantic wrapper to apply to a number: a scaled quantity, a point in time, a duration, a unit or a display. */ export type NumberWrapper = - | { decimals: number; kind: 'Amount'; unit?: string } - | { kind: 'DateTime' } - | { kind: 'SolAmount' }; + | { base?: 2 | 10; kind: 'FixedPoint'; scale: number; unit?: string } + | { decimals: InjectableIntegerValueNode; kind: 'AmountDisplay'; unit?: InjectableStringValueNode } + | { kind: 'DateTime'; ticksPerSecond?: number } + | { kind: 'Duration'; ticksPerSecond?: number } + | { kind: 'SolAmount' } + | { kind: 'Unit'; unit: string } + | { kind: 'UnitDisplay'; unit: InjectableStringValueNode }; type NumberWrapperMap = Record; +const INTEGER_ONLY_KINDS = ['AmountDisplay', 'DateTime', 'Duration', 'FixedPoint', 'SolAmount']; +const INTEGER_AND_FLOAT_KINDS = ['Unit', 'UnitDisplay']; + +/** + * Nodes whose integer children are sizes, prefixes or already-wrapped + * numbers rather than values, so they are never wrapped. + */ +const NON_VALUE_INTEGER_PARENTS: NodeKind[] = [ + 'booleanTypeNode', + 'dateTimeTypeNode', + 'durationTypeNode', + 'enumTypeNode', + 'fixedPointTypeNode', + 'prefixedCountNode', + 'sizePrefixTransformNode', +]; + +/** + * Give semantic meaning to the numbers matching the given selectors, e.g. + * turn a `u64` into a token amount or a timestamp. + * + * - `FixedPoint` and `SolAmount` wrap the integer in a `fixedPointTypeNode` + * (`SolAmount` being a fixed point of scale 9 in `SOL`). + * - `DateTime` and `Duration` wrap the integer in a `dateTimeTypeNode` or + * a `durationTypeNode`. + * - `Unit` sets the unit of an integer or a float. + * - `AmountDisplay` and `UnitDisplay` set the display of an integer (or of + * a float, for `UnitDisplay`) to an `amountNumberDisplayNode` or a + * `unitNumberDisplayNode`. + * + * Wrappers carry the `transforms` of the number they wrap. Integers used as + * sizes or prefixes (e.g. an enum size or a size prefix), and numbers within + * the type of a constant (e.g. a hidden prefix), are left untouched. + * + * @throws {CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER} if the wrapper + * kind is unknown, if a fixed point has a zero scale or wraps a `shortU16`, + * or if a wrapped integer already carries a unit or display. + * + * @example + * ```ts + * setNumberWrappersVisitor({ + * 'mint.supply': { kind: 'FixedPoint', scale: 6, unit: 'USDC' }, + * lamports: { kind: 'SolAmount' }, + * createdAt: { kind: 'DateTime' }, + * 'transfer.amount': { decimals: injectedValueNode({ key: 'decimals' }), kind: 'AmountDisplay' }, + * }); + * ``` + */ export function setNumberWrappersVisitor(map: NumberWrapperMap) { return bottomUpTransformerVisitor( - Object.entries(map).map(([selectorStack, wrapper]): BottomUpNodeTransformerWithSelector => ({ - select: `${selectorStack}.[numberTypeNode]`, - transform: node => { - assertIsNestedTypeNode(node, 'numberTypeNode'); - switch (wrapper.kind) { - case 'DateTime': - return dateTimeTypeNode(node); - case 'SolAmount': - return solAmountTypeNode(node); - case 'Amount': - return amountTypeNode(node, wrapper.decimals, wrapper.unit); - default: - throw new CodamaError(CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER, { wrapper }); - } - }, - })), + Object.entries(map).map(([selector, wrapper]): BottomUpNodeTransformerWithSelector => { + assertValidWrapper(wrapper); + const kinds = INTEGER_AND_FLOAT_KINDS.includes(wrapper.kind) + ? '[integerTypeNode|floatTypeNode]' + : '[integerTypeNode]'; + return { + select: [`${selector}.${kinds}`, path => isValueNumber(path)], + transform: node => { + if (isNode(node, 'floatTypeNode')) return wrapFloat(node, wrapper); + if (isNode(node, 'integerTypeNode')) return wrapInteger(node, wrapper); + return node; + }, + }; + }), ); } + +/** Nodes whose descendants are the types of constants, which are never wrapped. */ +const CONSTANT_ANCESTORS: NodeKind[] = ['constantPdaSeedNode', 'constantValueNode']; + +/** + * Whether the number at the end of the path is a value rather than a size, + * a prefix, a wrapped number or (part of) the type of a constant. + */ +function isValueNumber(path: NodePath): boolean { + if (path.some(node => CONSTANT_ANCESTORS.includes(node.kind))) return false; + const number = path[path.length - 1]; + const parent = path[path.length - 2] as Node | undefined; + if (!parent) return true; + if (NON_VALUE_INTEGER_PARENTS.includes(parent.kind)) return false; + return !(isNode(parent, 'optionTypeNode') && parent.prefix === number && parent.item !== number); +} + +function assertValidWrapper(wrapper: NumberWrapper): void { + if (![...INTEGER_ONLY_KINDS, ...INTEGER_AND_FLOAT_KINDS].includes(wrapper.kind)) { + throw invalidWrapper(wrapper, 'unknown wrapper kind'); + } + if (wrapper.kind === 'FixedPoint' && wrapper.scale === 0) { + throw invalidWrapper(wrapper, 'a fixed point must have a non-zero scale; use a `Unit` wrapper instead'); + } +} + +function wrapInteger(number: IntegerTypeNode, wrapper: NumberWrapper): TypeNode { + switch (wrapper.kind) { + case 'Unit': + return integerTypeNode(number.format, { ...number, unit: wrapper.unit }); + case 'UnitDisplay': + return integerTypeNode(number.format, { + ...number, + display: unitNumberDisplayNode({ unit: wrapper.unit }), + }); + case 'AmountDisplay': + return integerTypeNode(number.format, { + ...number, + display: amountNumberDisplayNode({ decimals: wrapper.decimals, unit: wrapper.unit }), + }); + default: + return wrapIntegerInTypeNode(number, wrapper); + } +} + +/** Wrap an integer in a fixed point, date-time or duration, which carries its transforms. */ +function wrapIntegerInTypeNode( + number: IntegerTypeNode, + wrapper: Extract, +): TypeNode { + if (number.unit !== undefined || number.display !== undefined) { + throw invalidWrapper(wrapper, 'the wrapped integer must not carry a unit or display'); + } + const { transforms } = number; + const inner = integerTypeNode(number.format, { ...number, transforms: undefined }); + switch (wrapper.kind) { + case 'DateTime': + return dateTimeTypeNode(inner, { ticksPerSecond: wrapper.ticksPerSecond, transforms }); + case 'Duration': + return durationTypeNode(inner, { ticksPerSecond: wrapper.ticksPerSecond, transforms }); + case 'FixedPoint': + case 'SolAmount': { + if (number.format === 'shortU16') { + throw invalidWrapper(wrapper, 'a fixed point cannot wrap a variable-size `shortU16` integer'); + } + return wrapper.kind === 'SolAmount' + ? fixedPointTypeNode(inner, 9, { transforms, unit: 'SOL' }) + : fixedPointTypeNode(inner, wrapper.scale, { base: wrapper.base, transforms, unit: wrapper.unit }); + } + } +} + +function wrapFloat(number: FloatTypeNode, wrapper: NumberWrapper): TypeNode { + switch (wrapper.kind) { + case 'Unit': + return floatTypeNode(number.format, { ...number, unit: wrapper.unit }); + case 'UnitDisplay': + return floatTypeNode(number.format, { ...number, display: unitNumberDisplayNode({ unit: wrapper.unit }) }); + default: + return number; + } +} + +function invalidWrapper(wrapper: NumberWrapper, reason: string): CodamaError { + return new CodamaError(CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER, { kind: wrapper.kind, reason, wrapper }); +} diff --git a/packages/visitors/test/setInstructionAccountDefaultValuesVisitor.test.ts b/packages/visitors/test/setInstructionAccountDefaultValuesVisitor.test.ts new file mode 100644 index 000000000..06d6536ad --- /dev/null +++ b/packages/visitors/test/setInstructionAccountDefaultValuesVisitor.test.ts @@ -0,0 +1,167 @@ +import { + accountValueNode, + assertIsNode, + identityValueNode, + instructionAccountNode, + InstructionAccountNode, + instructionNode, + InstructionNode, + Node, + payerValueNode, + pdaNode, + pdaSeedValueNode, + pdaValueNode, + programNode, + publicKeyTypeNode, + publicKeyValueNode, + variablePdaSeedNode, +} from '@codama/nodes'; +import { visit } from '@codama/visitors-core'; +import { expect, test } from 'vitest'; + +import { + getCommonInstructionAccountDefaultRules, + InstructionAccountDefaultRule, + setInstructionAccountDefaultValuesVisitor, +} from '../src'; + +const account = (identifier: string, options: Partial = {}) => + instructionAccountNode({ identifier, isSigner: false, isWritable: false, ...options }); + +const programWith = (...instructions: InstructionNode[]) => + programNode({ + identifier: 'myProgram', + instructions, + pdas: [pdaNode({ identifier: 'vault', seeds: [variablePdaSeedNode('owner', publicKeyTypeNode())] })], + publicKey: '1111', + }); + +const getAccounts = (node: Node | null, index = 0) => { + assertIsNode(node, 'programNode'); + return node.instructions?.[index].accounts ?? []; +}; + +test('it sets the default values of matching accounts', () => { + // Given an instruction with a payer and a system program. + const node = programWith( + instructionNode({ accounts: [account('payer'), account('systemProgram')], identifier: 'create' }), + ); + + // When we apply the common rules. + const result = visit(node, setInstructionAccountDefaultValuesVisitor(getCommonInstructionAccountDefaultRules())); + + // Then both accounts get their default values. + expect(getAccounts(result).map(a => a.defaultValue)).toStrictEqual([ + payerValueNode(), + publicKeyValueNode('11111111111111111111111111111111', { identifier: 'splSystem' }), + ]); +}); + +test('the common rules match snake_case identifiers', () => { + // Given an instruction with snake_case accounts. + const node = programWith( + instructionNode({ + accounts: [account('fee_payer'), account('token_program'), account('sysvar_instructions_account')], + identifier: 'create', + }), + ); + + // When we apply the common rules, then every account gets a default value. + const result = visit(node, setInstructionAccountDefaultValuesVisitor(getCommonInstructionAccountDefaultRules())); + expect(getAccounts(result).map(a => a.defaultValue?.kind)).toStrictEqual([ + 'payerValueNode', + 'publicKeyValueNode', + 'publicKeyValueNode', + ]); +}); + +test('it matches string identifiers exactly', () => { + // Given an instruction with a snake_case account. + const node = programWith(instructionNode({ accounts: [account('my_program')], identifier: 'create' })); + const defaultValue = publicKeyValueNode('11111111111111111111111111111111'); + + // When a rule uses another casing, then nothing changes. + expect( + visit(node, setInstructionAccountDefaultValuesVisitor([{ account: 'myProgram', defaultValue }])), + ).toStrictEqual(node); + + // When a rule uses the exact identifier, then the default value is set. + const result = visit(node, setInstructionAccountDefaultValuesVisitor([{ account: 'my_program', defaultValue }])); + expect(getAccounts(result)[0].defaultValue).toStrictEqual(defaultValue); +}); + +test('it ignores optional or defaulted accounts when requested', () => { + // Given an optional account and an account with a default value. + const node = programWith( + instructionNode({ + accounts: [ + account('authority', { isOptional: true }), + account('payer', { defaultValue: identityValueNode() }), + ], + identifier: 'create', + }), + ); + + // When we apply the common rules, then neither account changes. + expect( + visit(node, setInstructionAccountDefaultValuesVisitor(getCommonInstructionAccountDefaultRules())), + ).toStrictEqual(node); +}); + +test('it gives precedence to rules restricted to an instruction without mutating the rules', () => { + // Given two instructions with an `authority` account. + const node = programWith( + instructionNode({ accounts: [account('authority')], identifier: 'create' }), + instructionNode({ accounts: [account('authority')], identifier: 'close' }), + ); + + // And a global rule declared before an instruction-specific one. + const rules: InstructionAccountDefaultRule[] = [ + { account: 'authority', defaultValue: identityValueNode() }, + { account: 'authority', defaultValue: payerValueNode(), instruction: 'close' }, + ]; + const rulesCopy = [...rules]; + + // When we apply them. + const result = visit(node, setInstructionAccountDefaultValuesVisitor(rules)); + + // Then the instruction-specific rule wins for its instruction only. + expect(getAccounts(result, 0)[0].defaultValue).toStrictEqual(identityValueNode()); + expect(getAccounts(result, 1)[0].defaultValue).toStrictEqual(payerValueNode()); + expect(rules).toStrictEqual(rulesCopy); +}); + +test('it fills PDA seeds and skips rules whose seeds cannot be filled', () => { + // Given an instruction with an owner and one without. + const node = programWith( + instructionNode({ accounts: [account('owner'), account('vault')], identifier: 'deposit' }), + instructionNode({ accounts: [account('vault')], identifier: 'close' }), + ); + + // When we default the vault account to its PDA. + const result = visit( + node, + setInstructionAccountDefaultValuesVisitor([{ account: 'vault', defaultValue: pdaValueNode('vault') }]), + ); + + // Then the seed is filled where possible and the rule is skipped otherwise. + expect(getAccounts(result, 0)[1].defaultValue).toStrictEqual( + pdaValueNode('vault', { seeds: [pdaSeedValueNode('owner', accountValueNode('owner'))] }), + ); + expect(getAccounts(result, 1)[0].defaultValue).toBeUndefined(); +}); + +test('it sets the default values of sub-instruction accounts', () => { + // Given an instruction with a sub-instruction. + const node = programWith( + instructionNode({ + identifier: 'parent', + subInstructions: [instructionNode({ accounts: [account('payer')], identifier: 'child' })], + }), + ); + + // When we apply the common rules, then the sub-instruction account gets a default value. + const result = visit(node, setInstructionAccountDefaultValuesVisitor(getCommonInstructionAccountDefaultRules())); + assertIsNode(result, 'programNode'); + expect(result.instructions?.[0].subInstructions?.[0].accounts?.[0].defaultValue).toStrictEqual(payerValueNode()); +}); diff --git a/packages/visitors/test/setInstructionDiscriminatorsVisitor.test.ts b/packages/visitors/test/setInstructionDiscriminatorsVisitor.test.ts new file mode 100644 index 000000000..c5c16341c --- /dev/null +++ b/packages/visitors/test/setInstructionDiscriminatorsVisitor.test.ts @@ -0,0 +1,200 @@ +import { + CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR, + CODAMA_ERROR__VISITORS__UNRECOGNIZED_UPDATE_KEYS, + CodamaError, +} from '@codama/errors'; +import { + assertIsNode, + constantDiscriminatorNode, + constantValueNode, + definedTypeLinkNode, + definedTypeNode, + fieldDiscriminatorNode, + hiddenPrefixTransformNode, + InstructionNode, + instructionNode, + integerTypeNode, + integerValueNode, + programNode, + sizeDiscriminatorNode, + stringTypeNode, + structFieldTypeNode, + structTypeNode, +} from '@codama/nodes'; +import { visit } from '@codama/visitors-core'; +import { expect, test } from 'vitest'; + +import { setInstructionDiscriminatorsVisitor } from '../src'; + +const u64Field = (identifier: string) => structFieldTypeNode({ identifier, type: integerTypeNode('u64') }); +const discriminatorField = (identifier = 'discriminator', value = '3') => + structFieldTypeNode({ + defaultValue: integerValueNode(value), + defaultValueStrategy: 'omitted', + identifier, + type: integerTypeNode('u8'), + }); + +test('it adds a discriminator field to inline struct data', () => { + // Given an instruction with struct data. + const node = instructionNode({ data: structTypeNode([u64Field('amount')]), identifier: 'transfer' }); + + // When we set its discriminator. + const result = visit(node, setInstructionDiscriminatorsVisitor({ transfer: { value: integerValueNode('3') } })); + + // Then a discriminator field is prepended and discriminates the instruction. + expect(result).toStrictEqual( + instructionNode({ + data: structTypeNode([discriminatorField(), u64Field('amount')]), + discriminators: [fieldDiscriminatorNode('discriminator')], + identifier: 'transfer', + }), + ); +}); + +test('it adds a discriminator field to instructions without data', () => { + // Given an instruction without data. + const node = instructionNode({ identifier: 'ping' }); + + // When we set its discriminator with a custom identifier. + const result = visit( + node, + setInstructionDiscriminatorsVisitor({ ping: { identifier: 'kind', value: integerValueNode('7') } }), + ); + + // Then the data becomes a struct with the discriminator field. + expect(result).toStrictEqual( + instructionNode({ + data: structTypeNode([discriminatorField('kind', '7')]), + discriminators: [fieldDiscriminatorNode('kind')], + identifier: 'ping', + }), + ); +}); + +test('it adds a hidden prefix to linked data', () => { + // Given an instruction whose data links to a shared defined type. + const node = programNode({ + definedTypes: [definedTypeNode({ identifier: 'args', type: structTypeNode([u64Field('amount')]) })], + identifier: 'myProgram', + instructions: [instructionNode({ data: definedTypeLinkNode('args'), identifier: 'transfer' })], + publicKey: '1111', + }); + + // When we set its discriminator. + const result = visit(node, setInstructionDiscriminatorsVisitor({ transfer: { value: integerValueNode('3') } })); + + // Then the data is prefixed with a hidden constant and the defined type is untouched. + const constant = constantValueNode(integerTypeNode('u8'), integerValueNode('3')); + assertIsNode(result, 'programNode'); + expect(result.definedTypes).toStrictEqual(node.definedTypes); + expect(result.instructions?.[0]).toStrictEqual( + instructionNode({ + data: definedTypeLinkNode('args', { transforms: [hiddenPrefixTransformNode([constant])] }), + discriminators: [constantDiscriminatorNode(constant, { offset: 0 })], + identifier: 'transfer', + }), + ); +}); + +test('it shifts existing discriminators by the size of the new one', () => { + // Given an instruction with field, constant and size discriminators. + const constant = constantValueNode(integerTypeNode('u16'), integerValueNode('1')); + const node = instructionNode({ + data: structTypeNode([u64Field('amount')]), + discriminators: [ + fieldDiscriminatorNode('amount', { offset: 0 }), + constantDiscriminatorNode(constant, { offset: 2 }), + sizeDiscriminatorNode(8), + ], + identifier: 'transfer', + }); + + // When we set a u32 discriminator. + const result = visit( + node, + setInstructionDiscriminatorsVisitor({ + transfer: { type: integerTypeNode('u32'), value: integerValueNode('3') }, + }), + ); + + // Then the existing discriminators are shifted by 4 bytes. + assertIsNode(result, 'instructionNode'); + expect(result.discriminators).toStrictEqual([ + fieldDiscriminatorNode('discriminator'), + fieldDiscriminatorNode('amount', { offset: 4 }), + constantDiscriminatorNode(constant, { offset: 6 }), + sizeDiscriminatorNode(12), + ]); +}); + +test('it throws when the discriminator cannot be set', () => { + const withField = instructionNode({ data: structTypeNode([u64Field('discriminator')]), identifier: 'ix' }); + const withLink = instructionNode({ data: definedTypeLinkNode('args'), identifier: 'ix' }); + const value = integerValueNode('3'); + const cannotSet = (instruction: InstructionNode, reason: string) => + new CodamaError(CODAMA_ERROR__VISITORS__CANNOT_SET_INSTRUCTION_DISCRIMINATOR, { + instruction, + instructionName: instruction.identifier, + reason, + }); + + // When the identifier already exists in the data. + expect(() => visit(withField, setInstructionDiscriminatorsVisitor({ ix: { value } }))).toThrow( + cannotSet(withField, 'the data already has a field named `discriminator`'), + ); + + // When the optional strategy is used with a hidden prefix. + expect(() => visit(withLink, setInstructionDiscriminatorsVisitor({ ix: { strategy: 'optional', value } }))).toThrow( + cannotSet( + withLink, + 'the `optional` strategy is not supported when the discriminator is added as a hidden prefix', + ), + ); + + // When the discriminator does not have a fixed size. + expect(() => + visit(withField, setInstructionDiscriminatorsVisitor({ ix: { type: stringTypeNode('utf8'), value } })), + ).toThrow(cannotSet(withField, 'the discriminator type must have a fixed size')); +}); + +test('it throws on unrecognized keys', () => { + // When we use the v1 `name` key, then we expect an error when creating the visitor. + expect(() => + setInstructionDiscriminatorsVisitor({ ix: { name: 'kind', value: integerValueNode('3') } as never }), + ).toThrow( + new CodamaError(CODAMA_ERROR__VISITORS__UNRECOGNIZED_UPDATE_KEYS, { + allowedKeys: ['docs', 'identifier', 'strategy', 'type', 'value'], + selector: 'ix', + unrecognizedKeys: ['name'], + }), + ); +}); + +test('it adds a hidden prefix to struct data that carries transforms', () => { + // Given struct data with a hidden prefix discriminated by a constant at offset 0. + const magic = constantValueNode(integerTypeNode('u8'), integerValueNode('9')); + const node = instructionNode({ + data: structTypeNode([u64Field('amount')], { transforms: [hiddenPrefixTransformNode([magic])] }), + discriminators: [constantDiscriminatorNode(magic, { offset: 0 })], + identifier: 'transfer', + }); + + // When we set its discriminator. + const result = visit(node, setInstructionDiscriminatorsVisitor({ transfer: { value: integerValueNode('3') } })); + + // Then a new outermost hidden prefix is added and the existing discriminator is shifted after it. + const constant = constantValueNode(integerTypeNode('u8'), integerValueNode('3')); + expect(result).toStrictEqual( + instructionNode({ + data: structTypeNode([u64Field('amount')], { + transforms: [hiddenPrefixTransformNode([magic]), hiddenPrefixTransformNode([constant])], + }), + discriminators: [ + constantDiscriminatorNode(constant, { offset: 0 }), + constantDiscriminatorNode(magic, { offset: 1 }), + ], + identifier: 'transfer', + }), + ); +}); diff --git a/packages/visitors/test/setNumberWrappersVisitor.test.ts b/packages/visitors/test/setNumberWrappersVisitor.test.ts new file mode 100644 index 000000000..97a236903 --- /dev/null +++ b/packages/visitors/test/setNumberWrappersVisitor.test.ts @@ -0,0 +1,171 @@ +import { CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER, CodamaError, isCodamaError } from '@codama/errors'; +import { + amountNumberDisplayNode, + assertIsNode, + constantValueNode, + dateTimeTypeNode, + durationTypeNode, + enumTypeNode, + fixedPointTypeNode, + fixedSizeTransformNode, + floatTypeNode, + hiddenPrefixTransformNode, + hiddenSuffixTransformNode, + injectedValueNode, + integerTypeNode, + integerValueNode, + optionTypeNode, + sizePrefixTransformNode, + stringTypeNode, + stringValueNode, + structFieldTypeNode, + structFieldValueNode, + structTypeNode, + structValueNode, + TypeNode, + unitNumberDisplayNode, +} from '@codama/nodes'; +import { visit } from '@codama/visitors-core'; +import { expect, test } from 'vitest'; + +import { NumberWrapper, setNumberWrappersVisitor } from '../src'; + +const struct = (fields: Record) => + structTypeNode(Object.entries(fields).map(([identifier, type]) => structFieldTypeNode({ identifier, type }))); +const wrapField = (type: TypeNode, wrapper: NumberWrapper) => { + const result = visit(struct({ value: type }), setNumberWrappersVisitor({ value: wrapper })); + assertIsNode(result, 'structTypeNode'); + return result.fields?.[0].type; +}; + +test('it wraps integers in fixed points, date-times and durations', () => { + const u64 = integerTypeNode('u64'); + const i64 = integerTypeNode('i64'); + expect(wrapField(u64, { kind: 'FixedPoint', scale: 6, unit: 'USDC' })).toStrictEqual( + fixedPointTypeNode(u64, 6, { unit: 'USDC' }), + ); + expect(wrapField(u64, { base: 2, kind: 'FixedPoint', scale: 32 })).toStrictEqual( + fixedPointTypeNode(u64, 32, { base: 2 }), + ); + expect(wrapField(u64, { kind: 'SolAmount' })).toStrictEqual(fixedPointTypeNode(u64, 9, { unit: 'SOL' })); + expect(wrapField(i64, { kind: 'DateTime' })).toStrictEqual(dateTimeTypeNode(i64)); + expect(wrapField(i64, { kind: 'Duration', ticksPerSecond: 1000 })).toStrictEqual( + durationTypeNode(i64, { ticksPerSecond: 1000 }), + ); +}); + +test('it sets units and displays on integers', () => { + const decimals = injectedValueNode({ key: 'decimals' }); + expect(wrapField(integerTypeNode('u32'), { kind: 'Unit', unit: 'bytes' })).toStrictEqual( + integerTypeNode('u32', { unit: 'bytes' }), + ); + expect(wrapField(integerTypeNode('u64'), { decimals, kind: 'AmountDisplay' })).toStrictEqual( + integerTypeNode('u64', { display: amountNumberDisplayNode({ decimals }) }), + ); + expect(wrapField(integerTypeNode('u16'), { kind: 'UnitDisplay', unit: stringValueNode('bps') })).toStrictEqual( + integerTypeNode('u16', { display: unitNumberDisplayNode({ unit: stringValueNode('bps') }) }), + ); +}); + +test('it sets units and unit displays on floats only', () => { + const f64 = floatTypeNode('f64'); + expect(wrapField(f64, { kind: 'Unit', unit: 'USD' })).toStrictEqual(floatTypeNode('f64', { unit: 'USD' })); + expect(wrapField(f64, { kind: 'UnitDisplay', unit: stringValueNode('USD') })).toStrictEqual( + floatTypeNode('f64', { display: unitNumberDisplayNode({ unit: stringValueNode('USD') }) }), + ); + expect(wrapField(f64, { kind: 'SolAmount' })).toStrictEqual(f64); +}); + +test('it moves the transforms of the integer onto the wrapper', () => { + const transforms = [fixedSizeTransformNode(16)]; + expect(wrapField(integerTypeNode('u64', { transforms }), { kind: 'SolAmount' })).toStrictEqual( + fixedPointTypeNode(integerTypeNode('u64'), 9, { transforms, unit: 'SOL' }), + ); +}); + +test('it does not wrap integers used as sizes or prefixes', () => { + // Given a field whose integers are an option prefix, a size prefix and an enum size. + const label = stringTypeNode('utf8', { transforms: [sizePrefixTransformNode(integerTypeNode('u32'))] }); + const kind = enumTypeNode([], { size: integerTypeNode('u16') }); + const node = struct({ + kind, + label, + value: optionTypeNode(integerTypeNode('u64'), { prefix: integerTypeNode('u8') }), + }); + + // When we wrap every number of the struct. + const result = visit(node, setNumberWrappersVisitor({ '[structTypeNode]': { kind: 'Unit', unit: 'x' } })); + + // Then only the value of the option is wrapped. + expect(result).toStrictEqual( + struct({ + kind, + label, + value: optionTypeNode(integerTypeNode('u64', { unit: 'x' }), { prefix: integerTypeNode('u8') }), + }), + ); +}); + +test('it does not wrap numbers that are already wrapped', () => { + // Given a field that is already a fixed point. + const node = struct({ value: fixedPointTypeNode(integerTypeNode('u64'), 6) }); + + // When we wrap it again, then nothing changes. + expect(visit(node, setNumberWrappersVisitor({ value: { kind: 'DateTime' } }))).toStrictEqual(node); +}); + +test('it throws on invalid wrappers', () => { + const expectInvalid = (fn: () => unknown) => { + let error: unknown; + try { + fn(); + } catch (e) { + error = e; + } + expect(isCodamaError(error, CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER)).toBe(true); + }; + + // An unknown kind or a zero scale throws when creating the visitor. + expectInvalid(() => setNumberWrappersVisitor({ value: { kind: 'Amount' } as never })); + expectInvalid(() => setNumberWrappersVisitor({ value: { kind: 'FixedPoint', scale: 0 } })); + + // A shortU16 fixed point or a wrapped integer with a unit throws when visiting. + expectInvalid(() => wrapField(integerTypeNode('shortU16'), { kind: 'SolAmount' })); + expectInvalid(() => wrapField(integerTypeNode('i64', { unit: 's' }), { kind: 'DateTime' })); +}); + +test('it reports the wrapper kind and reason', () => { + expect(() => setNumberWrappersVisitor({ value: { kind: 'FixedPoint', scale: 0 } })).toThrow( + new CodamaError(CODAMA_ERROR__VISITORS__INVALID_NUMBER_WRAPPER, { + kind: 'FixedPoint', + reason: 'a fixed point must have a non-zero scale; use a `Unit` wrapper instead', + wrapper: { kind: 'FixedPoint', scale: 0 }, + }), + ); +}); + +test('it does not wrap numbers within the types of constants', () => { + // Given a field with a hidden prefix constant and a constant whose type is a struct. + const prefix = constantValueNode(integerTypeNode('u64'), integerValueNode('1')); + const structConstant = constantValueNode( + struct({ inner: integerTypeNode('u64') }), + structValueNode([structFieldValueNode('inner', integerValueNode('2'))]), + ); + const node = struct({ + value: integerTypeNode('i64', { + transforms: [hiddenPrefixTransformNode([prefix]), hiddenSuffixTransformNode([structConstant])], + }), + }); + + // When we wrap every number of the struct. + const result = visit(node, setNumberWrappersVisitor({ '[structTypeNode]': { kind: 'DateTime' } })); + + // Then only the field itself is wrapped and the constants are untouched. + expect(result).toStrictEqual( + struct({ + value: dateTimeTypeNode(integerTypeNode('i64'), { + transforms: [hiddenPrefixTransformNode([prefix]), hiddenSuffixTransformNode([structConstant])], + }), + }), + ); +});