All files / api/src/ai-providers/bedrock image-generate.ts

0% Statements 0/88
0% Branches 0/1
0% Functions 0/1
0% Lines 0/88

Press n or j to go to the next uncovered block, b, p or k for the previous block.

1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108                                                                                                                                                                                                                       
import { generateInvalidProviderResponseError } from '@api/utils/ai-provider';
import type {
  AIProviderFunctionConfig,
  ResponseTransformFunction,
} from '@shared/types/ai-providers/config';
import type { GenerateImageRequestBody } from '@shared/types/api/routes/images-api';
import { AIProvider } from '@shared/types/constants';
import { StabilityAIImageGenerateV2Config } from '../stability-ai/image-generate-v2';
import { bedrockErrorResponseTransform } from './chat-complete';
 
export const bedrockStabilityAIImageGenerateV1Config: AIProviderFunctionConfig =
  {
    prompt: {
      param: 'text_prompts',
      required: true,
      transform: (saRequestBody: GenerateImageRequestBody) => {
        return [
          {
            text: saRequestBody.prompt,
            weight: 1,
          },
        ];
      },
    },
    n: {
      param: 'samples',
      min: 1,
      max: 10,
    },
    size: [
      {
        param: 'height',
        transform: (saRequestBody: GenerateImageRequestBody): number =>
          parseInt(saRequestBody.size?.toLowerCase().split('x')[1] || '0', 10),
        min: 320,
      },
      {
        param: 'width',
        transform: (saRequestBody: GenerateImageRequestBody): number =>
          parseInt(saRequestBody.size?.toLowerCase().split('x')[0] || '0', 10),
        min: 320,
      },
    ],
    style: {
      param: 'style_preset',
    },
  };
 
interface ImageArtifact {
  base64: string;
  finishReason: 'CONTENT_FILTERED' | 'ERROR' | 'SUCCESS';
  seed: number;
}
 
export const bedrockStabilityAIImageGenerateV1ResponseTransform: ResponseTransformFunction =
  (aiProviderResponseBody, aiProviderResponseStatus) => {
    if (aiProviderResponseStatus !== 200) {
      const errorResponse = bedrockErrorResponseTransform(
        aiProviderResponseBody,
      );
      if (errorResponse) return errorResponse;
    }
 
    if ('artifacts' in aiProviderResponseBody) {
      const artifacts =
        aiProviderResponseBody.artifacts as unknown as ImageArtifact[];
      return {
        created: Math.floor(Date.now() / 1000),
        data: artifacts.map((art) => ({ b64_json: art.base64 })),
        provider: AIProvider.BEDROCK,
      };
    }
 
    return generateInvalidProviderResponseError(
      aiProviderResponseBody,
      AIProvider.BEDROCK,
    );
  };
 
export const bedrockStabilityAIImageGenerateV2Config =
  StabilityAIImageGenerateV2Config;
 
export const bedrockStabilityAIImageGenerateV2ResponseTransform: ResponseTransformFunction =
  (aiProviderResponseBody, aiProviderResponseStatus) => {
    if (aiProviderResponseStatus !== 200) {
      const errorResponse = bedrockErrorResponseTransform(
        aiProviderResponseBody,
      );
      if (errorResponse) return errorResponse;
    }
 
    if ('images' in aiProviderResponseBody) {
      const images = aiProviderResponseBody.images as unknown as string[];
      return {
        created: Math.floor(Date.now() / 1000),
        data: images.map((image) => ({
          b64_json: image,
        })),
        provider: AIProvider.BEDROCK,
      };
    }
 
    return generateInvalidProviderResponseError(
      aiProviderResponseBody,
      AIProvider.BEDROCK,
    );
  };