feat: 1) Present 'maxtokens' as properties tied to a single model. 2) Remove the original author's implementation of the send verification logic and replace it with a user input validator. Pre-verification 3) Provides the ability to pull the 'User Visible modellist' provided by 'provider' 4) Provider-related parameters are passed in the constructor of 'providerClient'. Not passed in the 'chat' method

This commit is contained in:
Dean-YZG
2024-05-17 21:11:21 +08:00
parent 74a6e1260e
commit 8093d1ffba
30 changed files with 883 additions and 581 deletions

View File

@@ -0,0 +1,5 @@
export * from "./types";
export * from "./locale";
export * from "./utils";

View File

@@ -1,5 +1,7 @@
import { RequestMessage } from "../api"; import { RequestMessage } from "../api";
export { type RequestMessage };
// ===================================== LLM Types start ====================================== // ===================================== LLM Types start ======================================
export interface ModelConfig { export interface ModelConfig {
@@ -10,35 +12,50 @@ export interface ModelConfig {
max_tokens: number; max_tokens: number;
} }
export type Model = { export interface ModelSettings extends Omit<ModelConfig, "max_tokens"> {
global_max_tokens: number;
}
export type ModelTemplate = {
name: string; // id of model in a provider name: string; // id of model in a provider
displayName: string; displayName: string;
isVisionModel?: boolean; isVisionModel?: boolean;
isDefaultActive: boolean; // model is initialized to be active isDefaultActive: boolean; // model is initialized to be active
isDefaultSelected?: boolean; // model is initialized to be as default used model isDefaultSelected?: boolean; // model is initialized to be as default used model
providerTemplateName: string; max_tokens?: number;
}; };
export interface Model extends Omit<ModelTemplate, "isDefaultActive"> {
providerTemplateName: string;
isActive: boolean;
providerName: string;
available: boolean;
customized: boolean; // Only customized model is allowed to be modified
}
export interface ModelInfo extends Pick<ModelTemplate, "name"> {
[k: string]: any;
}
// ===================================== LLM Types end ====================================== // ===================================== LLM Types end ======================================
// ===================================== Chat Request Types start ====================================== // ===================================== Chat Request Types start ======================================
export interface ChatRequestPayload<SettingKeys extends string = ""> { export interface ChatRequestPayload {
messages: RequestMessage[]; messages: RequestMessage[];
providerConfig: Record<SettingKeys, string>;
context: { context: {
isApp: boolean; isApp: boolean;
}; };
} }
export interface StandChatRequestPayload<SettingKeys extends string = ""> export interface StandChatRequestPayload extends ChatRequestPayload {
extends ChatRequestPayload<SettingKeys> {
modelConfig: ModelConfig; modelConfig: ModelConfig;
model: string; model: string;
} }
export interface InternalChatRequestPayload<SettingKeys extends string = ""> export interface InternalChatRequestPayload<SettingKeys extends string = "">
extends StandChatRequestPayload<SettingKeys> { extends StandChatRequestPayload {
providerConfig: Partial<Record<SettingKeys, string>>;
isVisionModel: Model["isVisionModel"]; isVisionModel: Model["isVisionModel"];
stream: boolean; stream: boolean;
} }
@@ -50,12 +67,18 @@ export interface ProviderRequestPayload {
method: string; method: string;
} }
export interface ChatHandlers { export interface InternalChatHandlers {
onProgress: (message: string, chunk: string) => void; onProgress: (message: string, chunk: string) => void;
onFinish: (message: string) => void; onFinish: (message: string) => void;
onError: (err: Error) => void; onError: (err: Error) => void;
} }
export interface ChatHandlers extends InternalChatHandlers {
onProgress: (chunk: string) => void;
onFinish: () => void;
onFlash: (message: string) => void;
}
// ===================================== Chat Request Types end ====================================== // ===================================== Chat Request Types end ======================================
// ===================================== Chat Response Types start ====================================== // ===================================== Chat Response Types start ======================================
@@ -75,7 +98,8 @@ export type Validator =
| "number" | "number"
| "string" | "string"
| NumberRange | NumberRange
| NumberRange[]; | NumberRange[]
| ((v: any) => Promise<string | void>);
export type CommonSettingItem<SettingKeys extends string> = { export type CommonSettingItem<SettingKeys extends string> = {
name: SettingKeys; name: SettingKeys;
@@ -141,22 +165,20 @@ export interface IProviderTemplate<
displayName: string; displayName: string;
settingItems: SettingItem<SettingKeys>[]; settingItems: SettingItem<SettingKeys>[];
}; };
readonly models: Model[]; readonly defaultModels: ModelTemplate[];
// formatChatPayload(payload: InternalChatRequestPayload<SettingKeys>): ProviderRequestPayload;
// readWholeMessageResponseBody(res: WholeMessageResponseBody): StandChatReponseMessage;
streamChat( streamChat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
onProgress?: (message: string, chunk: string) => void, handlers: ChatHandlers,
onFinish?: (message: string) => void,
onError?: (err: Error) => void,
): AbortController; ): AbortController;
chat( chat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
): Promise<StandChatReponseMessage>; ): Promise<StandChatReponseMessage>;
getAvailableModels?(
providerConfig: InternalChatRequestPayload<SettingKeys>["providerConfig"],
): Promise<ModelInfo[]>;
} }
export interface Serializable<Snapshot> { export interface Serializable<Snapshot> {

View File

@@ -0,0 +1,26 @@
import { RequestMessage } from "./types";
export function getMessageTextContent(message: RequestMessage) {
if (typeof message.content === "string") {
return message.content;
}
for (const c of message.content) {
if (c.type === "text") {
return c.text ?? "";
}
}
return "";
}
export function getMessageImages(message: RequestMessage): string[] {
if (typeof message.content === "string") {
return [];
}
const urls: string[] = [];
for (const c of message.content) {
if (c.type === "image_url") {
urls.push(c.image_url?.url ?? "");
}
}
return urls;
}

View File

@@ -1,9 +1,9 @@
export * from "./types"; export * from "../common/types";
export * from "./providerClient"; export * from "./providerClient";
export * from "./modelClient"; export * from "./modelClient";
export * from "./locale"; export * from "../common/locale";
export * from "./shim"; export * from "./shim";

View File

@@ -1,23 +1,28 @@
import { ChatRequestPayload, Model, ModelConfig, ChatHandlers } from "./types"; import {
import { ProviderClient, ProviderTemplateName } from "./providerClient"; ChatRequestPayload,
Model,
ModelSettings,
InternalChatHandlers,
} from "../common";
import { Provider, ProviderClient } from "./providerClient";
export class ModelClient { export class ModelClient {
static getAllProvidersDefaultModels = () => {
return ProviderClient.getAllProvidersDefaultModels();
};
constructor( constructor(
private model: Model, private model: Model,
private modelConfig: ModelConfig, private modelSettings: ModelSettings,
private providerClient: ProviderClient, private providerClient: ProviderClient,
) {} ) {}
chat(payload: ChatRequestPayload, handlers: ChatHandlers) { chat(payload: ChatRequestPayload, handlers: InternalChatHandlers) {
try { try {
return this.providerClient.streamChat( return this.providerClient.streamChat(
{ {
...payload, ...payload,
modelConfig: this.modelConfig, modelConfig: {
...this.modelSettings,
max_tokens:
this.model.max_tokens ?? this.modelSettings.global_max_tokens,
},
model: this.model.name, model: this.model.name,
}, },
handlers, handlers,
@@ -31,7 +36,11 @@ export class ModelClient {
try { try {
return this.providerClient.chat({ return this.providerClient.chat({
...payload, ...payload,
modelConfig: this.modelConfig, modelConfig: {
...this.modelSettings,
max_tokens:
this.model.max_tokens ?? this.modelSettings.global_max_tokens,
},
model: this.model.name, model: this.model.name,
}); });
} catch (e) { } catch (e) {
@@ -40,7 +49,50 @@ export class ModelClient {
} }
} }
export function ModelClientFactory(model: Model, modelConfig: ModelConfig) { // must generate new ModelClient during every chat
const providerClient = new ProviderClient(model.providerTemplateName); export function ModelClientFactory(
return new ModelClient(model, modelConfig, providerClient); model: Model,
provider: Provider,
modelSettings: ModelSettings,
) {
const providerClient = new ProviderClient(provider);
return new ModelClient(model, modelSettings, providerClient);
}
export function getFiltertModels(
models: readonly Model[],
customModels: string,
) {
const modelTable: Record<string, Model> = {};
// default models
models.forEach((m) => {
modelTable[m.name] = m;
});
// server custom models
customModels
.split(",")
.filter((v) => !!v && v.length > 0)
.forEach((m) => {
const available = !m.startsWith("-");
const nameConfig =
m.startsWith("+") || m.startsWith("-") ? m.slice(1) : m;
const [name, displayName] = nameConfig.split("=");
// enable or disable all models
if (name === "all") {
Object.values(modelTable).forEach(
(model) => (model.available = available),
);
} else {
modelTable[name] = {
...modelTable[name],
displayName,
available,
};
}
});
return modelTable;
} }

View File

@@ -1,118 +1,182 @@
import { import {
ChatHandlers,
IProviderTemplate, IProviderTemplate,
InternalChatHandlers,
Model, Model,
ModelTemplate,
StandChatReponseMessage, StandChatReponseMessage,
StandChatRequestPayload, StandChatRequestPayload,
} from "./types"; } from "../common";
import * as ProviderTemplates from "@/app/client/providers"; import * as ProviderTemplates from "@/app/client/providers";
import { cloneDeep } from "lodash-es"; import { nanoid } from "nanoid";
export type ProviderTemplate = export type ProviderTemplate = IProviderTemplate<any, any, any>;
(typeof ProviderTemplates)[keyof typeof ProviderTemplates];
export type ProviderTemplateName = export type ProviderTemplateName =
(typeof ProviderTemplates)[keyof typeof ProviderTemplates]["prototype"]["name"]; (typeof ProviderTemplates)[keyof typeof ProviderTemplates]["prototype"]["name"];
export interface Provider<
Providerconfig extends Record<string, any> = Record<string, any>,
> {
name: string; // id of provider
isActive: boolean;
providerTemplateName: ProviderTemplateName;
providerConfig: Providerconfig;
isDefault: boolean; // Not allow to modify models of default provider
updated: boolean; // provider initial is finished
displayName: string;
models: Model[];
}
const providerTemplates = Object.values(ProviderTemplates).reduce(
(r, t) => ({
...r,
[t.prototype.name]: new t(),
}),
{} as Record<ProviderTemplateName, ProviderTemplate>,
);
export class ProviderClient { export class ProviderClient {
provider: IProviderTemplate<any, any, any>; providerTemplate: IProviderTemplate<any, any, any>;
static ProviderTemplates = ProviderTemplates; static ProviderTemplates = providerTemplates;
static getAllProvidersDefaultModels = () => {
return Object.values(ProviderClient.ProviderTemplates).reduce(
(r, p) => ({
...r,
[p.prototype.name]: cloneDeep(p.prototype.models),
}),
{} as Record<ProviderTemplateName, Model[]>,
);
};
static getAllProviderTemplates = () => { static getAllProviderTemplates = () => {
return Object.values(ProviderClient.ProviderTemplates).reduce( return Object.values(providerTemplates).reduce(
(r, p) => ({ (r, t) => ({
...r, ...r,
[p.prototype.name]: p, [t.name]: t,
}), }),
{} as Record<ProviderTemplateName, ProviderTemplate>, {} as Record<ProviderTemplateName, ProviderTemplate>,
); );
}; };
static getProviderTemplateList = () => { static getProviderTemplateMetaList = () => {
return Object.values(ProviderClient.ProviderTemplates); return Object.values(providerTemplates).map((t) => ({
...t.providerMeta,
name: t.name,
}));
}; };
constructor(providerTemplateName: string) { constructor(private provider: Provider) {
this.provider = this.getProviderTemplate(providerTemplateName); const { providerTemplateName } = provider;
} this.providerTemplate = this.getProviderTemplate(providerTemplateName);
get settingItems() {
const { providerMeta } = this.provider;
const { settingItems } = providerMeta;
return settingItems;
} }
private getProviderTemplate(providerTemplateName: string) { private getProviderTemplate(providerTemplateName: string) {
const providerTemplate = const providerTemplate = Object.values(providerTemplates).find(
Object.values(ProviderTemplates).find( (template) => template.name === providerTemplateName,
(template) => template.prototype.name === providerTemplateName, );
) || ProviderTemplates.NextChatProvider;
return new providerTemplate(); return providerTemplate || providerTemplates.openai;
} }
getModelConfig(modelName: string) { private getModelConfig(modelName: string) {
const { models } = this.provider; const { models } = this.provider;
return ( return (
models.find((config) => config.name === modelName) || models.find((m) => m.name === modelName) ||
models.find((config) => config.isDefaultSelected) models.find((m) => m.isDefaultSelected)
); );
} }
getAvailableModels() {
return Promise.resolve(
this.providerTemplate.getAvailableModels?.(this.provider.providerConfig),
)
.then((res) => {
const { defaultModels } = this.providerTemplate;
const availableModelsSet = new Set(
(res ?? defaultModels).map((o) => o.name),
);
return defaultModels.filter((m) => availableModelsSet.has(m.name));
})
.catch(() => {
return this.providerTemplate.defaultModels;
});
}
async chat( async chat(
payload: StandChatRequestPayload<string>, payload: StandChatRequestPayload,
): Promise<StandChatReponseMessage> { ): Promise<StandChatReponseMessage> {
return this.provider.chat({ return this.providerTemplate.chat({
...payload, ...payload,
stream: false, stream: false,
isVisionModel: this.getModelConfig(payload.model)?.isVisionModel, isVisionModel: this.getModelConfig(payload.model)?.isVisionModel,
providerConfig: this.provider.providerConfig,
}); });
} }
streamChat(payload: StandChatRequestPayload<string>, handlers: ChatHandlers) { streamChat(payload: StandChatRequestPayload, handlers: InternalChatHandlers) {
return this.provider.streamChat( let responseText = "";
let remainText = "";
const timer = this.providerTemplate.streamChat(
{ {
...payload, ...payload,
stream: true, stream: true,
isVisionModel: this.getModelConfig(payload.model)?.isVisionModel, isVisionModel: this.getModelConfig(payload.model)?.isVisionModel,
providerConfig: this.provider.providerConfig,
},
{
onProgress: (chunk) => {
remainText += chunk;
},
onError: (err) => {
handlers.onError(err);
},
onFinish: () => {},
onFlash: (message: string) => {
handlers.onFinish(message);
},
}, },
handlers.onProgress,
handlers.onFinish,
handlers.onError,
); );
timer.signal.onabort = () => {
const message = responseText + remainText;
remainText = "";
handlers.onFinish(message);
};
const animateResponseText = () => {
if (remainText.length > 0) {
const fetchCount = Math.max(1, Math.round(remainText.length / 60));
const fetchText = remainText.slice(0, fetchCount);
responseText += fetchText;
remainText = remainText.slice(fetchCount);
handlers.onProgress(responseText, fetchText);
}
requestAnimationFrame(animateResponseText);
};
// start animaion
animateResponseText();
return timer;
} }
} }
export interface Provider { type Params = Omit<Provider, "providerTemplateName" | "name" | "isDefault">;
name: string; // id of provider
displayName: string;
isActive: boolean;
providerTemplateName: ProviderTemplateName;
models: Model[];
}
function createProvider( function createProvider(
provider: ProviderTemplateName, provider: ProviderTemplateName,
params?: Omit<Provider, "providerTemplateName">, isDefault: true,
): Provider;
function createProvider(provider: ProviderTemplate, isDefault: true): Provider;
function createProvider(
provider: ProviderTemplateName,
isDefault: false,
params: Params,
): Provider; ): Provider;
function createProvider( function createProvider(
provider: ProviderTemplate, provider: ProviderTemplate,
params?: Omit<Provider, "providerTemplateName">, isDefault: false,
params: Params,
): Provider; ): Provider;
function createProvider( function createProvider(
provider: ProviderTemplate | ProviderTemplateName, provider: ProviderTemplate | ProviderTemplateName,
params?: Omit<Provider, "providerTemplateName">, isDefault: boolean,
params?: Params,
): Provider { ): Provider {
let providerTemplate: ProviderTemplate; let providerTemplate: ProviderTemplate;
if (typeof provider === "string") { if (typeof provider === "string") {
@@ -120,17 +184,41 @@ function createProvider(
} else { } else {
providerTemplate = provider; providerTemplate = provider;
} }
const name = `${providerTemplate.name}__${nanoid()}`;
const { const {
name = providerTemplate.prototype.name, displayName = providerTemplate.providerMeta.displayName,
displayName = providerTemplate.prototype.providerMeta.displayName, models = providerTemplate.defaultModels.map((m) =>
models = providerTemplate.prototype.models, createModelFromModelTemplate(m, providerTemplate, name),
),
providerConfig,
} = params ?? {}; } = params ?? {};
return { return {
name, name,
displayName, displayName,
isActive: true, isActive: true,
models, models,
providerTemplateName: providerTemplate.prototype.name, providerTemplateName: providerTemplate.name,
providerConfig: isDefault ? {} : providerConfig!,
isDefault,
updated: true,
};
}
function createModelFromModelTemplate(
m: ModelTemplate,
p: ProviderTemplate,
providerName: string,
) {
return {
...m,
providerTemplateName: p.name,
providerName,
isActive: m.isDefaultActive,
available: true,
customized: false,
}; };
} }

View File

@@ -1,4 +1,4 @@
import { SettingItem } from "../../core/types"; import { SettingItem } from "../../common";
import Locale from "./locale"; import Locale from "./locale";
export type SettingKeys = export type SettingKeys =
@@ -13,6 +13,12 @@ export const AnthropicMetas = {
Vision: "2023-06-01", Vision: "2023-06-01",
}; };
export const ClaudeMapper = {
assistant: "assistant",
user: "user",
system: "user",
} as const;
export const modelConfigs = [ export const modelConfigs = [
{ {
name: "claude-instant-1.2", name: "claude-instant-1.2",
@@ -58,6 +64,8 @@ export const modelConfigs = [
}, },
]; ];
const defaultEndpoint = "/api/anthropic";
export const settingItems: SettingItem<SettingKeys>[] = [ export const settingItems: SettingItem<SettingKeys>[] = [
{ {
name: "anthropicUrl", name: "anthropicUrl",
@@ -65,7 +73,22 @@ export const settingItems: SettingItem<SettingKeys>[] = [
description: Locale.Endpoint.SubTitle + AnthropicMetas.ExampleEndpoint, description: Locale.Endpoint.SubTitle + AnthropicMetas.ExampleEndpoint,
placeholder: AnthropicMetas.ExampleEndpoint, placeholder: AnthropicMetas.ExampleEndpoint,
type: "input", type: "input",
validators: ["required"], defaultValue: defaultEndpoint,
validators: [
"required",
async (v: any) => {
if (typeof v === "string" && !v.startsWith(defaultEndpoint)) {
try {
new URL(v);
} catch (e) {
return Locale.Endpoint.Error.IllegalURL;
}
}
if (typeof v === "string" && v.endsWith("/")) {
return Locale.Endpoint.Error.EndWithBackslash;
}
},
],
}, },
{ {
name: "anthropicApiKey", name: "anthropicApiKey",
@@ -74,7 +97,7 @@ export const settingItems: SettingItem<SettingKeys>[] = [
placeholder: Locale.ApiKey.Placeholder, placeholder: Locale.ApiKey.Placeholder,
type: "input", type: "input",
inputType: "password", inputType: "password",
validators: ["required"], // validators: ["required"],
}, },
{ {
name: "anthropicApiVersion", name: "anthropicApiVersion",
@@ -82,6 +105,6 @@ export const settingItems: SettingItem<SettingKeys>[] = [
description: Locale.ApiVerion.SubTitle, description: Locale.ApiVerion.SubTitle,
placeholder: AnthropicMetas.Vision, placeholder: AnthropicMetas.Vision,
type: "input", type: "input",
validators: ["required"], // validators: ["required"],
}, },
]; ];

View File

@@ -1,29 +1,27 @@
import { getMessageTextContent } from "@/app/utils";
import { import {
AnthropicMetas, AnthropicMetas,
ClaudeMapper,
SettingKeys, SettingKeys,
modelConfigs, modelConfigs,
settingItems, settingItems,
} from "./config"; } from "./config";
import { import {
ChatHandlers,
InternalChatRequestPayload, InternalChatRequestPayload,
IProviderTemplate, IProviderTemplate,
} from "../../core/types"; getMessageTextContent,
RequestMessage,
} from "../../common";
import { import {
EventStreamContentType, EventStreamContentType,
fetchEventSource, fetchEventSource,
} from "@fortaine/fetch-event-source"; } from "@fortaine/fetch-event-source";
import Locale from "@/app/locales"; import Locale from "@/app/locales";
import { prettyObject } from "@/app/utils/format"; import { getAuthKey, trimEnd, prettyObject } from "./utils";
import { cloneDeep } from "lodash-es";
export type AnthropicProviderSettingKeys = SettingKeys; export type AnthropicProviderSettingKeys = SettingKeys;
const ClaudeMapper = {
assistant: "assistant",
user: "user",
system: "user",
} as const;
export type MultiBlockContent = { export type MultiBlockContent = {
type: "image" | "text"; type: "image" | "text";
source?: { source?: {
@@ -75,64 +73,25 @@ export default class AnthropicProvider
settingItems, settingItems,
}; };
models = modelConfigs.map((c) => ({ ...c, providerTemplateName: this.name })); defaultModels = modelConfigs;
readonly REQUEST_TIMEOUT_MS = 60000; readonly REQUEST_TIMEOUT_MS = 60000;
private path(payload: InternalChatRequestPayload<SettingKeys>) { private path(payload: InternalChatRequestPayload<SettingKeys>) {
const { const {
providerConfig: { anthropicUrl }, providerConfig: { anthropicUrl },
context: { isApp },
} = payload; } = payload;
let baseUrl: string = anthropicUrl; return `${trimEnd(anthropicUrl!)}/${AnthropicMetas.ChatPath}`;
// if endpoint is empty, use default endpoint
if (baseUrl.trim().length === 0) {
baseUrl = "/api/anthropic";
}
if (!baseUrl.startsWith("http") && !baseUrl.startsWith("/api")) {
baseUrl = "https://" + baseUrl;
}
baseUrl = trimEnd(baseUrl, "/");
return `${baseUrl}/${AnthropicMetas.ChatPath}`;
} }
private formatChatPayload(payload: InternalChatRequestPayload<SettingKeys>) { private formatMessage(
const { messages: RequestMessage[],
messages, payload: InternalChatRequestPayload<SettingKeys>,
isVisionModel, ) {
model, const { isVisionModel } = payload;
stream,
modelConfig,
providerConfig,
} = payload;
const { anthropicApiKey, anthropicApiVersion, anthropicUrl } =
providerConfig;
const { temperature, top_p, max_tokens } = modelConfig;
const keys = ["system", "user"]; return messages
// roles must alternate between "user" and "assistant" in claude, so add a fake assistant message between two user messages
for (let i = 0; i < messages.length - 1; i++) {
const message = messages[i];
const nextMessage = messages[i + 1];
if (keys.includes(message.role) && keys.includes(nextMessage.role)) {
messages[i] = [
message,
{
role: "assistant",
content: ";",
},
] as any;
}
}
const prompt = messages
.flat() .flat()
.filter((v) => { .filter((v) => {
if (!v.content) return false; if (!v.content) return false;
@@ -180,6 +139,40 @@ export default class AnthropicProvider
}), }),
}; };
}); });
}
private formatChatPayload(payload: InternalChatRequestPayload<SettingKeys>) {
const {
messages: outsideMessages,
model,
stream,
modelConfig,
providerConfig,
} = payload;
const { anthropicApiKey, anthropicApiVersion } = providerConfig;
const { temperature, top_p, max_tokens } = modelConfig;
const keys = ["system", "user"];
// roles must alternate between "user" and "assistant" in claude, so add a fake assistant message between two user messages
const messages = cloneDeep(outsideMessages);
for (let i = 0; i < messages.length - 1; i++) {
const message = messages[i];
const nextMessage = messages[i + 1];
if (keys.includes(message.role) && keys.includes(nextMessage.role)) {
messages[i] = [
message,
{
role: "assistant",
content: ";",
},
] as any;
}
}
const prompt = this.formatMessage(messages, payload);
const requestBody: AnthropicChatRequest = { const requestBody: AnthropicChatRequest = {
messages: prompt, messages: prompt,
@@ -196,7 +189,7 @@ export default class AnthropicProvider
"Content-Type": "application/json", "Content-Type": "application/json",
Accept: "application/json", Accept: "application/json",
"x-api-key": anthropicApiKey ?? "", "x-api-key": anthropicApiKey ?? "",
"anthropic-version": anthropicApiVersion, "anthropic-version": anthropicApiVersion ?? "",
Authorization: getAuthKey(anthropicApiKey), Authorization: getAuthKey(anthropicApiKey),
}, },
body: JSON.stringify(requestBody), body: JSON.stringify(requestBody),
@@ -204,6 +197,7 @@ export default class AnthropicProvider
url: this.path(payload), url: this.path(payload),
}; };
} }
private readWholeMessageResponseBody(res: any) { private readWholeMessageResponseBody(res: any) {
return { return {
message: res?.content?.[0]?.text ?? "", message: res?.content?.[0]?.text ?? "",
@@ -259,50 +253,12 @@ export default class AnthropicProvider
streamChat( streamChat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
onProgress: (message: string, chunk: string) => void, handlers: ChatHandlers,
onFinish: (message: string) => void,
onError: (err: Error) => void,
) { ) {
const requestPayload = this.formatChatPayload(payload); const requestPayload = this.formatChatPayload(payload);
let responseText = "";
let remainText = "";
let finished = false;
const timer = this.getTimer(); const timer = this.getTimer();
// animate response to make it looks smooth
const animateResponseText = () => {
if (finished || timer.signal.aborted) {
responseText += remainText;
console.log("[Response Animation] finished");
if (responseText?.length === 0) {
onError(new Error("empty response from server"));
}
return;
}
if (remainText.length > 0) {
const fetchCount = Math.max(1, Math.round(remainText.length / 60));
const fetchText = remainText.slice(0, fetchCount);
responseText += fetchText;
remainText = remainText.slice(fetchCount);
onProgress(responseText, fetchText);
}
requestAnimationFrame(animateResponseText);
};
// start animaion
animateResponseText();
const finish = () => {
if (!finished) {
finished = true;
onFinish(responseText + remainText);
}
};
fetchEventSource(requestPayload.url, { fetchEventSource(requestPayload.url, {
...requestPayload, ...requestPayload,
async onopen(res) { async onopen(res) {
@@ -311,8 +267,8 @@ export default class AnthropicProvider
console.log("[OpenAI] request response content type: ", contentType); console.log("[OpenAI] request response content type: ", contentType);
if (contentType?.startsWith("text/plain")) { if (contentType?.startsWith("text/plain")) {
responseText = await res.clone().text(); const responseText = await res.clone().text();
return finish(); return handlers.onFlash(responseText);
} }
if ( if (
@@ -322,29 +278,29 @@ export default class AnthropicProvider
?.startsWith(EventStreamContentType) || ?.startsWith(EventStreamContentType) ||
res.status !== 200 res.status !== 200
) { ) {
const responseTexts = [responseText]; const responseTexts = [];
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
let extraInfo = await res.clone().text(); let extraInfo = await res.clone().text();
try { try {
const resJson = await res.clone().json(); const resJson = await res.clone().json();
extraInfo = prettyObject(resJson); extraInfo = prettyObject(resJson);
} catch {} } catch {}
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
if (extraInfo) { if (extraInfo) {
responseTexts.push(extraInfo); responseTexts.push(extraInfo);
} }
responseText = responseTexts.join("\n\n"); const responseText = responseTexts.join("\n\n");
return finish(); return handlers.onFlash(responseText);
} }
}, },
onmessage(msg) { onmessage(msg) {
if (msg.data === "[DONE]" || finished) { if (msg.data === "[DONE]") {
return finish(); return;
} }
const text = msg.data; const text = msg.data;
try { try {
@@ -353,20 +309,19 @@ export default class AnthropicProvider
delta: { content: string }; delta: { content: string };
}>; }>;
const delta = choices[0]?.delta?.content; const delta = choices[0]?.delta?.content;
const textmoderation = json?.prompt_filter_results;
if (delta) { if (delta) {
remainText += delta; handlers.onProgress(delta);
} }
} catch (e) { } catch (e) {
console.error("[Request] parse error", text, msg); console.error("[Request] parse error", text, msg);
} }
}, },
onclose() { onclose() {
finish(); handlers.onFinish();
}, },
onerror(e) { onerror(e) {
onError(e); handlers.onError(e);
throw e; throw e;
}, },
openWhenHidden: true, openWhenHidden: true,
@@ -375,28 +330,3 @@ export default class AnthropicProvider
return timer; return timer;
} }
} }
function trimEnd(s: string, end = " ") {
if (end.length === 0) return s;
while (s.endsWith(end)) {
s = s.slice(0, -end.length);
}
return s;
}
function bearer(value: string) {
return `Bearer ${value.trim()}`;
}
function getAuthKey(apiKey = "") {
let authKey = "";
if (apiKey) {
// use user's api key first
authKey = bearer(apiKey);
}
return authKey;
}

View File

@@ -1,4 +1,4 @@
import { getLocaleText } from "../../core/locale"; import { getLocaleText } from "../../common";
export default getLocaleText< export default getLocaleText<
{ {
@@ -10,6 +10,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: string; Title: string;
SubTitle: string; SubTitle: string;
Error: {
EndWithBackslash: string;
IllegalURL: string;
};
}; };
ApiVerion: { ApiVerion: {
Title: string; Title: string;
@@ -29,6 +33,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "接口地址", Title: "接口地址",
SubTitle: "样例:", SubTitle: "样例:",
Error: {
EndWithBackslash: "不能以「/」结尾",
IllegalURL: "请输入一个完整可用的url",
},
}, },
ApiVerion: { ApiVerion: {
@@ -47,6 +55,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Endpoint Address", Title: "Endpoint Address",
SubTitle: "Example:", SubTitle: "Example:",
Error: {
EndWithBackslash: "Cannot end with '/'",
IllegalURL: "Please enter a complete available url",
},
}, },
ApiVerion: { ApiVerion: {
@@ -64,6 +76,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Endpoint Address", Title: "Endpoint Address",
SubTitle: "Exemplo: ", SubTitle: "Exemplo: ",
Error: {
EndWithBackslash: "Não é possível terminar com '/'",
IllegalURL: "Insira um URL completo disponível",
},
}, },
ApiVerion: { ApiVerion: {
@@ -81,6 +97,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Adresa koncového bodu", Title: "Adresa koncového bodu",
SubTitle: "Príklad:", SubTitle: "Príklad:",
Error: {
EndWithBackslash: "Nemôže končiť znakom „/“",
IllegalURL: "Zadajte úplnú dostupnú adresu URL",
},
}, },
ApiVerion: { ApiVerion: {
@@ -98,6 +118,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "終端地址", Title: "終端地址",
SubTitle: "範例:", SubTitle: "範例:",
Error: {
EndWithBackslash: "不能以「/」結尾",
IllegalURL: "請輸入一個完整可用的url",
},
}, },
ApiVerion: { ApiVerion: {

View File

@@ -0,0 +1,38 @@
export function trimEnd(s: string, end = " ") {
if (end.length === 0) return s;
while (s.endsWith(end)) {
s = s.slice(0, -end.length);
}
return s;
}
export function bearer(value: string) {
return `Bearer ${value.trim()}`;
}
export function getAuthKey(apiKey = "") {
let authKey = "";
if (apiKey) {
// use user's api key first
authKey = bearer(apiKey);
}
return authKey;
}
export function prettyObject(msg: any) {
const obj = msg;
if (typeof msg !== "string") {
msg = JSON.stringify(msg, null, " ");
}
if (msg === "{}") {
return obj.toString();
}
if (msg.startsWith("```json")) {
return msg;
}
return ["```json", msg, "```"].join("\n");
}

View File

@@ -1,12 +1,11 @@
import Locale from "./locale"; import Locale from "./locale";
import { SettingItem } from "../../core/types"; import { SettingItem } from "../../common";
import { modelConfigs as openaiModelConfigs } from "../openai/config"; import { modelConfigs as openaiModelConfigs } from "../openai/config";
export const AzureMetas = { export const AzureMetas = {
ExampleEndpoint: "https://{resource-url}/openai/deployments/{deploy-id}", ExampleEndpoint: "https://{resource-url}/openai/deployments/{deploy-id}",
ChatPath: "v1/chat/completions", ChatPath: "v1/chat/completions",
OpenAI: "/api/openai",
}; };
export type SettingKeys = "azureUrl" | "azureApiKey" | "azureApiVersion"; export type SettingKeys = "azureUrl" | "azureApiKey" | "azureApiVersion";
@@ -20,6 +19,21 @@ export const settingItems: SettingItem<SettingKeys>[] = [
description: Locale.Endpoint.SubTitle + AzureMetas.ExampleEndpoint, description: Locale.Endpoint.SubTitle + AzureMetas.ExampleEndpoint,
placeholder: AzureMetas.ExampleEndpoint, placeholder: AzureMetas.ExampleEndpoint,
type: "input", type: "input",
validators: [
async (v: any) => {
if (typeof v === "string") {
try {
new URL(v);
} catch (e) {
return Locale.Endpoint.Error.IllegalURL;
}
}
if (typeof v === "string" && v.endsWith("/")) {
return Locale.Endpoint.Error.EndWithBackslash;
}
},
"required",
],
}, },
{ {
name: "azureApiKey", name: "azureApiKey",

View File

@@ -1,15 +1,17 @@
import { settingItems, SettingKeys, modelConfigs, AzureMetas } from "./config"; import { settingItems, SettingKeys, modelConfigs, AzureMetas } from "./config";
import { import {
ChatHandlers,
InternalChatRequestPayload, InternalChatRequestPayload,
IProviderTemplate, IProviderTemplate,
} from "../../core/types"; ModelInfo,
import { getMessageTextContent } from "@/app/utils"; getMessageTextContent,
} from "../../common";
import { import {
EventStreamContentType, EventStreamContentType,
fetchEventSource, fetchEventSource,
} from "@fortaine/fetch-event-source"; } from "@fortaine/fetch-event-source";
import { prettyObject } from "@/app/utils/format";
import Locale from "@/app/locales"; import Locale from "@/app/locales";
import { makeAzurePath, makeBearer, prettyObject, validString } from "./utils";
export type AzureProviderSettingKeys = SettingKeys; export type AzureProviderSettingKeys = SettingKeys;
@@ -43,13 +45,30 @@ interface RequestPayload {
max_tokens?: number; max_tokens?: number;
} }
interface ModelList {
object: "list";
data: Array<{
capabilities: {
fine_tune: boolean;
inference: boolean;
completion: boolean;
chat_completion: boolean;
embeddings: boolean;
};
lifecycle_status: "generally-available";
id: string;
created_at: number;
object: "model";
}>;
}
export default class Azure export default class Azure
implements IProviderTemplate<SettingKeys, "azure", typeof AzureMetas> implements IProviderTemplate<SettingKeys, "azure", typeof AzureMetas>
{ {
name = "azure" as const; name = "azure" as const;
metas = AzureMetas; metas = AzureMetas;
models = modelConfigs.map((c) => ({ ...c, providerTemplateName: this.name })); defaultModels = modelConfigs;
providerMeta = { providerMeta = {
displayName: "Azure", displayName: "Azure",
@@ -62,25 +81,11 @@ export default class Azure
const { const {
providerConfig: { azureUrl, azureApiVersion }, providerConfig: { azureUrl, azureApiVersion },
} = payload; } = payload;
const path = makeAzurePath(AzureMetas.ChatPath, azureApiVersion!);
const path = makeAzurePath(AzureMetas.ChatPath, azureApiVersion); console.log("[Proxy Endpoint] ", azureUrl, path);
let baseUrl = azureUrl; return [azureUrl!, path].join("/");
if (!baseUrl) {
baseUrl = "/api/openai";
}
if (baseUrl.endsWith("/")) {
baseUrl = baseUrl.slice(0, baseUrl.length - 1);
}
if (!baseUrl.startsWith("http") && !baseUrl.startsWith(AzureMetas.OpenAI)) {
baseUrl = "https://" + baseUrl;
}
console.log("[Proxy Endpoint] ", baseUrl, path);
return [baseUrl, path].join("/");
} }
private getHeaders(payload: InternalChatRequestPayload<SettingKeys>) { private getHeaders(payload: InternalChatRequestPayload<SettingKeys>) {
@@ -90,14 +95,9 @@ export default class Azure
"Content-Type": "application/json", "Content-Type": "application/json",
Accept: "application/json", Accept: "application/json",
}; };
const authHeader = "Authorization";
const makeBearer = (s: string) => `Bearer ${s.trim()}`;
const validString = (x?: string): x is string => Boolean(x && x.length > 0);
// when using google api in app, not set auth header
if (validString(azureApiKey)) { if (validString(azureApiKey)) {
headers[authHeader] = makeBearer(azureApiKey); headers["Authorization"] = makeBearer(azureApiKey);
} }
return headers; return headers;
@@ -197,52 +197,12 @@ export default class Azure
streamChat( streamChat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
onProgress: (message: string, chunk: string) => void, handlers: ChatHandlers,
onFinish: (message: string) => void,
onError: (err: Error) => void,
) { ) {
const requestPayload = this.formatChatPayload(payload); const requestPayload = this.formatChatPayload(payload);
const timer = this.getTimer(); const timer = this.getTimer();
let responseText = "";
let remainText = "";
let finished = false;
// animate response to make it looks smooth
const animateResponseText = () => {
if (finished || timer.signal.aborted) {
responseText += remainText;
console.log("[Response Animation] finished");
if (responseText?.length === 0) {
onError(new Error("empty response from server"));
}
return;
}
if (remainText.length > 0) {
const fetchCount = Math.max(1, Math.round(remainText.length / 60));
const fetchText = remainText.slice(0, fetchCount);
responseText += fetchText;
remainText = remainText.slice(fetchCount);
onProgress(responseText, fetchText);
}
requestAnimationFrame(animateResponseText);
};
// start animaion
animateResponseText();
const finish = () => {
if (!finished) {
finished = true;
onFinish(responseText + remainText);
}
};
timer.signal.onabort = finish;
fetchEventSource(requestPayload.url, { fetchEventSource(requestPayload.url, {
...requestPayload, ...requestPayload,
async onopen(res) { async onopen(res) {
@@ -251,8 +211,8 @@ export default class Azure
console.log("[OpenAI] request response content type: ", contentType); console.log("[OpenAI] request response content type: ", contentType);
if (contentType?.startsWith("text/plain")) { if (contentType?.startsWith("text/plain")) {
responseText = await res.clone().text(); const responseText = await res.clone().text();
return finish(); return handlers.onFlash(responseText);
} }
if ( if (
@@ -262,29 +222,29 @@ export default class Azure
?.startsWith(EventStreamContentType) || ?.startsWith(EventStreamContentType) ||
res.status !== 200 res.status !== 200
) { ) {
const responseTexts = [responseText]; const responseTexts = [];
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
let extraInfo = await res.clone().text(); let extraInfo = await res.clone().text();
try { try {
const resJson = await res.clone().json(); const resJson = await res.clone().json();
extraInfo = prettyObject(resJson); extraInfo = prettyObject(resJson);
} catch {} } catch {}
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
if (extraInfo) { if (extraInfo) {
responseTexts.push(extraInfo); responseTexts.push(extraInfo);
} }
responseText = responseTexts.join("\n\n"); const responseText = responseTexts.join("\n\n");
return finish(); return handlers.onFlash(responseText);
} }
}, },
onmessage(msg) { onmessage(msg) {
if (msg.data === "[DONE]" || finished) { if (msg.data === "[DONE]") {
return finish(); return;
} }
const text = msg.data; const text = msg.data;
try { try {
@@ -293,34 +253,41 @@ export default class Azure
delta: { content: string }; delta: { content: string };
}>; }>;
const delta = choices[0]?.delta?.content; const delta = choices[0]?.delta?.content;
const textmoderation = json?.prompt_filter_results;
if (delta) { if (delta) {
remainText += delta; handlers.onProgress(delta);
} }
} catch (e) { } catch (e) {
console.error("[Request] parse error", text, msg); console.error("[Request] parse error", text, msg);
} }
}, },
onclose() { onclose() {
finish(); handlers.onFinish();
}, },
onerror(e) { onerror(e) {
onError(e); handlers.onError(e);
throw e; throw e;
}, },
openWhenHidden: true, openWhenHidden: true,
}); });
return timer; return timer;
} }
}
async getAvailableModels(
function makeAzurePath(path: string, apiVersion: string) { providerConfig: Record<SettingKeys, string>,
// should omit /v1 prefix ): Promise<ModelInfo[]> {
path = path.replaceAll("v1/", ""); const { azureApiKey, azureUrl } = providerConfig;
const res = await fetch(`${azureUrl}/vi/models`, {
// should add api-key to query string headers: {
path += `${path.includes("?") ? "&" : "?"}api-version=${apiVersion}`; Authorization: `Bearer ${azureApiKey}`,
},
return path; method: "GET",
});
const data: ModelList = await res.json();
return data.data.map((o) => ({
name: o.id,
}));
}
} }

View File

@@ -1,4 +1,4 @@
import { getLocaleText } from "../../core/locale"; import { getLocaleText } from "../../common";
export default getLocaleText< export default getLocaleText<
{ {
@@ -10,6 +10,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: string; Title: string;
SubTitle: string; SubTitle: string;
Error: {
EndWithBackslash: string;
IllegalURL: string;
};
}; };
ApiVerion: { ApiVerion: {
Title: string; Title: string;
@@ -29,6 +33,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "接口地址", Title: "接口地址",
SubTitle: "样例:", SubTitle: "样例:",
Error: {
EndWithBackslash: "不能以「/」结尾",
IllegalURL: "请输入一个完整可用的url",
},
}, },
ApiVerion: { ApiVerion: {
@@ -46,6 +54,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Azure Endpoint", Title: "Azure Endpoint",
SubTitle: "Example: ", SubTitle: "Example: ",
Error: {
EndWithBackslash: "Cannot end with '/'",
IllegalURL: "Please enter a complete available url",
},
}, },
ApiVerion: { ApiVerion: {
@@ -63,6 +75,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Endpoint Azure", Title: "Endpoint Azure",
SubTitle: "Exemplo: ", SubTitle: "Exemplo: ",
Error: {
EndWithBackslash: "Não é possível terminar com '/'",
IllegalURL: "Insira um URL completo disponível",
},
}, },
ApiVerion: { ApiVerion: {
@@ -80,6 +96,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Koncový bod Azure", Title: "Koncový bod Azure",
SubTitle: "Príklad: ", SubTitle: "Príklad: ",
Error: {
EndWithBackslash: "Nemôže končiť znakom „/“",
IllegalURL: "Zadajte úplnú dostupnú adresu URL",
},
}, },
ApiVerion: { ApiVerion: {
@@ -97,6 +117,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "介面(Endpoint) 地址", Title: "介面(Endpoint) 地址",
SubTitle: "樣例:", SubTitle: "樣例:",
Error: {
EndWithBackslash: "不能以「/」結尾",
IllegalURL: "請輸入一個完整可用的url",
},
}, },
ApiVerion: { ApiVerion: {

View File

@@ -0,0 +1,27 @@
export function makeAzurePath(path: string, apiVersion: string) {
// should omit /v1 prefix
path = path.replaceAll("v1/", "");
// should add api-key to query string
path += `${path.includes("?") ? "&" : "?"}api-version=${apiVersion}`;
return path;
}
export function prettyObject(msg: any) {
const obj = msg;
if (typeof msg !== "string") {
msg = JSON.stringify(msg, null, " ");
}
if (msg === "{}") {
return obj.toString();
}
if (msg.startsWith("```json")) {
return msg;
}
return ["```json", msg, "```"].join("\n");
}
export const makeBearer = (s: string) => `Bearer ${s.trim()}`;
export const validString = (x?: string): x is string =>
Boolean(x && x.length > 0);

View File

@@ -1,11 +1,9 @@
import { SettingItem } from "../../core/types"; import { SettingItem } from "../../common";
import Locale from "./locale"; import Locale from "./locale";
export const GoogleMetas = { export const GoogleMetas = {
ExampleEndpoint: "https://generativelanguage.googleapis.com/", ExampleEndpoint: "https://generativelanguage.googleapis.com/",
ChatPath: (modelName: string) => `v1beta/models/${modelName}:generateContent`, ChatPath: (modelName: string) => `v1beta/models/${modelName}:generateContent`,
VisionChatPath: (modelName: string) =>
`v1beta/models/${modelName}:generateContent`,
}; };
export type SettingKeys = "googleUrl" | "googleApiKey" | "googleApiVersion"; export type SettingKeys = "googleUrl" | "googleApiKey" | "googleApiVersion";
@@ -41,7 +39,20 @@ export const settingItems: SettingItem<SettingKeys>[] = [
description: Locale.Endpoint.SubTitle + GoogleMetas.ExampleEndpoint, description: Locale.Endpoint.SubTitle + GoogleMetas.ExampleEndpoint,
placeholder: GoogleMetas.ExampleEndpoint, placeholder: GoogleMetas.ExampleEndpoint,
type: "input", type: "input",
validators: ["required"], validators: [
async (v: any) => {
if (typeof v === "string") {
try {
new URL(v);
} catch (e) {
return Locale.Endpoint.Error.IllegalURL;
}
}
if (typeof v === "string" && v.endsWith("/")) {
return Locale.Endpoint.Error.EndWithBackslash;
}
},
],
}, },
{ {
name: "googleApiKey", name: "googleApiKey",
@@ -50,7 +61,7 @@ export const settingItems: SettingItem<SettingKeys>[] = [
placeholder: Locale.ApiKey.Placeholder, placeholder: Locale.ApiKey.Placeholder,
type: "input", type: "input",
inputType: "password", inputType: "password",
validators: ["required"], // validators: ["required"],
}, },
{ {
name: "googleApiVersion", name: "googleApiVersion",
@@ -58,6 +69,6 @@ export const settingItems: SettingItem<SettingKeys>[] = [
description: Locale.ApiVersion.SubTitle, description: Locale.ApiVersion.SubTitle,
placeholder: "2023-08-01-preview", placeholder: "2023-08-01-preview",
type: "input", type: "input",
validators: ["required"], // validators: ["required"],
}, },
]; ];

View File

@@ -1,13 +1,34 @@
import { getMessageImages, getMessageTextContent } from "@/app/utils";
import { SettingKeys, modelConfigs, settingItems, GoogleMetas } from "./config"; import { SettingKeys, modelConfigs, settingItems, GoogleMetas } from "./config";
import { import {
ChatHandlers,
InternalChatRequestPayload, InternalChatRequestPayload,
IProviderTemplate, IProviderTemplate,
ModelInfo,
StandChatReponseMessage, StandChatReponseMessage,
} from "../../core/types"; getMessageTextContent,
getMessageImages,
} from "../../common";
import { ensureProperEnding, makeBearer, validString } from "./utils";
export type GoogleProviderSettingKeys = SettingKeys; export type GoogleProviderSettingKeys = SettingKeys;
interface ModelList {
models: Array<{
name: string;
baseModelId: string;
version: string;
displayName: string;
description: string;
inputTokenLimit: number; // Integer
outputTokenLimit: number; // Integer
supportedGenerationMethods: [string];
temperature: number;
topP: number;
topK: number; // Integer
}>;
nextPageToken: string;
}
export default class GoogleProvider export default class GoogleProvider
implements IProviderTemplate<SettingKeys, "google", typeof GoogleMetas> implements IProviderTemplate<SettingKeys, "google", typeof GoogleMetas>
{ {
@@ -18,7 +39,7 @@ export default class GoogleProvider
displayName: "Google", displayName: "Google",
settingItems, settingItems,
}; };
models = modelConfigs.map((c) => ({ ...c, providerTemplateName: this.name })); defaultModels = modelConfigs;
readonly REQUEST_TIMEOUT_MS = 60000; readonly REQUEST_TIMEOUT_MS = 60000;
@@ -33,19 +54,8 @@ export default class GoogleProvider
Accept: "application/json", Accept: "application/json",
}; };
const authHeader = "Authorization"; if (!isApp && validString(googleApiKey)) {
headers["Authorization"] = makeBearer(googleApiKey);
const makeBearer = (s: string) => `Bearer ${s.trim()}`;
const validString = (x?: string): x is string => Boolean(x && x.length > 0);
// when using google api in app, not set auth header
if (!isApp) {
// use user's api key first
if (validString(googleApiKey)) {
headers[authHeader] = makeBearer(googleApiKey);
} else {
throw new Error("no apiKey when chat through google");
}
} }
return headers; return headers;
@@ -135,15 +145,9 @@ export default class GoogleProvider
], ],
}; };
let baseUrl = googleUrl; let googleChatPath = GoogleMetas.ChatPath(model);
let googleChatPath = isVisionModel let baseUrl = googleUrl ?? "/api/google/" + googleChatPath;
? GoogleMetas.VisionChatPath(model)
: GoogleMetas.ChatPath(model);
if (!baseUrl) {
baseUrl = "/api/google/" + googleChatPath;
}
if (isApp) { if (isApp) {
baseUrl += `?key=${googleApiKey}`; baseUrl += `?key=${googleApiKey}`;
@@ -193,44 +197,13 @@ export default class GoogleProvider
streamChat( streamChat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
onProgress: (message: string, chunk: string) => void, handlers: ChatHandlers,
onFinish: (message: string) => void,
onError: (err: Error) => void,
) { ) {
const requestPayload = this.formatChatPayload(payload); const requestPayload = this.formatChatPayload(payload);
let responseText = "";
let remainText = "";
let finished = false;
const timer = this.getTimer(); const timer = this.getTimer();
let existingTexts: string[] = []; let existingTexts: string[] = [];
const finish = () => {
finished = true;
onFinish(existingTexts.join(""));
};
// animate response to make it looks smooth
const animateResponseText = () => {
if (finished || timer.signal.aborted) {
responseText += remainText;
finish();
return;
}
if (remainText.length > 0) {
const fetchCount = Math.max(1, Math.round(remainText.length / 60));
const fetchText = remainText.slice(0, fetchCount);
responseText += fetchText;
remainText = remainText.slice(fetchCount);
onProgress(responseText, fetchText);
}
requestAnimationFrame(animateResponseText);
};
// start animaion
animateResponseText();
fetch(requestPayload.url, { fetch(requestPayload.url, {
...requestPayload, ...requestPayload,
@@ -250,18 +223,16 @@ export default class GoogleProvider
try { try {
let data = JSON.parse(ensureProperEnding(partialData)); let data = JSON.parse(ensureProperEnding(partialData));
if (data && data[0].error) { if (data && data[0].error) {
onError(new Error(data[0].error.message)); handlers.onError(new Error(data[0].error.message));
} else { } else {
onError(new Error("Request failed")); handlers.onError(new Error("Request failed"));
} }
} catch (_) { } catch (_) {
onError(new Error("Request failed")); handlers.onError(new Error("Request failed"));
} }
} }
console.log("Stream complete"); console.log("Stream complete");
// options.onFinish(responseText + remainText);
finished = true;
return Promise.resolve(); return Promise.resolve();
} }
@@ -285,7 +256,7 @@ export default class GoogleProvider
if (textArray.length > existingTexts.length) { if (textArray.length > existingTexts.length) {
const deltaArray = textArray.slice(existingTexts.length); const deltaArray = textArray.slice(existingTexts.length);
existingTexts = textArray; existingTexts = textArray;
remainText += deltaArray.join(""); handlers.onProgress(deltaArray.join(""));
} }
} catch (error) { } catch (error) {
// console.log("[Response Animation] error: ", error,partialData); // console.log("[Response Animation] error: ", error,partialData);
@@ -300,6 +271,7 @@ export default class GoogleProvider
}); });
return timer; return timer;
} }
async chat( async chat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
): Promise<StandChatReponseMessage> { ): Promise<StandChatReponseMessage> {
@@ -328,11 +300,19 @@ export default class GoogleProvider
return message; return message;
} }
}
function ensureProperEnding(str: string) { async getAvailableModels(
if (str.startsWith("[") && !str.endsWith("]")) { providerConfig: Record<SettingKeys, string>,
return str + "]"; ): Promise<ModelInfo[]> {
const { googleApiKey, googleUrl } = providerConfig;
const res = await fetch(`${googleUrl}/v1beta/models?key=${googleApiKey}`, {
headers: {
Authorization: `Bearer ${googleApiKey}`,
},
method: "GET",
});
const data: ModelList = await res.json();
return data.models;
} }
return str;
} }

