From ccb3bed91ee8a88c692b67666bfd546d90941f5d Mon Sep 17 00:00:00 2001 From: DarkSky <25152247+darkskygit@users.noreply.github.com> Date: Wed, 17 Apr 2024 22:05:38 +0800 Subject: [PATCH] feat: add blob upload support for copilot (#6584) --- .../backend/server/src/core/quota/storage.ts | 66 ++++++++++++-- .../src/core/workspaces/resolvers/blob.ts | 45 ++------- .../src/fundamentals/config/storage/index.ts | 5 + .../server/src/plugins/copilot/controller.ts | 46 +++++++++- .../server/src/plugins/copilot/index.ts | 10 +- .../server/src/plugins/copilot/resolver.ts | 37 +++++++- .../server/src/plugins/copilot/session.ts | 12 +++ .../server/src/plugins/copilot/storage.ts | 91 +++++++++++++++++++ packages/backend/server/src/schema.gql | 1 + packages/frontend/graphql/src/schema.ts | 1 + 10 files changed, 260 insertions(+), 54 deletions(-) create mode 100644 packages/backend/server/src/plugins/copilot/storage.ts diff --git a/packages/backend/server/src/core/quota/storage.ts b/packages/backend/server/src/core/quota/storage.ts index 7ec0d510e..8ba20532b 100644 --- a/packages/backend/server/src/core/quota/storage.ts +++ b/packages/backend/server/src/core/quota/storage.ts @@ -7,7 +7,10 @@ import { OneGB } from './constant'; import { QuotaService } from './service'; import { formatSize, QuotaQueryType } from './types'; -type QuotaBusinessType = QuotaQueryType & { businessBlobLimit: number }; +type QuotaBusinessType = QuotaQueryType & { + businessBlobLimit: number; + unlimited: boolean; +}; @Injectable() export class QuotaManagementService { @@ -59,6 +62,52 @@ export class QuotaManagementService { }, 0); } + private generateQuotaCalculator( + quota: number, + blobLimit: number, + usedSize: number, + unlimited = false + ) { + const checkExceeded = (recvSize: number) => { + const total = usedSize + recvSize; + // only skip total storage check if workspace has unlimited feature + if (total > quota && !unlimited) { + this.logger.log(`storage size limit exceeded: ${total} > ${quota}`); + return true; + } else if (recvSize > blobLimit) { + this.logger.log(`blob size limit exceeded: ${recvSize} > ${blobLimit}`); + return true; + } else { + return false; + } + }; + return checkExceeded; + } + + async getQuotaCalculator(userId: string) { + const quota = await this.getUserQuota(userId); + const { storageQuota, businessBlobLimit } = quota; + const usedSize = await this.getUserUsage(userId); + + return this.generateQuotaCalculator( + storageQuota, + businessBlobLimit, + usedSize + ); + } + + async getQuotaCalculatorByWorkspace(workspaceId: string) { + const { storageQuota, usedSize, businessBlobLimit, unlimited } = + await this.getWorkspaceUsage(workspaceId); + + return this.generateQuotaCalculator( + storageQuota, + businessBlobLimit, + usedSize, + unlimited + ); + } + // get workspace's owner quota and total size of used // quota was apply to owner's account async getWorkspaceUsage(workspaceId: string): Promise { @@ -79,6 +128,12 @@ export class QuotaManagementService { } = await this.quota.getUserQuota(owner.id); // get all workspaces size of owner used const usedSize = await this.getUserUsage(owner.id); + // relax restrictions if workspace has unlimited feature + // todo(@darkskygit): need a mechanism to allow feature as a middleware to edit quota + const unlimited = await this.feature.hasWorkspaceFeature( + workspaceId, + FeatureType.UnlimitedWorkspace + ); const quota = { name, @@ -90,15 +145,10 @@ export class QuotaManagementService { copilotActionLimit, humanReadable, usedSize, + unlimited, }; - // relax restrictions if workspace has unlimited feature - // todo(@darkskygit): need a mechanism to allow feature as a middleware to edit quota - const unlimited = await this.feature.hasWorkspaceFeature( - workspaceId, - FeatureType.UnlimitedWorkspace - ); - if (unlimited) { + if (quota.unlimited) { return this.mergeUnlimitedQuota(quota); } diff --git a/packages/backend/server/src/core/workspaces/resolvers/blob.ts b/packages/backend/server/src/core/workspaces/resolvers/blob.ts index 335fea36b..9717dddb7 100644 --- a/packages/backend/server/src/core/workspaces/resolvers/blob.ts +++ b/packages/backend/server/src/core/workspaces/resolvers/blob.ts @@ -1,8 +1,4 @@ -import { - ForbiddenException, - Logger, - PayloadTooLargeException, -} from '@nestjs/common'; +import { Logger, PayloadTooLargeException, UseGuards } from '@nestjs/common'; import { Args, Int, @@ -16,20 +12,23 @@ import { SafeIntResolver } from 'graphql-scalars'; import GraphQLUpload from 'graphql-upload/GraphQLUpload.mjs'; import type { FileUpload } from '../../../fundamentals'; -import { MakeCache, PreventCache } from '../../../fundamentals'; +import { + CloudThrottlerGuard, + MakeCache, + PreventCache, +} from '../../../fundamentals'; import { CurrentUser } from '../../auth'; -import { FeatureManagementService, FeatureType } from '../../features'; import { QuotaManagementService } from '../../quota'; import { WorkspaceBlobStorage } from '../../storage'; import { PermissionService } from '../permission'; import { Permission, WorkspaceBlobSizes, WorkspaceType } from '../types'; +@UseGuards(CloudThrottlerGuard) @Resolver(() => WorkspaceType) export class WorkspaceBlobResolver { logger = new Logger(WorkspaceBlobResolver.name); constructor( private readonly permissions: PermissionService, - private readonly feature: FeatureManagementService, private readonly quota: QuotaManagementService, private readonly storage: WorkspaceBlobStorage ) {} @@ -124,34 +123,8 @@ export class WorkspaceBlobResolver { Permission.Write ); - const { storageQuota, usedSize, businessBlobLimit } = - await this.quota.getWorkspaceUsage(workspaceId); - - const unlimited = await this.feature.hasWorkspaceFeature( - workspaceId, - FeatureType.UnlimitedWorkspace - ); - - const checkExceeded = (recvSize: number) => { - if (!storageQuota) { - throw new ForbiddenException('Cannot find user quota.'); - } - const total = usedSize + recvSize; - // only skip total storage check if workspace has unlimited feature - if (total > storageQuota && !unlimited) { - this.logger.log( - `storage size limit exceeded: ${total} > ${storageQuota}` - ); - return true; - } else if (recvSize > businessBlobLimit) { - this.logger.log( - `blob size limit exceeded: ${recvSize} > ${businessBlobLimit}` - ); - return true; - } else { - return false; - } - }; + const checkExceeded = + await this.quota.getQuotaCalculatorByWorkspace(workspaceId); if (checkExceeded(0)) { throw new PayloadTooLargeException( diff --git a/packages/backend/server/src/fundamentals/config/storage/index.ts b/packages/backend/server/src/fundamentals/config/storage/index.ts index 32f989f86..a56404277 100644 --- a/packages/backend/server/src/fundamentals/config/storage/index.ts +++ b/packages/backend/server/src/fundamentals/config/storage/index.ts @@ -19,6 +19,7 @@ export type StorageConfig = { export interface StoragesConfig { avatar: StorageConfig<{ publicLinkFactory: (key: string) => string }>; blob: StorageConfig; + copilot: StorageConfig; } export interface AFFiNEStorageConfig { @@ -51,6 +52,10 @@ export function getDefaultAFFiNEStorageConfig(): AFFiNEStorageConfig { provider: 'fs', bucket: 'blobs', }, + copilot: { + provider: 'fs', + bucket: 'copilot', + }, }, }; } diff --git a/packages/backend/server/src/plugins/copilot/controller.ts b/packages/backend/server/src/plugins/copilot/controller.ts index bb5d128b6..d78c1a210 100644 --- a/packages/backend/server/src/plugins/copilot/controller.ts +++ b/packages/backend/server/src/plugins/copilot/controller.ts @@ -3,6 +3,8 @@ import { Controller, Get, InternalServerErrorException, + Logger, + NotFoundException, Param, Query, Req, @@ -17,15 +19,18 @@ import { from, map, merge, + mergeMap, Observable, switchMap, toArray, } from 'rxjs'; +import { Public } from '../../core/auth'; import { CurrentUser } from '../../core/auth/current-user'; import { Config } from '../../fundamentals'; import { CopilotProviderService } from './providers'; import { ChatSession, ChatSessionService } from './session'; +import { CopilotStorage } from './storage'; import { CopilotCapability } from './types'; export interface ChatEvent { @@ -36,10 +41,13 @@ export interface ChatEvent { @Controller('/api/copilot') export class CopilotController { + private readonly logger = new Logger(CopilotController.name); + constructor( private readonly config: Config, private readonly chatSession: ChatSessionService, - private readonly provider: CopilotProviderService + private readonly provider: CopilotProviderService, + private readonly storage: CopilotStorage ) {} private async hasAttachment(sessionId: string, messageId?: string) { @@ -230,12 +238,19 @@ export class CopilotController { delete params.message; delete params.messageId; + const handleRemoteLink = this.storage.handleRemoteLink.bind( + this.storage, + user.id, + sessionId + ); + return from( provider.generateImagesStream(session.finish(params), session.model, { signal: this.getSignal(req), user: user.id, }) ).pipe( + mergeMap(handleRemoteLink), connect(shared$ => merge( // actual chat event stream @@ -294,4 +309,33 @@ export class CopilotController { res.status(response.status).send(await response.json()); } + + @Public() + @Get('/blob/:userId/:workspaceId/:key') + async getBlob( + @Res() res: Response, + @Param('userId') userId: string, + @Param('workspaceId') workspaceId: string, + @Param('key') key: string + ) { + const { body, metadata } = await this.storage.get(userId, workspaceId, key); + + if (!body) { + throw new NotFoundException( + `Blob not found in ${userId}'s workspace ${workspaceId}: ${key}` + ); + } + + // metadata should always exists if body is not null + if (metadata) { + res.setHeader('content-type', metadata.contentType); + res.setHeader('last-modified', metadata.lastModified.toUTCString()); + res.setHeader('content-length', metadata.contentLength); + } else { + this.logger.warn(`Blob ${workspaceId}/${key} has no metadata`); + } + + res.setHeader('cache-control', 'public, max-age=2592000, immutable'); + body.pipe(res); + } } diff --git a/packages/backend/server/src/plugins/copilot/index.ts b/packages/backend/server/src/plugins/copilot/index.ts index 6d65f5f19..f3ecca2f9 100644 --- a/packages/backend/server/src/plugins/copilot/index.ts +++ b/packages/backend/server/src/plugins/copilot/index.ts @@ -1,6 +1,6 @@ import { ServerFeature } from '../../core/config'; -import { FeatureManagementService, FeatureService } from '../../core/features'; -import { QuotaService } from '../../core/quota'; +import { FeatureModule } from '../../core/features'; +import { QuotaModule } from '../../core/quota'; import { PermissionService } from '../../core/workspaces/permission'; import { Plugin } from '../registry'; import { CopilotController } from './controller'; @@ -15,23 +15,23 @@ import { } from './providers'; import { CopilotResolver, UserCopilotResolver } from './resolver'; import { ChatSessionService } from './session'; +import { CopilotStorage } from './storage'; registerCopilotProvider(FalProvider); registerCopilotProvider(OpenAIProvider); @Plugin({ name: 'copilot', + imports: [FeatureModule, QuotaModule], providers: [ PermissionService, - FeatureService, - FeatureManagementService, - QuotaService, ChatSessionService, CopilotResolver, ChatMessageCache, UserCopilotResolver, PromptService, CopilotProviderService, + CopilotStorage, ], controllers: [CopilotController], contributesTo: ServerFeature.Copilot, diff --git a/packages/backend/server/src/plugins/copilot/resolver.ts b/packages/backend/server/src/plugins/copilot/resolver.ts index 18a774d6c..f277c471c 100644 --- a/packages/backend/server/src/plugins/copilot/resolver.ts +++ b/packages/backend/server/src/plugins/copilot/resolver.ts @@ -1,4 +1,4 @@ -import { Logger } from '@nestjs/common'; +import { BadRequestException, Logger } from '@nestjs/common'; import { Args, Field, @@ -12,12 +12,18 @@ import { Resolver, } from '@nestjs/graphql'; import { GraphQLJSON, SafeIntResolver } from 'graphql-scalars'; +import GraphQLUpload from 'graphql-upload/GraphQLUpload.mjs'; import { CurrentUser } from '../../core/auth'; import { UserType } from '../../core/user'; import { PermissionService } from '../../core/workspaces/permission'; -import { MutexService, TooManyRequestsException } from '../../fundamentals'; +import { + FileUpload, + MutexService, + TooManyRequestsException, +} from '../../fundamentals'; import { ChatSessionService } from './session'; +import { CopilotStorage } from './storage'; import { AvailableModels, type ChatHistory, @@ -28,7 +34,7 @@ import { registerEnumType(AvailableModels, { name: 'CopilotModel' }); -const COPILOT_LOCKER = 'copilot'; +export const COPILOT_LOCKER = 'copilot'; // ================== Input Types ================== @@ -57,6 +63,9 @@ class CreateChatMessageInput implements Omit { @Field(() => [String], { nullable: true }) attachments!: string[] | undefined; + @Field(() => [GraphQLUpload], { nullable: true }) + blobs!: FileUpload[] | undefined; + @Field(() => GraphQLJSON, { nullable: true }) params!: Record | undefined; } @@ -140,7 +149,8 @@ export class CopilotResolver { constructor( private readonly permissions: PermissionService, private readonly mutex: MutexService, - private readonly chatSession: ChatSessionService + private readonly chatSession: ChatSessionService, + private readonly storage: CopilotStorage ) {} @ResolveField(() => CopilotQuotaType, { @@ -260,6 +270,25 @@ export class CopilotResolver { if (!lock) { return new TooManyRequestsException('Server is busy'); } + const session = await this.chatSession.get(options.sessionId); + if (!session) return new BadRequestException('Session not found'); + + if (options.blobs) { + options.attachments = options.attachments || []; + const { workspaceId } = session.config; + + for (const blob of options.blobs) { + const uploaded = await this.storage.handleUpload(user.id, blob); + const link = await this.storage.put( + user.id, + workspaceId, + uploaded.filename, + uploaded.buffer + ); + options.attachments.push(link); + } + } + try { return await this.chatSession.createMessage(options); } catch (e: any) { diff --git a/packages/backend/server/src/plugins/copilot/session.ts b/packages/backend/server/src/plugins/copilot/session.ts index 90014b3b7..93a887f50 100644 --- a/packages/backend/server/src/plugins/copilot/session.ts +++ b/packages/backend/server/src/plugins/copilot/session.ts @@ -35,6 +35,18 @@ export class ChatSession implements AsyncDisposable { return this.state.prompt.model; } + get config() { + const { + sessionId, + userId, + workspaceId, + docId, + prompt: { name: promptName }, + } = this.state; + + return { sessionId, userId, workspaceId, docId, promptName }; + } + push(message: ChatMessage) { if ( this.state.prompt.action && diff --git a/packages/backend/server/src/plugins/copilot/storage.ts b/packages/backend/server/src/plugins/copilot/storage.ts new file mode 100644 index 000000000..ed49b63c6 --- /dev/null +++ b/packages/backend/server/src/plugins/copilot/storage.ts @@ -0,0 +1,91 @@ +import { createHash } from 'node:crypto'; + +import { Injectable, PayloadTooLargeException } from '@nestjs/common'; + +import { QuotaManagementService } from '../../core/quota'; +import { + type BlobInputType, + Config, + type FileUpload, + type StorageProvider, + StorageProviderFactory, +} from '../../fundamentals'; + +@Injectable() +export class CopilotStorage { + public readonly provider: StorageProvider; + + constructor( + private readonly config: Config, + private readonly storageFactory: StorageProviderFactory, + private readonly quota: QuotaManagementService + ) { + this.provider = this.storageFactory.create('copilot'); + } + + async put( + userId: string, + workspaceId: string, + key: string, + blob: BlobInputType + ) { + const name = `${userId}/${workspaceId}/${key}`; + await this.provider.put(name, blob); + return `${this.config.baseUrl}/api/copilot/blob/${name}`; + } + + async get(userId: string, workspaceId: string, key: string) { + return this.provider.get(`${userId}/${workspaceId}/${key}`); + } + + async delete(userId: string, workspaceId: string, key: string) { + return this.provider.delete(`${userId}/${workspaceId}/${key}`); + } + + async handleUpload(userId: string, blob: FileUpload) { + const checkExceeded = await this.quota.getQuotaCalculator(userId); + + if (checkExceeded(0)) { + throw new PayloadTooLargeException( + 'Storage or blob size limit exceeded.' + ); + } + const buffer = await new Promise((resolve, reject) => { + const stream = blob.createReadStream(); + const chunks: Uint8Array[] = []; + stream.on('data', chunk => { + chunks.push(chunk); + + // check size after receive each chunk to avoid unnecessary memory usage + const bufferSize = chunks.reduce((acc, cur) => acc + cur.length, 0); + if (checkExceeded(bufferSize)) { + reject( + new PayloadTooLargeException('Storage or blob size limit exceeded.') + ); + } + }); + stream.on('error', reject); + stream.on('end', () => { + const buffer = Buffer.concat(chunks); + + if (checkExceeded(buffer.length)) { + reject(new PayloadTooLargeException('Storage limit exceeded.')); + } else { + resolve(buffer); + } + }); + }); + + return { + buffer, + filename: blob.filename, + }; + } + + async handleRemoteLink(userId: string, workspaceId: string, link: string) { + const response = await fetch(link); + const buffer = new Uint8Array(await response.arrayBuffer()); + const filename = createHash('sha256').update(buffer).digest('base64url'); + return this.put(userId, workspaceId, filename, Buffer.from(buffer)); + } +} diff --git a/packages/backend/server/src/schema.gql b/packages/backend/server/src/schema.gql index 7915e8b1f..da7ee6499 100644 --- a/packages/backend/server/src/schema.gql +++ b/packages/backend/server/src/schema.gql @@ -40,6 +40,7 @@ type CopilotQuota { input CreateChatMessageInput { attachments: [String!] + blobs: [Upload!] content: String params: JSON sessionId: String! diff --git a/packages/frontend/graphql/src/schema.ts b/packages/frontend/graphql/src/schema.ts index 23e8f5f0a..3d490692d 100644 --- a/packages/frontend/graphql/src/schema.ts +++ b/packages/frontend/graphql/src/schema.ts @@ -38,6 +38,7 @@ export interface Scalars { export interface CreateChatMessageInput { attachments: InputMaybe>; + blobs: InputMaybe>; content: InputMaybe; params: InputMaybe; sessionId: Scalars['String']['input'];