feat: allow custom seed (#6709)
This commit is contained in:
@@ -100,6 +100,17 @@ export class CopilotController {
|
|||||||
return controller.signal;
|
return controller.signal;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private parseNumber(value: string | string[] | undefined) {
|
||||||
|
if (!value) {
|
||||||
|
return undefined;
|
||||||
|
}
|
||||||
|
const num = Number.parseInt(Array.isArray(value) ? value[0] : value, 10);
|
||||||
|
if (Number.isNaN(num)) {
|
||||||
|
return undefined;
|
||||||
|
}
|
||||||
|
return num;
|
||||||
|
}
|
||||||
|
|
||||||
private handleError(err: any) {
|
private handleError(err: any) {
|
||||||
if (err instanceof Error) {
|
if (err instanceof Error) {
|
||||||
const ret = {
|
const ret = {
|
||||||
@@ -256,6 +267,7 @@ export class CopilotController {
|
|||||||
|
|
||||||
return from(
|
return from(
|
||||||
provider.generateImagesStream(session.finish(params), session.model, {
|
provider.generateImagesStream(session.finish(params), session.model, {
|
||||||
|
seed: this.parseNumber(params.seed),
|
||||||
signal: this.getSignal(req),
|
signal: this.getSignal(req),
|
||||||
user: user.id,
|
user: user.id,
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import assert from 'node:assert';
|
|||||||
|
|
||||||
import {
|
import {
|
||||||
CopilotCapability,
|
CopilotCapability,
|
||||||
|
CopilotImageOptions,
|
||||||
CopilotImageToImageProvider,
|
CopilotImageToImageProvider,
|
||||||
CopilotProviderType,
|
CopilotProviderType,
|
||||||
CopilotTextToImageProvider,
|
CopilotTextToImageProvider,
|
||||||
@@ -57,10 +58,7 @@ export class FalProvider
|
|||||||
async generateImages(
|
async generateImages(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string = this.availableModels[0],
|
model: string = this.availableModels[0],
|
||||||
options: {
|
options: CopilotImageOptions = {}
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
} = {}
|
|
||||||
): Promise<Array<string>> {
|
): Promise<Array<string>> {
|
||||||
const { content, attachments } = messages.pop() || {};
|
const { content, attachments } = messages.pop() || {};
|
||||||
if (!this.availableModels.includes(model)) {
|
if (!this.availableModels.includes(model)) {
|
||||||
@@ -82,7 +80,7 @@ export class FalProvider
|
|||||||
image_url: attachments?.[0],
|
image_url: attachments?.[0],
|
||||||
prompt: content,
|
prompt: content,
|
||||||
sync_mode: true,
|
sync_mode: true,
|
||||||
seed: 42,
|
seed: options.seed || 42,
|
||||||
enable_safety_checks: false,
|
enable_safety_checks: false,
|
||||||
}),
|
}),
|
||||||
signal: options.signal,
|
signal: options.signal,
|
||||||
@@ -100,10 +98,7 @@ export class FalProvider
|
|||||||
async *generateImagesStream(
|
async *generateImagesStream(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string = this.availableModels[0],
|
model: string = this.availableModels[0],
|
||||||
options: {
|
options: CopilotImageOptions = {}
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
} = {}
|
|
||||||
): AsyncIterable<string> {
|
): AsyncIterable<string> {
|
||||||
const ret = await this.generateImages(messages, model, options);
|
const ret = await this.generateImages(messages, model, options);
|
||||||
for (const url of ret) {
|
for (const url of ret) {
|
||||||
|
|||||||
@@ -5,6 +5,9 @@ import { ClientOptions, OpenAI } from 'openai';
|
|||||||
import {
|
import {
|
||||||
ChatMessageRole,
|
ChatMessageRole,
|
||||||
CopilotCapability,
|
CopilotCapability,
|
||||||
|
CopilotChatOptions,
|
||||||
|
CopilotEmbeddingOptions,
|
||||||
|
CopilotImageOptions,
|
||||||
CopilotImageToTextProvider,
|
CopilotImageToTextProvider,
|
||||||
CopilotProviderType,
|
CopilotProviderType,
|
||||||
CopilotTextToEmbeddingProvider,
|
CopilotTextToEmbeddingProvider,
|
||||||
@@ -147,12 +150,7 @@ export class OpenAIProvider
|
|||||||
async generateText(
|
async generateText(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string = 'gpt-3.5-turbo',
|
model: string = 'gpt-3.5-turbo',
|
||||||
options: {
|
options: CopilotChatOptions = {}
|
||||||
temperature?: number;
|
|
||||||
maxTokens?: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
} = {}
|
|
||||||
): Promise<string> {
|
): Promise<string> {
|
||||||
this.checkParams({ messages, model });
|
this.checkParams({ messages, model });
|
||||||
const result = await this.instance.chat.completions.create(
|
const result = await this.instance.chat.completions.create(
|
||||||
@@ -175,12 +173,7 @@ export class OpenAIProvider
|
|||||||
async *generateTextStream(
|
async *generateTextStream(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string = 'gpt-3.5-turbo',
|
model: string = 'gpt-3.5-turbo',
|
||||||
options: {
|
options: CopilotChatOptions = {}
|
||||||
temperature?: number;
|
|
||||||
maxTokens?: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
} = {}
|
|
||||||
): AsyncIterable<string> {
|
): AsyncIterable<string> {
|
||||||
this.checkParams({ messages, model });
|
this.checkParams({ messages, model });
|
||||||
const result = await this.instance.chat.completions.create(
|
const result = await this.instance.chat.completions.create(
|
||||||
@@ -214,11 +207,7 @@ export class OpenAIProvider
|
|||||||
async generateEmbedding(
|
async generateEmbedding(
|
||||||
messages: string | string[],
|
messages: string | string[],
|
||||||
model: string,
|
model: string,
|
||||||
options: {
|
options: CopilotEmbeddingOptions = { dimensions: DEFAULT_DIMENSIONS }
|
||||||
dimensions: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
} = { dimensions: DEFAULT_DIMENSIONS }
|
|
||||||
): Promise<number[][]> {
|
): Promise<number[][]> {
|
||||||
messages = Array.isArray(messages) ? messages : [messages];
|
messages = Array.isArray(messages) ? messages : [messages];
|
||||||
this.checkParams({ embeddings: messages, model });
|
this.checkParams({ embeddings: messages, model });
|
||||||
@@ -236,10 +225,7 @@ export class OpenAIProvider
|
|||||||
async generateImages(
|
async generateImages(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string = 'dall-e-3',
|
model: string = 'dall-e-3',
|
||||||
options: {
|
options: CopilotImageOptions = {}
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
} = {}
|
|
||||||
): Promise<Array<string>> {
|
): Promise<Array<string>> {
|
||||||
const { content: prompt } = messages.pop() || {};
|
const { content: prompt } = messages.pop() || {};
|
||||||
if (!prompt) {
|
if (!prompt) {
|
||||||
@@ -261,10 +247,7 @@ export class OpenAIProvider
|
|||||||
async *generateImagesStream(
|
async *generateImagesStream(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string = 'dall-e-3',
|
model: string = 'dall-e-3',
|
||||||
options: {
|
options: CopilotImageOptions = {}
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
} = {}
|
|
||||||
): AsyncIterable<string> {
|
): AsyncIterable<string> {
|
||||||
const ret = await this.generateImages(messages, model, options);
|
const ret = await this.generateImages(messages, model, options);
|
||||||
for (const url of ret) {
|
for (const url of ret) {
|
||||||
|
|||||||
@@ -143,6 +143,32 @@ export enum CopilotCapability {
|
|||||||
ImageToText = 'image-to-text',
|
ImageToText = 'image-to-text',
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const CopilotProviderOptionsSchema = z.object({
|
||||||
|
signal: z.instanceof(AbortSignal).optional(),
|
||||||
|
user: z.string().optional(),
|
||||||
|
});
|
||||||
|
|
||||||
|
const CopilotChatOptionsSchema = CopilotProviderOptionsSchema.extend({
|
||||||
|
temperature: z.number().optional(),
|
||||||
|
maxTokens: z.number().optional(),
|
||||||
|
}).optional();
|
||||||
|
|
||||||
|
export type CopilotChatOptions = z.infer<typeof CopilotChatOptionsSchema>;
|
||||||
|
|
||||||
|
const CopilotEmbeddingOptionsSchema = CopilotProviderOptionsSchema.extend({
|
||||||
|
dimensions: z.number(),
|
||||||
|
}).optional();
|
||||||
|
|
||||||
|
export type CopilotEmbeddingOptions = z.infer<
|
||||||
|
typeof CopilotEmbeddingOptionsSchema
|
||||||
|
>;
|
||||||
|
|
||||||
|
const CopilotImageOptionsSchema = CopilotProviderOptionsSchema.extend({
|
||||||
|
seed: z.number().optional(),
|
||||||
|
}).optional();
|
||||||
|
|
||||||
|
export type CopilotImageOptions = z.infer<typeof CopilotImageOptionsSchema>;
|
||||||
|
|
||||||
export interface CopilotProvider {
|
export interface CopilotProvider {
|
||||||
readonly type: CopilotProviderType;
|
readonly type: CopilotProviderType;
|
||||||
getCapabilities(): CopilotCapability[];
|
getCapabilities(): CopilotCapability[];
|
||||||
@@ -153,22 +179,12 @@ export interface CopilotTextToTextProvider extends CopilotProvider {
|
|||||||
generateText(
|
generateText(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model?: string,
|
model?: string,
|
||||||
options?: {
|
options?: CopilotChatOptions
|
||||||
temperature?: number;
|
|
||||||
maxTokens?: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): Promise<string>;
|
): Promise<string>;
|
||||||
generateTextStream(
|
generateTextStream(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model?: string,
|
model?: string,
|
||||||
options?: {
|
options?: CopilotChatOptions
|
||||||
temperature?: number;
|
|
||||||
maxTokens?: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): AsyncIterable<string>;
|
): AsyncIterable<string>;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -176,11 +192,7 @@ export interface CopilotTextToEmbeddingProvider extends CopilotProvider {
|
|||||||
generateEmbedding(
|
generateEmbedding(
|
||||||
messages: string[] | string,
|
messages: string[] | string,
|
||||||
model: string,
|
model: string,
|
||||||
options: {
|
options?: CopilotEmbeddingOptions
|
||||||
dimensions: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): Promise<number[][]>;
|
): Promise<number[][]>;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -188,18 +200,12 @@ export interface CopilotTextToImageProvider extends CopilotProvider {
|
|||||||
generateImages(
|
generateImages(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string,
|
model: string,
|
||||||
options: {
|
options?: CopilotImageOptions
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): Promise<Array<string>>;
|
): Promise<Array<string>>;
|
||||||
generateImagesStream(
|
generateImagesStream(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model?: string,
|
model?: string,
|
||||||
options?: {
|
options?: CopilotImageOptions
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): AsyncIterable<string>;
|
): AsyncIterable<string>;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -207,22 +213,12 @@ export interface CopilotImageToTextProvider extends CopilotProvider {
|
|||||||
generateText(
|
generateText(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string,
|
model: string,
|
||||||
options: {
|
options?: CopilotChatOptions
|
||||||
temperature?: number;
|
|
||||||
maxTokens?: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): Promise<string>;
|
): Promise<string>;
|
||||||
generateTextStream(
|
generateTextStream(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string,
|
model: string,
|
||||||
options: {
|
options?: CopilotChatOptions
|
||||||
temperature?: number;
|
|
||||||
maxTokens?: number;
|
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): AsyncIterable<string>;
|
): AsyncIterable<string>;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -230,18 +226,12 @@ export interface CopilotImageToImageProvider extends CopilotProvider {
|
|||||||
generateImages(
|
generateImages(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model: string,
|
model: string,
|
||||||
options: {
|
options?: CopilotImageOptions
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): Promise<Array<string>>;
|
): Promise<Array<string>>;
|
||||||
generateImagesStream(
|
generateImagesStream(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
model?: string,
|
model?: string,
|
||||||
options?: {
|
options?: CopilotImageOptions
|
||||||
signal?: AbortSignal;
|
|
||||||
user?: string;
|
|
||||||
}
|
|
||||||
): AsyncIterable<string>;
|
): AsyncIterable<string>;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user