All files / api/src/ai-providers/together-ai embed.ts

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

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                                                                                                                                                                                                     
import { generateInvalidProviderResponseError } from '@api/utils/ai-provider';
import type {
  AIProviderFunctionConfig,
  ResponseTransformFunction,
} from '@shared/types/ai-providers/config';
import type {
  CreateEmbeddingsRequestBody,
  CreateEmbeddingsResponseBody,
} from '@shared/types/api/routes/embeddings-api';
import { AIProvider } from '@shared/types/constants';
import { togetherAIErrorResponseTransform } from './chat-complete';
 
export const togetherAIEmbedConfig: AIProviderFunctionConfig = {
  model: {
    param: 'model',
    required: true,
    default: 'mistral-embed',
  },
  input: {
    param: 'input',
    required: true,
    transform: (saRequestBody: CreateEmbeddingsRequestBody): string[] => {
      if ('input' in saRequestBody) {
        if (Array.isArray(saRequestBody.input)) {
          return saRequestBody.input as string[];
        }
        return [saRequestBody.input as string];
      }
      throw new Error('Invalid input for embedding');
    },
  },
  user: {
    param: 'user',
  },
  encoding_format: {
    param: 'encoding_format',
  },
  dimensions: {
    param: 'dimensions',
  },
};
 
export interface TogetherAIEmbedResponse {
  object: string;
  data: {
    object: string;
    embedding: number[];
    index: number;
  }[];
  model: string;
  usage: {
    prompt_tokens: number;
    total_tokens: number;
  };
}
 
export const togetherAIEmbedResponseTransform: ResponseTransformFunction = (
  aiProviderResponseBody,
  aiProviderResponseStatus,
  _responseHeaders,
  _strictOpenAiCompliance,
  saRequestData,
) => {
  if (aiProviderResponseStatus !== 200) {
    const errorResponse = togetherAIErrorResponseTransform(
      aiProviderResponseBody,
    );
    if (errorResponse) return errorResponse;
  }
 
  if ('data' in aiProviderResponseBody) {
    const response =
      aiProviderResponseBody as unknown as TogetherAIEmbedResponse;
    const _requestBody =
      saRequestData.requestBody as unknown as CreateEmbeddingsRequestBody;
 
    const responseBody: CreateEmbeddingsResponseBody = {
      object: response.object as 'list',
      data: response.data.map((d) => ({
        object: d.object as 'embedding',
        embedding: d.embedding,
        index: d.index,
      })),
      model: response.model,
      usage: {
        prompt_tokens: response.usage?.prompt_tokens || 0,
        total_tokens: response.usage?.total_tokens || 0,
      },
    };
 
    return responseBody;
  }
 
  return generateInvalidProviderResponseError(
    aiProviderResponseBody,
    AIProvider.TOGETHER_AI,
  );
};