View File

@@ -1,4 +1,4 @@
import { getLocaleText } from "../../core/locale"; import { getLocaleText } from "../../common";
export default getLocaleText< export default getLocaleText<
{ {
@@ -10,6 +10,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: string; Title: string;
SubTitle: string; SubTitle: string;
Error: {
EndWithBackslash: string;
IllegalURL: string;
};
}; };
ApiVersion: { ApiVersion: {
Title: string; Title: string;
@@ -29,6 +33,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "终端地址", Title: "终端地址",
SubTitle: "示例:", SubTitle: "示例:",
Error: {
EndWithBackslash: "不能以「/」结尾",
IllegalURL: "请输入一个完整可用的url",
},
}, },
ApiVersion: { ApiVersion: {
@@ -46,6 +54,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Endpoint Address", Title: "Endpoint Address",
SubTitle: "Example:", SubTitle: "Example:",
Error: {
EndWithBackslash: "Cannot end with '/'",
IllegalURL: "Please enter a complete available url",
},
}, },
ApiVersion: { ApiVersion: {
@@ -64,6 +76,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Adresa koncového bodu", Title: "Adresa koncového bodu",
SubTitle: "Príklad:", SubTitle: "Príklad:",
Error: {
EndWithBackslash: "Nemôže končiť znakom „/“",
IllegalURL: "Zadajte úplnú dostupnú adresu URL",
},
}, },
ApiVersion: { ApiVersion: {
@@ -81,6 +97,10 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "終端地址", Title: "終端地址",
SubTitle: "範例:", SubTitle: "範例:",
Error: {
EndWithBackslash: "不能以「/」結尾",
IllegalURL: "請輸入一個完整可用的url",
},
}, },
ApiVersion: { ApiVersion: {

View File

@@ -0,0 +1,10 @@
export const makeBearer = (s: string) => `Bearer ${s.trim()}`;
export const validString = (x?: string): x is string =>
Boolean(x && x.length > 0);
export function ensureProperEnding(str: string) {
if (str.startsWith("[") && !str.endsWith("]")) {
return str + "]";
}
return str;
}

View File

@@ -1,4 +1,4 @@
import { SettingItem } from "../../core/types"; import { SettingItem } from "../../common";
import { isVisionModel } from "@/app/utils"; import { isVisionModel } from "@/app/utils";
import Locale from "@/app/locales"; import Locale from "@/app/locales";

View File

@@ -4,19 +4,21 @@ import {
SettingKeys, SettingKeys,
NextChatMetas, NextChatMetas,
} from "./config"; } from "./config";
import { getMessageTextContent } from "@/app/utils";
import { ACCESS_CODE_PREFIX } from "@/app/constant"; import { ACCESS_CODE_PREFIX } from "@/app/constant";
import { import {
ChatHandlers,
getMessageTextContent,
InternalChatRequestPayload, InternalChatRequestPayload,
IProviderTemplate, IProviderTemplate,
StandChatReponseMessage, StandChatReponseMessage,
} from "../../core/types"; } from "../../common";
import { import {
EventStreamContentType, EventStreamContentType,
fetchEventSource, fetchEventSource,
} from "@fortaine/fetch-event-source"; } from "@fortaine/fetch-event-source";
import { prettyObject } from "@/app/utils/format"; import { prettyObject } from "@/app/utils/format";
import Locale from "@/app/locales"; import Locale from "@/app/locales";
import { makeBearer, validString } from "./utils";
export type NextChatProviderSettingKeys = SettingKeys; export type NextChatProviderSettingKeys = SettingKeys;
@@ -56,7 +58,7 @@ export default class NextChatProvider
name = "nextchat" as const; name = "nextchat" as const;
metas = NextChatMetas; metas = NextChatMetas;
models = modelConfigs.map((c) => ({ ...c, providerTemplateName: this.name })); defaultModels = modelConfigs;
providerMeta = { providerMeta = {
displayName: "NextChat", displayName: "NextChat",
@@ -82,14 +84,9 @@ export default class NextChatProvider
"Content-Type": "application/json", "Content-Type": "application/json",
Accept: "application/json", Accept: "application/json",
}; };
const authHeader = "Authorization";
const makeBearer = (s: string) => `Bearer ${s.trim()}`;
const validString = (x?: string): x is string => Boolean(x && x.length > 0);
// when using google api in app, not set auth header
if (validString(accessCode)) { if (validString(accessCode)) {
headers[authHeader] = makeBearer(ACCESS_CODE_PREFIX + accessCode); headers["Authorization"] = makeBearer(ACCESS_CODE_PREFIX + accessCode);
} }
return headers; return headers;
@@ -160,52 +157,12 @@ export default class NextChatProvider
streamChat( streamChat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
onProgress: (message: string, chunk: string) => void, handlers: ChatHandlers,
onFinish: (message: string) => void,
onError: (err: Error) => void,
) { ) {
const requestPayload = this.formatChatPayload(payload); const requestPayload = this.formatChatPayload(payload);
let responseText = "";
let remainText = "";
let finished = false;
const timer = this.getTimer(); const timer = this.getTimer();
// animate response to make it looks smooth
const animateResponseText = () => {
if (finished || timer.signal.aborted) {
responseText += remainText;
console.log("[Response Animation] finished");
if (responseText?.length === 0) {
onError(new Error("empty response from server"));
}
return;
}
if (remainText.length > 0) {
const fetchCount = Math.max(1, Math.round(remainText.length / 60));
const fetchText = remainText.slice(0, fetchCount);
responseText += fetchText;
remainText = remainText.slice(fetchCount);
onProgress(responseText, fetchText);
}
requestAnimationFrame(animateResponseText);
};
// start animaion
animateResponseText();
const finish = () => {
if (!finished) {
finished = true;
onFinish(responseText + remainText);
}
};
timer.signal.onabort = finish;
fetchEventSource(requestPayload.url, { fetchEventSource(requestPayload.url, {
...requestPayload, ...requestPayload,
async onopen(res) { async onopen(res) {
@@ -214,8 +171,8 @@ export default class NextChatProvider
console.log("[OpenAI] request response content type: ", contentType); console.log("[OpenAI] request response content type: ", contentType);
if (contentType?.startsWith("text/plain")) { if (contentType?.startsWith("text/plain")) {
responseText = await res.clone().text(); const responseText = await res.clone().text();
return finish(); return handlers.onFlash(responseText);
} }
if ( if (
@@ -225,29 +182,29 @@ export default class NextChatProvider
?.startsWith(EventStreamContentType) || ?.startsWith(EventStreamContentType) ||
res.status !== 200 res.status !== 200
) { ) {
const responseTexts = [responseText]; const responseTexts = [];
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
let extraInfo = await res.clone().text(); let extraInfo = await res.clone().text();
try { try {
const resJson = await res.clone().json(); const resJson = await res.clone().json();
extraInfo = prettyObject(resJson); extraInfo = prettyObject(resJson);
} catch {} } catch {}
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
if (extraInfo) { if (extraInfo) {
responseTexts.push(extraInfo); responseTexts.push(extraInfo);
} }
responseText = responseTexts.join("\n\n"); const responseText = responseTexts.join("\n\n");
return finish(); return handlers.onFlash(responseText);
} }
}, },
onmessage(msg) { onmessage(msg) {
if (msg.data === "[DONE]" || finished) { if (msg.data === "[DONE]") {
return finish(); return;
} }
const text = msg.data; const text = msg.data;
try { try {
@@ -256,20 +213,19 @@ export default class NextChatProvider
delta: { content: string }; delta: { content: string };
}>; }>;
const delta = choices[0]?.delta?.content; const delta = choices[0]?.delta?.content;
const textmoderation = json?.prompt_filter_results;
if (delta) { if (delta) {
remainText += delta; handlers.onProgress(delta);
} }
} catch (e) { } catch (e) {
console.error("[Request] parse error", text, msg); console.error("[Request] parse error", text, msg);
} }
}, },
onclose() { onclose() {
finish(); handlers.onFinish();
}, },
onerror(e) { onerror(e) {
onError(e); handlers.onError(e);
throw e; throw e;
}, },
openWhenHidden: true, openWhenHidden: true,
@@ -277,6 +233,7 @@ export default class NextChatProvider
return timer; return timer;
} }
async chat( async chat(
payload: InternalChatRequestPayload<"accessCode">, payload: InternalChatRequestPayload<"accessCode">,
): Promise<StandChatReponseMessage> { ): Promise<StandChatReponseMessage> {

View File

@@ -0,0 +1,18 @@
export const makeBearer = (s: string) => `Bearer ${s.trim()}`;
export const validString = (x?: string): x is string =>
Boolean(x && x.length > 0);
export function prettyObject(msg: any) {
const obj = msg;
if (typeof msg !== "string") {
msg = JSON.stringify(msg, null, " ");
}
if (msg === "{}") {
return obj.toString();
}
if (msg.startsWith("```json")) {
return msg;
}
return ["```json", msg, "```"].join("\n");
}

View File

@@ -1,8 +1,10 @@
import { SettingItem } from "../../core/types"; import { SettingItem } from "../../common";
import Locale from "./locale"; import Locale from "./locale";
export const OPENAI_BASE_URL = "https://api.openai.com"; export const OPENAI_BASE_URL = "https://api.openai.com";
export const ROLES = ["system", "user", "assistant"] as const;
export const OpenaiMetas = { export const OpenaiMetas = {
ChatPath: "v1/chat/completions", ChatPath: "v1/chat/completions",
UsagePath: "dashboard/billing/usage", UsagePath: "dashboard/billing/usage",
@@ -12,15 +14,20 @@ export const OpenaiMetas = {
export type SettingKeys = "openaiUrl" | "openaiApiKey"; export type SettingKeys = "openaiUrl" | "openaiApiKey";
export const defaultModal = "gpt-3.5-turbo";
export const modelConfigs = [ export const modelConfigs = [
{
name: "gpt-4o",
displayName: "gpt-4o",
isVision: false,
isDefaultActive: true,
isDefaultSelected: true,
},
{ {
name: "gpt-3.5-turbo", name: "gpt-3.5-turbo",
displayName: "gpt-3.5-turbo", displayName: "gpt-3.5-turbo",
isVision: false, isVision: false,
isDefaultActive: true, isDefaultActive: true,
isDefaultSelected: true, isDefaultSelected: false,
}, },
{ {
name: "gpt-3.5-turbo-0301", name: "gpt-3.5-turbo-0301",
@@ -150,13 +157,30 @@ export const modelConfigs = [
}, },
]; ];
const defaultEndpoint = "/api/openai";
export const settingItems: SettingItem<SettingKeys>[] = [ export const settingItems: SettingItem<SettingKeys>[] = [
{ {
name: "openaiUrl", name: "openaiUrl",
title: Locale.Endpoint.Title, title: Locale.Endpoint.Title,
description: Locale.Endpoint.SubTitle, description: Locale.Endpoint.SubTitle,
defaultValue: OPENAI_BASE_URL, defaultValue: defaultEndpoint,
type: "input", type: "input",
validators: [
"required",
async (v: any) => {
if (typeof v === "string" && v.endsWith("/")) {
return Locale.Endpoint.Error.EndWithBackslash;
}
if (
typeof v === "string" &&
!v.startsWith(defaultEndpoint) &&
!v.startsWith("http")
) {
return Locale.Endpoint.SubTitle;
}
},
],
}, },
{ {
name: "openaiApiKey", name: "openaiApiKey",

View File

@@ -1,19 +1,26 @@
import { modelConfigs, settingItems, SettingKeys, OpenaiMetas } from "./config";
import { getMessageTextContent } from "@/app/utils";
import { import {
ChatHandlers,
InternalChatRequestPayload, InternalChatRequestPayload,
IProviderTemplate, IProviderTemplate,
} from "../../core/types"; ModelInfo,
getMessageTextContent,
} from "../../common";
import { import {
EventStreamContentType, EventStreamContentType,
fetchEventSource, fetchEventSource,
} from "@fortaine/fetch-event-source"; } from "@fortaine/fetch-event-source";
import { prettyObject } from "@/app/utils/format";
import Locale from "@/app/locales"; import Locale from "@/app/locales";
import { makeBearer, validString, prettyObject } from "./utils";
import {
modelConfigs,
settingItems,
SettingKeys,
OpenaiMetas,
ROLES,
} from "./config";
export type OpenAIProviderSettingKeys = SettingKeys; export type OpenAIProviderSettingKeys = SettingKeys;
export const ROLES = ["system", "user", "assistant"] as const;
export type MessageRole = (typeof ROLES)[number]; export type MessageRole = (typeof ROLES)[number];
export interface MultimodalContent { export interface MultimodalContent {
@@ -28,7 +35,6 @@ export interface RequestMessage {
role: MessageRole; role: MessageRole;
content: string | MultimodalContent[]; content: string | MultimodalContent[];
} }
interface RequestPayload { interface RequestPayload {
messages: { messages: {
role: "system" | "user" | "assistant"; role: "system" | "user" | "assistant";
@@ -43,6 +49,16 @@ interface RequestPayload {
max_tokens?: number; max_tokens?: number;
} }
interface ModelList {
object: "list";
data: Array<{
id: string;
object: "model";
created: number;
owned_by: "system" | "openai-internal";
}>;
}
class OpenAIProvider class OpenAIProvider
implements IProviderTemplate<SettingKeys, "openai", typeof OpenaiMetas> implements IProviderTemplate<SettingKeys, "openai", typeof OpenaiMetas>
{ {
@@ -51,7 +67,7 @@ class OpenAIProvider
readonly REQUEST_TIMEOUT_MS = 60000; readonly REQUEST_TIMEOUT_MS = 60000;
models = modelConfigs.map((c) => ({ ...c, providerTemplateName: this.name })); defaultModels = modelConfigs;
providerMeta = { providerMeta = {
displayName: "OpenAI", displayName: "OpenAI",
@@ -62,25 +78,11 @@ class OpenAIProvider
const { const {
providerConfig: { openaiUrl }, providerConfig: { openaiUrl },
} = payload; } = payload;
const path = OpenaiMetas.ChatPath; const path = OpenaiMetas.ChatPath;
let baseUrl = openaiUrl; console.log("[Proxy Endpoint] ", openaiUrl, path);
if (!baseUrl) { return [openaiUrl, path].join("/");
baseUrl = "/api/openai";
}
if (baseUrl.endsWith("/")) {
baseUrl = baseUrl.slice(0, baseUrl.length - 1);
}
if (!baseUrl.startsWith("http") && !baseUrl.startsWith("/api/openai")) {
baseUrl = "https://" + baseUrl;
}
console.log("[Proxy Endpoint] ", baseUrl, path);
return [baseUrl, path].join("/");
} }
private getHeaders(payload: InternalChatRequestPayload<SettingKeys>) { private getHeaders(payload: InternalChatRequestPayload<SettingKeys>) {
@@ -90,14 +92,9 @@ class OpenAIProvider
"Content-Type": "application/json", "Content-Type": "application/json",
Accept: "application/json", Accept: "application/json",
}; };
const authHeader = "Authorization";
const makeBearer = (s: string) => `Bearer ${s.trim()}`;
const validString = (x?: string): x is string => Boolean(x && x.length > 0);
// when using google api in app, not set auth header
if (validString(openaiApiKey)) { if (validString(openaiApiKey)) {
headers[authHeader] = makeBearer(openaiApiKey); headers["Authorization"] = makeBearer(openaiApiKey);
} }
return headers; return headers;
@@ -143,9 +140,11 @@ class OpenAIProvider
}; };
} }
private readWholeMessageResponseBody(res: any) { private readWholeMessageResponseBody(res: {
choices: { message: { content: any } }[];
}) {
return { return {
message: res.choices?.at(0)?.message?.content ?? "", message: res.choices?.[0]?.message?.content ?? "",
}; };
} }
@@ -190,52 +189,12 @@ class OpenAIProvider
streamChat( streamChat(
payload: InternalChatRequestPayload<SettingKeys>, payload: InternalChatRequestPayload<SettingKeys>,
onProgress: (message: string, chunk: string) => void, handlers: ChatHandlers,
onFinish: (message: string) => void,
onError: (err: Error) => void,
) { ) {
const requestPayload = this.formatChatPayload(payload); const requestPayload = this.formatChatPayload(payload);
const timer = this.getTimer(); const timer = this.getTimer();
let responseText = "";
let remainText = "";
let finished = false;
// animate response to make it looks smooth
const animateResponseText = () => {
if (finished || timer.signal.aborted) {
responseText += remainText;
console.log("[Response Animation] finished");
if (responseText?.length === 0) {
onError(new Error("empty response from server"));
}
return;
}
if (remainText.length > 0) {
const fetchCount = Math.max(1, Math.round(remainText.length / 60));
const fetchText = remainText.slice(0, fetchCount);
responseText += fetchText;
remainText = remainText.slice(fetchCount);
onProgress(responseText, fetchText);
}
requestAnimationFrame(animateResponseText);
};
// start animaion
animateResponseText();
const finish = () => {
if (!finished) {
finished = true;
onFinish(responseText + remainText);
}
};
timer.signal.onabort = finish;
fetchEventSource(requestPayload.url, { fetchEventSource(requestPayload.url, {
...requestPayload, ...requestPayload,
async onopen(res) { async onopen(res) {
@@ -244,8 +203,8 @@ class OpenAIProvider
console.log("[OpenAI] request response content type: ", contentType); console.log("[OpenAI] request response content type: ", contentType);
if (contentType?.startsWith("text/plain")) { if (contentType?.startsWith("text/plain")) {
responseText = await res.clone().text(); const responseText = await res.clone().text();
return finish(); return handlers.onFlash(responseText);
} }
if ( if (
@@ -255,29 +214,29 @@ class OpenAIProvider
?.startsWith(EventStreamContentType) || ?.startsWith(EventStreamContentType) ||
res.status !== 200 res.status !== 200
) { ) {
const responseTexts = [responseText]; const responseTexts = [];
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
let extraInfo = await res.clone().text(); let extraInfo = await res.clone().text();
try { try {
const resJson = await res.clone().json(); const resJson = await res.clone().json();
extraInfo = prettyObject(resJson); extraInfo = prettyObject(resJson);
} catch {} } catch {}
if (res.status === 401) {
responseTexts.push(Locale.Error.Unauthorized);
}
if (extraInfo) { if (extraInfo) {
responseTexts.push(extraInfo); responseTexts.push(extraInfo);
} }
responseText = responseTexts.join("\n\n"); const responseText = responseTexts.join("\n\n");
return finish(); return handlers.onFlash(responseText);
} }
}, },
onmessage(msg) { onmessage(msg) {
if (msg.data === "[DONE]" || finished) { if (msg.data === "[DONE]") {
return finish(); return;
} }
const text = msg.data; const text = msg.data;
try { try {
@@ -286,20 +245,19 @@ class OpenAIProvider
delta: { content: string }; delta: { content: string };
}>; }>;
const delta = choices[0]?.delta?.content; const delta = choices[0]?.delta?.content;
const textmoderation = json?.prompt_filter_results;
if (delta) { if (delta) {
remainText += delta; handlers.onProgress(delta);
} }
} catch (e) { } catch (e) {
console.error("[Request] parse error", text, msg); console.error("[Request] parse error", text, msg);
} }
}, },
onclose() { onclose() {
finish(); handlers.onFinish();
}, },
onerror(e) { onerror(e) {
onError(e); handlers.onError(e);
throw e; throw e;
}, },
openWhenHidden: true, openWhenHidden: true,
@@ -307,6 +265,23 @@ class OpenAIProvider
return timer; return timer;
} }
async getAvailableModels(
providerConfig: Record<SettingKeys, string>,
): Promise<ModelInfo[]> {
const { openaiApiKey, openaiUrl } = providerConfig;
const res = await fetch(`${openaiUrl}/vi/models`, {
headers: {
Authorization: `Bearer ${openaiApiKey}`,
},
method: "GET",
});
const data: ModelList = await res.json();
return data.data.map((o) => ({
name: o.id,
}));
}
} }
export default OpenAIProvider; export default OpenAIProvider;

View File

@@ -1,4 +1,4 @@
import { getLocaleText } from "../../core/locale"; import { getLocaleText } from "../../common/locale";
export default getLocaleText< export default getLocaleText<
{ {
@@ -11,6 +11,9 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: string; Title: string;
SubTitle: string; SubTitle: string;
Error: {
EndWithBackslash: string;
};
}; };
}, },
"en" "en"
@@ -26,6 +29,9 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "接口地址", Title: "接口地址",
SubTitle: "除默认地址外,必须包含 http(s)://", SubTitle: "除默认地址外,必须包含 http(s)://",
Error: {
EndWithBackslash: "不能以「/」结尾",
},
}, },
}, },
en: { en: {
@@ -38,6 +44,9 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "OpenAI Endpoint", Title: "OpenAI Endpoint",
SubTitle: "Must starts with http(s):// or use /api/openai as default", SubTitle: "Must starts with http(s):// or use /api/openai as default",
Error: {
EndWithBackslash: "Cannot end with '/'",
},
}, },
}, },
pt: { pt: {
@@ -50,6 +59,9 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "Endpoint OpenAI", Title: "Endpoint OpenAI",
SubTitle: "Deve começar com http(s):// ou usar /api/openai como padrão", SubTitle: "Deve começar com http(s):// ou usar /api/openai como padrão",
Error: {
EndWithBackslash: "Não é possível terminar com '/'",
},
}, },
}, },
sk: { sk: {
@@ -63,6 +75,9 @@ export default getLocaleText<
Title: "Koncový bod OpenAI", Title: "Koncový bod OpenAI",
SubTitle: SubTitle:
"Musí začínať http(s):// alebo použiť /api/openai ako predvolený", "Musí začínať http(s):// alebo použiť /api/openai ako predvolený",
Error: {
EndWithBackslash: "Nemôže končiť znakom „/“",
},
}, },
}, },
tw: { tw: {
@@ -75,6 +90,9 @@ export default getLocaleText<
Endpoint: { Endpoint: {
Title: "介面(Endpoint) 地址", Title: "介面(Endpoint) 地址",
SubTitle: "除預設地址外,必須包含 http(s)://", SubTitle: "除預設地址外,必須包含 http(s)://",
Error: {
EndWithBackslash: "不能以「/」結尾",
},
}, },
}, },
}, },

