bedrock.ts 3.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119
  1. import { createAmazonBedrock } from '@ai-sdk/amazon-bedrock'
  2. import { createCredentialChain, fromNodeProviderChain } from '@aws-sdk/credential-providers'
  3. import { CredentialsProviderError } from '@smithy/property-provider'
  4. import { awsCredentialsProvider } from '@vercel/functions/oidc'
  5. import { LanguageModel } from 'ai'
  6. import { BedrockModel } from './model.utils'
  7. import { selectWeightedKey } from './util'
  8. const credentialProvider = createCredentialChain(
  9. // Vercel OIDC provider will be used for staging/production
  10. vercelOidcProvider,
  11. // AWS profile will be used for local development
  12. fromNodeProviderChain({
  13. profile: process.env.AWS_BEDROCK_PROFILE,
  14. })
  15. )
  16. /**
  17. * Creates a Vercel OIDC provider for AWS credentials.
  18. *
  19. * Wraps `awsCredentialsProvider` to properly handle errors
  20. * so that it can be used in a credential chain.
  21. */
  22. async function vercelOidcProvider() {
  23. try {
  24. return await awsCredentialsProvider({
  25. roleArn: process.env.AWS_BEDROCK_ROLE_ARN!,
  26. })()
  27. } catch (error) {
  28. const message = error instanceof Error ? error.message : 'Failed to create Vercel OIDC provider'
  29. // Re-throw using the correct error type and `tryNextLink` option
  30. throw new CredentialsProviderError(message, {
  31. tryNextLink: true,
  32. })
  33. }
  34. }
  35. export async function checkAwsCredentials() {
  36. try {
  37. const credentials = await credentialProvider()
  38. return !!credentials
  39. } catch (error) {
  40. return false
  41. }
  42. }
  43. export const bedrockRegionMap = {
  44. use1: 'us-east-1',
  45. use2: 'us-east-2',
  46. usw2: 'us-west-2',
  47. euc1: 'eu-central-1',
  48. } as const
  49. export type BedrockRegion = keyof typeof bedrockRegionMap
  50. export const regionPrefixMap: Record<BedrockRegion, string> = {
  51. use1: 'us',
  52. use2: 'us',
  53. usw2: 'us',
  54. euc1: 'eu',
  55. }
  56. export type RegionWeights = Record<BedrockRegion, number>
  57. /**
  58. * Weights for distributing requests across Bedrock regions.
  59. * Weights are proportional to our rate limits per model per region.
  60. */
  61. const modelRegionWeights: Record<BedrockModel, RegionWeights> = {
  62. ['anthropic.claude-3-7-sonnet-20250219-v1:0']: {
  63. use1: 40,
  64. use2: 10,
  65. usw2: 10,
  66. euc1: 10,
  67. },
  68. ['openai.gpt-oss-120b-1:0']: {
  69. use1: 0,
  70. use2: 0,
  71. usw2: 30,
  72. euc1: 0,
  73. },
  74. }
  75. /**
  76. * Creates a Bedrock client that routes requests to different regions
  77. * based on a routing key, with optional OpenAI support.
  78. *
  79. * Used to load balance requests across multiple regions depending on
  80. * their capacities.
  81. */
  82. export function createRoutedBedrock(routingKey?: string) {
  83. return async (modelId: BedrockModel): Promise<LanguageModel> => {
  84. const regionWeights = modelRegionWeights[modelId]
  85. // Select the Bedrock region based on the routing key and the model
  86. const bedrockRegion = routingKey
  87. ? await selectWeightedKey(routingKey, regionWeights)
  88. : // There's a few places where getModel is called without a routing key
  89. // Will cause disproportionate load on use1 region
  90. regionWeights['use1'] > 0
  91. ? 'use1'
  92. : 'usw2'
  93. const bedrock = createAmazonBedrock({
  94. credentialProvider,
  95. region: bedrockRegionMap[bedrockRegion],
  96. })
  97. // Cross-region models require the region prefix
  98. const activeRegions = Object.values(regionWeights).filter((weight) => weight > 0).length
  99. const modelName = activeRegions > 1 ? `${regionPrefixMap[bedrockRegion]}.${modelId}` : modelId
  100. const model = bedrock(modelName)
  101. return model
  102. }
  103. }