All files / api/src/ai-providers/bedrock get-batch-output.ts

0% Statements 0/137
100% Branches 1/1
100% Functions 1/1
0% Lines 0/137

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 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178                                                                                                                                                                                                                                                                                                                                                                   
import type { BedrockGetBatchResponse } from '@api/ai-providers/bedrock/types';
import { getOctetStreamToOctetStreamTransformer } from '@api/handlers/stream-handler-utils';
import type { AppContext } from '@api/types/hono';
import type { SuperAgentsRequestData } from '@shared/types/api/request';
import type { SuperAgentsTarget } from '@shared/types/api/request/headers';
import { AIProvider } from '@shared/types/constants';
import bedrockAPIConfig from './api';
import { BedrockUploadFileResponseTransforms } from './upload-file-utils';
 
// Define a more specific type for the transform functions
type BedrockResponseTransformFunction = (modelOutput: {
  id: string;
}) => Record<string, unknown>;
 
const getModelProvider = (modelId: string): string => {
  let provider = '';
  if (modelId.includes('llama2')) provider = 'llama2';
  else if (modelId.includes('llama3')) provider = 'llama3';
  else if (modelId.includes('titan')) provider = 'titan';
  else if (modelId.includes('mistral')) provider = 'mistral';
  else if (modelId.includes('anthropic')) provider = 'anthropic';
  else if (modelId.includes('ai21')) provider = 'ai21';
  else if (modelId.includes('cohere')) provider = 'cohere';
  else throw new Error('Invalid model slug');
  return provider;
};
 
const getRowTransform = (
  modelId: string,
): ((row: Record<string, unknown>) => Record<string, unknown>) => {
  const provider = getModelProvider(modelId);
  return (row: Record<string, unknown>): Record<string, unknown> => {
    if (!row.modelOutput && row.error) {
      // Convert Error to Record<string, unknown> format
      if (row.error instanceof Error) {
        return {
          error: row.error.message,
          status: 'error',
        };
      }
      return row.error as Record<string, unknown>;
    }
 
    // Cast modelOutput to a type with the required properties
    const modelOutput = row.modelOutput as { id: string };
 
    // Cast the transform function to the correct type
    const transformFunction = BedrockUploadFileResponseTransforms[
      provider
    ] as BedrockResponseTransformFunction;
    const transformedResponse = transformFunction(modelOutput);
    transformedResponse.model = modelId;
 
    return {
      id: modelOutput.id,
      custom_id: row.recordId as string,
      response: {
        status_code: 200,
        request_id: modelOutput.id,
        body: transformedResponse,
      },
      error: null,
    };
  };
};
 
export const bedrockGetBatchOutputRequestHandler = async ({
  c,
  saTarget,
  saRequestData,
}: {
  c: AppContext;
  saTarget: SuperAgentsTarget;
  saRequestData: SuperAgentsRequestData;
}): Promise<Response> => {
  try {
    // get s3 file id from batch details
    // get file from s3
    const baseUrl = bedrockAPIConfig.getBaseURL({
      c,
      saTarget,
      saRequestData,
    });
    const batchId = saRequestData.url
      .split('/v1/batches/')[1]
      .replace('/output', '');
    const retrieveBatchURL = `${baseUrl}/model-invocation-job/${batchId}`;
    const retrieveBatchesHeaders = await bedrockAPIConfig.headers({
      c,
      saTarget,
      saRequestData,
    });
    const retrieveBatchesResponse = await fetch(retrieveBatchURL, {
      method: 'GET',
      headers: retrieveBatchesHeaders as HeadersInit,
    });
 
    const batchDetails: BedrockGetBatchResponse =
      await retrieveBatchesResponse.json();
    const outputFileId = batchDetails.outputDataConfig.s3OutputDataConfig.s3Uri;
 
    const { aws_region } = saTarget;
    const awsS3Bucket = outputFileId.replace('s3://', '').split('/')[0];
    const jobId = batchDetails.jobArn.split('/')[1];
    const inputS3URIParts =
      batchDetails.inputDataConfig.s3InputDataConfig.s3Uri.split('/');
 
    const primaryKey = outputFileId?.replace(`s3://${awsS3Bucket}/`, '') ?? '';
 
    const awsS3ObjectKey = `${primaryKey}${jobId}/${inputS3URIParts[inputS3URIParts.length - 1]}.out`;
    const awsModelProvider = batchDetails.modelId;
 
    const s3FileURL = `https://${awsS3Bucket}.s3.${aws_region}.amazonaws.com/${awsS3ObjectKey}`;
    const s3FileHeaders = await bedrockAPIConfig.headers({
      c,
      saTarget,
      saRequestData,
    });
    const s3FileResponse = await fetch(s3FileURL, {
      method: 'GET',
      headers: s3FileHeaders as HeadersInit,
    });
    let responseStream: ReadableStream;
    if (
      s3FileResponse.headers.get('content-type')?.includes('octet-stream') &&
      s3FileResponse?.body
    ) {
      responseStream = s3FileResponse?.body?.pipeThrough(
        getOctetStreamToOctetStreamTransformer(
          getRowTransform(awsModelProvider),
        ),
      );
      return new Response(responseStream, {
        headers: {
          'content-type': 'application/octet-stream',
        },
      });
    } else {
      const body = await s3FileResponse.text();
      throw new Error(body);
    }
  } catch (error: unknown) {
    let errorResponse: Record<string, unknown> & { provider?: string };
 
    try {
      errorResponse = JSON.parse((error as Error).message);
      errorResponse.provider = AIProvider.BEDROCK;
    } catch (_e) {
      errorResponse = {
        error: {
          message: (error as Error).message,
          type: null,
          param: null,
          code: 500,
        },
        provider: AIProvider.BEDROCK,
      };
    }
    return new Response(JSON.stringify(errorResponse), {
      status: 500,
      headers: {
        'Content-Type': 'application/json',
      },
    });
  }
};
 
// export const bedrockGetBatchOutputResponseTransform: ResponseTransformFunction =
//   (aiProviderResponseBody, aiProviderResponseStatus) => {
//     if (aiProviderResponseStatus !== 200) {
//       const errorResponse = bedrockErrorResponseTransform(
//         aiProviderResponseBody,
//       );
//       if (errorResponse) return errorResponse;
//     }
//     return aiProviderResponseBody as unknown as GetBatchOutputResponseBody;
//   }; // TODO: Add this back in when we have a way to handle the response