All files / api/src/ai-providers/palm complete.ts

52.74% Statements 48/91
100% Branches 0/0
0% Functions 0/2
52.74% Lines 48/91

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 1101x 1x                   1x       1x 1x 1x 1x 1x 1x 1x 1x 1x 1x             1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x 1x   1x                                                                                      
import { googleErrorResponseTransform } from '@api/ai-providers/google/chat-complete';
import { generateInvalidProviderResponseError } from '@api/utils/ai-provider';
import type {
  AIProviderFunctionConfig,
  ResponseTransformFunction,
} from '@shared/types/ai-providers/config';
import type {
  CompletionFinishReason,
  CompletionRequestBody,
  CompletionResponseBody,
} from '@shared/types/api/routes/completions-api';
import { AIProvider } from '@shared/types/constants';
 
// TODOS: this configuration does not enforce the maximum token limit for the input parameter. If you want to enforce this, you might need to add a custom validation function or a max property to the ParameterConfig interface, and then use it in the input configuration. However, this might be complex because the token count is not a simple length check, but depends on the specific tokenization method used by the model.
 
export const palmCompleteConfig: AIProviderFunctionConfig = {
  model: {
    param: 'model',
    required: true,
    default: 'model/text-bison-001',
  },
  prompt: {
    param: 'prompt',
    default: '',
    transform: (saRequestBody: CompletionRequestBody) => {
      const { prompt: text } = saRequestBody;
      const prompt = {
        text,
      };
      return prompt;
    },
  },
  temperature: {
    param: 'temperature',
    default: 1,
    min: 0,
    max: 1,
  },
  top_p: {
    param: 'topP',
    default: 1,
    min: 0,
    max: 1,
  },
  top_k: {
    param: 'topK',
    default: 1,
    min: 0,
    max: 1,
  },
  n: {
    param: 'candidateCount',
    default: 1,
    min: 1,
    max: 8,
  },
  max_tokens: {
    param: 'maxOutputTokens',
    default: 100,
    min: 1,
  },
  stop: {
    param: 'stopSequences',
  },
};
 
export const palmCompleteResponseTransform: ResponseTransformFunction = (
  aiProviderResponseBody,
  aiProviderResponseStatus,
) => {
  if (aiProviderResponseStatus !== 200) {
    const errorResponse = googleErrorResponseTransform(
      aiProviderResponseBody,
      AIProvider.PALM,
    );
    if (errorResponse) return errorResponse;
  }
 
  if ('candidates' in aiProviderResponseBody) {
    const candidates = aiProviderResponseBody.candidates as {
      output: string;
    }[];
 
    const responseBody: CompletionResponseBody = {
      id: Date.now().toString(),
      object: 'text_completion',
      created: Math.floor(Date.now() / 1000),
      model: 'Unknown',
      choices:
        candidates.map((generation: { output: string }, index: number) => ({
          text: generation.output,
          index: index,
          logprobs: {
            top_logprobs: [],
            tokens: [],
            token_logprobs: [],
            text_offset: [],
          },
          finish_reason: 'length' as CompletionFinishReason,
        })) ?? [],
    };
    return responseBody;
  }
 
  return generateInvalidProviderResponseError(
    aiProviderResponseBody,
    AIProvider.PALM,
  );
};