View File

@@ -0,0 +1,18 @@
export const makeBearer = (s: string) => `Bearer ${s.trim()}`;
export const validString = (x?: string): x is string =>
Boolean(x && x.length > 0);
export function prettyObject(msg: any) {
const obj = msg;
if (typeof msg !== "string") {
msg = JSON.stringify(msg, null, " ");
}
if (msg === "{}") {
return obj.toString();
}
if (msg.startsWith("```json")) {
return msg;
}
return ["```json", msg, "```"].join("\n");
}

View File

@@ -37,6 +37,8 @@ type Error =
error: false; error: false;
}; };
type Validate = (v: any) => Error | Promise<Error>;
export interface ListItemProps { export interface ListItemProps {
title: string; title: string;
subTitle?: string; subTitle?: string;
@@ -44,7 +46,7 @@ export interface ListItemProps {
className?: string; className?: string;
onClick?: () => void; onClick?: () => void;
nextline?: boolean; nextline?: boolean;
validator?: (v: any) => Error | Promise<Error>; validator?: Validate | Validate[];
} }
export const ListContext = createContext< export const ListContext = createContext<
@@ -92,7 +94,15 @@ export function ListItem(props: ListItemProps) {
}, []); }, []);
const handleValidate = useCallback((v: any) => { const handleValidate = useCallback((v: any) => {
const insideValidator = validator || (() => {}); let insideValidator;
if (!validator) {
insideValidator = () => {};
} else if (Array.isArray(validator)) {
insideValidator = (v: any) =>
Promise.race(validator.map((validate) => validate(v)));
} else {
insideValidator = validator;
}
Promise.resolve(insideValidator(v)).then((result) => { Promise.resolve(insideValidator(v)).then((result) => {
if (result && result.error) { if (result && result.error) {

View File

@@ -9,22 +9,37 @@ import {
import { StoreKey } from "../constant"; import { StoreKey } from "../constant";
import { createPersistStore } from "../utils/store"; import { createPersistStore } from "../utils/store";
export const DEFAULT_CONFIG = { const firstUpdate = Date.now();
lastUpdate: Date.now(), // timestamp, to merge state
providers: ProviderClient.getProviderTemplateList() function getDefaultConfig() {
.filter((p) => p !== NextChatProvider) const providers = Object.values(ProviderClient.ProviderTemplates)
.map((p) => createProvider(p)), .filter((t) => !(t instanceof NextChatProvider))
}; .map((t) => createProvider(t, true));
export type ProvidersConfig = typeof DEFAULT_CONFIG; const initProvider = providers[0];
const currentModel =
initProvider.models.find((m) => m.isDefaultSelected) ||
initProvider.models[0];
return {
lastUpdate: firstUpdate, // timestamp, to merge state
currentModel: currentModel.name,
currentProvider: initProvider.name,
providers,
};
}
export type ProvidersConfig = ReturnType<typeof getDefaultConfig>;
export const useProviders = createPersistStore( export const useProviders = createPersistStore(
{ ...DEFAULT_CONFIG }, { ...getDefaultConfig() },
(set, get) => { (set, get) => {
const methods = { const methods = {
reset() { reset() {
set(() => ({ ...DEFAULT_CONFIG })); set(() => getDefaultConfig());
}, },
addProvider(provider: Provider) { addProvider(provider: Provider) {
@@ -53,10 +68,14 @@ export const useProviders = createPersistStore(
return get().providers.find((p) => p.name === providerName); return get().providers.find((p) => p.name === providerName);
}, },
addModel(model: Omit<Model, "providerTemplateName">, provider: Provider) { addModel(
model: Omit<Model, "providerTemplateName" | "customized">,
provider: Provider,
) {
const newModel: Model = { const newModel: Model = {
providerTemplateName: provider.providerTemplateName,
...model, ...model,
providerTemplateName: provider.providerTemplateName,
customized: true,
}; };
return methods.updateProvider({ return methods.updateProvider({
...provider, ...provider,
@@ -80,6 +99,13 @@ export const useProviders = createPersistStore(
}); });
}, },
switchModel(model: Model, provider: Provider) {
set(() => ({
currentModel: model.name,
currentProvider: provider.name,
}));
},
getModel( getModel(
modelName: string, modelName: string,
providerName: string, providerName: string,

View File

@@ -1,6 +1,6 @@
import { useMemo } from "react"; import { useMemo } from "react";
import { useAccessStore, useAppConfig } from "../store"; import { useAccessStore, useAppConfig } from "../store";
import { collectModels, collectModelsWithDefaultModel } from "./model"; import { collectModelsWithDefaultModel } from "./model";
export function useAllModels() { export function useAllModels() {
const accessStore = useAccessStore(); const accessStore = useAccessStore();

View File

@@ -5312,16 +5312,11 @@ mz@^2.7.0:
object-assign "^4.0.1" object-assign "^4.0.1"
thenify-all "^1.0.0" thenify-all "^1.0.0"
nanoid@^3.3.6: nanoid@^3.3.6, nanoid@^3.3.7:
version "3.3.7" version "3.3.7"
resolved "https://registry.yarnpkg.com/nanoid/-/nanoid-3.3.7.tgz#d0c301a691bc8d54efa0a2226ccf3fe2fd656bd8" resolved "https://registry.yarnpkg.com/nanoid/-/nanoid-3.3.7.tgz#d0c301a691bc8d54efa0a2226ccf3fe2fd656bd8"
integrity sha512-eSRppjcPIatRIMC1U6UngP8XFcz8MQWGQdt1MTBQ7NaAmvXDfvNxbvWV3x2y6CdEUciCSsDHDQZbhYaB8QEo2g== integrity sha512-eSRppjcPIatRIMC1U6UngP8XFcz8MQWGQdt1MTBQ7NaAmvXDfvNxbvWV3x2y6CdEUciCSsDHDQZbhYaB8QEo2g==
nanoid@^3.3.7:
version "3.3.7"
resolved "https://registry.npmmirror.com/nanoid/-/nanoid-3.3.7.tgz#d0c301a691bc8d54efa0a2226ccf3fe2fd656bd8"
integrity sha512-eSRppjcPIatRIMC1U6UngP8XFcz8MQWGQdt1MTBQ7NaAmvXDfvNxbvWV3x2y6CdEUciCSsDHDQZbhYaB8QEo2g==
nanoid@^5.0.3: nanoid@^5.0.3:
version "5.0.3" version "5.0.3"
resolved "https://registry.yarnpkg.com/nanoid/-/nanoid-5.0.3.tgz#6c97f53d793a7a1de6a38ebb46f50f95bf9793c7" resolved "https://registry.yarnpkg.com/nanoid/-/nanoid-5.0.3.tgz#6c97f53d793a7a1de6a38ebb46f50f95bf9793c7"