Commit 54a53e7d by Archer Committed by GitHub

fix: bind dataset training operations to collection ownership (#7071)

* fix: bind dataset training operations to collection ownership

* fix: remove stale training state dataset prop
parent 82fa431a
......@@ -14,4 +14,6 @@ description: 'FastGPT V4.15.0-beta4 Release Notes'
## 🐛 Bug Fixes
1. Fixed a potential unauthorized access risk in training APIs.
## 🛠️ Code Improvements
......@@ -52,6 +52,7 @@ fastgpt-plugin:
## 🐛 修复
1. 模型获取多模态文件链接异常。
2. 修复 training 接口存在的潜在越权风险。
## 🛠️ 代码优化
......
......@@ -64,7 +64,6 @@ export type CreateCollectionResponseType = z.infer<typeof CreateCollectionRespon
* Route: POST /core/dataset/collection/create/reTrainingCollection
* ============================================================================ */
export const ReTrainingCollectionBodySchema = DatasetCollectionStoreDataSchema.extend({
datasetId: z.string().meta({ description: '数据集 ID' }),
collectionId: z.string().meta({ description: '需要重新训练的集合 ID' })
});
export type ReTrainingCollectionBodyType = z.infer<typeof ReTrainingCollectionBodySchema>;
......
......@@ -9,10 +9,6 @@ import { PaginationSchema, PaginationResponseSchema } from '../../../api';
* Route: PUT /api/core/dataset/training/updateTrainingData
* ============================================================================ */
export const UpdateTrainingDataBodySchema = z.object({
datasetId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a05',
description: '知识库 ID'
}),
collectionId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a06',
description: '集合 ID'
......@@ -63,10 +59,6 @@ export type RebuildEmbeddingResponse = z.infer<typeof RebuildEmbeddingResponseSc
* Route: POST /api/core/dataset/training/deleteTrainingData
* ============================================================================ */
export const DeleteTrainingDataBodySchema = z.object({
datasetId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a05',
description: '知识库 ID'
}),
collectionId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a06',
description: '集合 ID'
......@@ -86,10 +78,6 @@ export type DeleteTrainingDataResponse = z.infer<typeof DeleteTrainingDataRespon
* Route: POST /api/core/dataset/training/getTrainingDataDetail
* ============================================================================ */
export const GetTrainingDataDetailBodySchema = z.object({
datasetId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a05',
description: '知识库 ID'
}),
collectionId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a06',
description: '集合 ID'
......
......@@ -162,6 +162,11 @@ export async function authDatasetCollection({
isRoot: isRootFromHeader
});
// collection 与 dataset 必须属于同一团队;否则说明对象归属已经损坏,不能继续按 datasetId 授权。
if (String(collection.teamId) !== String(dataset.teamId)) {
return Promise.reject(DatasetErrEnum.unAuthDataset);
}
return {
userId,
teamId,
......
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { DatasetErrEnum } from '@fastgpt/global/common/error/code/dataset';
import { ReadPermissionVal } from '@fastgpt/global/support/permission/constant';
const {
mockParseHeaderCert,
mockGetCollectionWithDataset,
mockFindDataset,
mockGetTmbInfoByTmbId,
mockGetTmbPermission
} = vi.hoisted(() => ({
mockParseHeaderCert: vi.fn(),
mockGetCollectionWithDataset: vi.fn(),
mockFindDataset: vi.fn(),
mockGetTmbInfoByTmbId: vi.fn(),
mockGetTmbPermission: vi.fn()
}));
vi.mock('@fastgpt/service/support/permission/auth/common', () => ({
parseHeaderCert: mockParseHeaderCert
}));
vi.mock('@fastgpt/service/core/dataset/controller', () => ({
getCollectionWithDataset: mockGetCollectionWithDataset
}));
vi.mock('@fastgpt/service/core/dataset/schema', () => ({
MongoDataset: {
findOne: mockFindDataset
}
}));
vi.mock('@fastgpt/service/support/user/team/controller', () => ({
getTmbInfoByTmbId: mockGetTmbInfoByTmbId
}));
vi.mock('@fastgpt/service/support/permission/controller', () => ({
getTmbPermission: mockGetTmbPermission
}));
vi.mock('@fastgpt/service/core/dataset/data/schema', () => ({
MongoDatasetData: {
findById: vi.fn()
}
}));
import { authDatasetCollection } from '@fastgpt/service/support/permission/dataset/auth';
const datasetId = '507f1f77bcf86cd799439011';
const collectionId = '507f1f77bcf86cd799439012';
const mockDatasetQuery = (dataset: Record<string, any>) => {
mockFindDataset.mockReturnValue({
lean: vi.fn().mockResolvedValue(dataset)
});
};
describe('authDatasetCollection', () => {
beforeEach(() => {
vi.clearAllMocks();
mockParseHeaderCert.mockResolvedValue({
teamId: 'team-a',
tmbId: 'tmb-a',
userId: 'user-a',
isRoot: false
});
mockGetTmbInfoByTmbId.mockResolvedValue({
teamId: 'team-a',
permission: { isOwner: true }
});
mockGetTmbPermission.mockResolvedValue(0);
mockDatasetQuery({
_id: datasetId,
teamId: 'team-a',
tmbId: 'tmb-a',
inheritPermission: false
});
});
it('rejects a collection whose team does not match its dataset team', async () => {
mockGetCollectionWithDataset.mockResolvedValue({
_id: collectionId,
teamId: 'team-b',
datasetId
});
await expect(
authDatasetCollection({
req: {} as any,
authToken: true,
collectionId,
per: ReadPermissionVal
})
).rejects.toBe(DatasetErrEnum.unAuthDataset);
});
it('allows a collection whose team matches its dataset team', async () => {
mockGetCollectionWithDataset.mockResolvedValue({
_id: collectionId,
teamId: 'team-a',
datasetId
});
const result = await authDatasetCollection({
req: {} as any,
authToken: true,
collectionId,
per: ReadPermissionVal
});
expect(result.collection._id).toBe(collectionId);
});
});
......@@ -287,11 +287,9 @@ const ProgressView = ({
};
const ErrorView = ({
datasetId,
collectionId,
refreshTrainingDetail
}: {
datasetId: string;
collectionId: string;
refreshTrainingDetail: () => void;
}) => {
......@@ -321,7 +319,7 @@ const ErrorView = ({
});
const { runAsync: getData, loading: getDataLoading } = useRequest(
(data: { datasetId: string; collectionId: string; dataId: string }) => {
(data: { collectionId: string; dataId: string }) => {
return getTrainingDataDetail(data);
},
{
......@@ -332,7 +330,7 @@ const ErrorView = ({
}
);
const { runAsync: deleteData, loading: deleteLoading } = useRequest(
(data: { datasetId: string; collectionId: string; dataId: string }) => {
(data: { collectionId: string; dataId: string }) => {
return deleteTrainingData(data);
},
{
......@@ -343,7 +341,7 @@ const ErrorView = ({
}
);
const { runAsync: updateData, loading: updateLoading } = useRequest(
(data: { datasetId: string; collectionId: string; dataId: string; q?: string; a?: string }) => {
(data: { collectionId: string; dataId: string; q?: string; a?: string }) => {
return updateTrainingData(data);
},
{
......@@ -364,7 +362,6 @@ const ErrorView = ({
onCancel={() => setEditChunk(undefined)}
onSave={(data) => {
updateData({
datasetId,
collectionId,
dataId: editChunk._id,
...data
......@@ -407,7 +404,7 @@ const ErrorView = ({
color={'myGray.600'}
leftIcon={<MyIcon name={'common/confirm/restoreTip'} w={4} />}
fontSize={'mini'}
onClick={() => updateData({ datasetId, collectionId, dataId: item._id })}
onClick={() => updateData({ collectionId, dataId: item._id })}
>
{t('dataset:dataset.ReTrain')}
</Button>
......@@ -418,7 +415,7 @@ const ErrorView = ({
color={'myGray.600'}
leftIcon={<MyIcon name={'edit'} w={4} />}
fontSize={'mini'}
onClick={() => getData({ datasetId, collectionId, dataId: item._id })}
onClick={() => getData({ collectionId, dataId: item._id })}
>
{t('dataset:dataset.Edit_Chunk')}
</Button>
......@@ -430,7 +427,7 @@ const ErrorView = ({
leftIcon={<MyIcon name={'delete'} w={4} />}
fontSize={'mini'}
onClick={() => {
deleteData({ datasetId, collectionId, dataId: item._id });
deleteData({ collectionId, dataId: item._id });
}}
>
{t('dataset:dataset.Delete_Chunk')}
......@@ -509,12 +506,10 @@ const EditView = ({
};
const TrainingStates = ({
datasetId,
collectionId,
defaultTab = 'states',
onClose
}: {
datasetId: string;
collectionId: string;
defaultTab?: 'states' | 'errors';
onClose: () => void;
......@@ -534,7 +529,7 @@ const TrainingStates = ({
// All retry logic
const { runAsync: handleRetryAll, loading: retrying } = useRequest(
() => updateTrainingData({ datasetId, collectionId }),
() => updateTrainingData({ collectionId }),
{
manual: true,
onSuccess: () => {
......@@ -581,7 +576,6 @@ const TrainingStates = ({
{tab === 'states' && trainingDetail && <ProgressView trainingDetail={trainingDetail} />}
{tab === 'errors' && (
<ErrorView
datasetId={datasetId}
collectionId={collectionId}
refreshTrainingDetail={refreshTrainingDetail}
/>
......
......@@ -485,7 +485,6 @@ const CollectionCard = () => {
{!!trainingStatesCollection && (
<TrainingStates
datasetId={datasetDetail._id}
collectionId={trainingStatesCollection.collectionId}
onClose={() => setTrainingStatesCollection(undefined)}
/>
......
......@@ -464,7 +464,6 @@ const DataCard = () => {
)}
{errorModalId && (
<TrainingStates
datasetId={datasetId}
defaultTab={'errors'}
collectionId={errorModalId}
onClose={() => {
......
......@@ -127,8 +127,14 @@ const Upload = () => {
};
if (importSource === ImportDataSourceEnum.reTraining) {
const reTrainingParams: Omit<typeof commonParams, 'datasetId'> & {
datasetId?: string;
} = {
...commonParams
};
delete reTrainingParams.datasetId;
const res = await postReTrainingDatasetFileCollection({
...commonParams,
...reTrainingParams,
collectionId
});
retrainNewCollectionId.current = res.collectionId;
......
......@@ -9,6 +9,7 @@ import { addAuditLog } from '@fastgpt/service/support/user/audit/util';
import { AuditEventEnum } from '@fastgpt/global/support/user/audit/constants';
import { getI18nDatasetType } from '@fastgpt/service/support/user/audit/util';
import { collectionTagsToTagLabel } from '@fastgpt/service/core/dataset/collection/utils';
import { parseApiInput } from '@fastgpt/service/common/zod/requestParseError';
import {
ReTrainingCollectionBodySchema,
ReTrainingCollectionResponseSchema,
......@@ -16,9 +17,10 @@ import {
} from '@fastgpt/global/openapi/core/dataset/collection/createApi';
async function handler(req: ApiRequestProps): Promise<ReTrainingCollectionResponseType> {
const { collectionId: inputCollectionId, ...data } = ReTrainingCollectionBodySchema.parse(
req.body
);
const { collectionId: inputCollectionId, ...data } = parseApiInput({
req,
bodySchema: ReTrainingCollectionBodySchema
}).body;
const { collection, teamId, tmbId } = await authDatasetCollection({
req,
......@@ -41,6 +43,9 @@ async function handler(req: ApiRequestProps): Promise<ReTrainingCollectionRespon
createCollectionParams: {
...collection,
...data,
datasetId: collection.datasetId,
teamId: collection.teamId,
tmbId: collection.tmbId,
parentId: collection.parentId ?? undefined,
updateTime: new Date(),
tags: await collectionTagsToTagLabel({
......
......@@ -11,12 +11,12 @@ import {
} from '@fastgpt/global/openapi/core/dataset/training/api';
async function handler(req: ApiRequestProps): Promise<DeleteTrainingDataResponse> {
const { datasetId, collectionId, dataId } = parseApiInput({
const { collectionId, dataId } = parseApiInput({
req,
bodySchema: DeleteTrainingDataBodySchema
}).body;
const { teamId } = await authDatasetCollection({
const { collection } = await authDatasetCollection({
req,
authToken: true,
authApiKey: true,
......@@ -25,8 +25,9 @@ async function handler(req: ApiRequestProps): Promise<DeleteTrainingDataResponse
});
await MongoDatasetTraining.deleteOne({
teamId,
datasetId,
teamId: collection.teamId,
datasetId: collection.datasetId,
collectionId: collection._id,
_id: dataId
});
......
......@@ -14,12 +14,12 @@ import { S3Buckets } from '@fastgpt/service/common/s3/config/constants';
import { parseApiInput } from '@fastgpt/service/common/zod/requestParseError';
async function handler(req: ApiRequestProps): Promise<GetTrainingDataDetailResponse> {
const { datasetId, collectionId, dataId } = parseApiInput({
const { collectionId, dataId } = parseApiInput({
req,
bodySchema: GetTrainingDataDetailBodySchema
}).body;
const { teamId } = await authDatasetCollection({
const { collection } = await authDatasetCollection({
req,
authToken: true,
authApiKey: true,
......@@ -27,7 +27,12 @@ async function handler(req: ApiRequestProps): Promise<GetTrainingDataDetailRespo
per: ReadPermissionVal
});
const data = await MongoDatasetTraining.findOne({ teamId, datasetId, _id: dataId }).lean();
const data = await MongoDatasetTraining.findOne({
teamId: collection.teamId,
datasetId: collection.datasetId,
collectionId: collection._id,
_id: dataId
}).lean();
if (!data) {
return GetTrainingDataDetailResponseSchema.parse(null);
......
......@@ -9,13 +9,15 @@ import {
UpdateTrainingDataResponseSchema,
type UpdateTrainingDataResponse
} from '@fastgpt/global/openapi/core/dataset/training/api';
import { parseApiInput } from '@fastgpt/service/common/zod/requestParseError';
async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse> {
const { datasetId, collectionId, dataId, q, a, chunkIndex } = UpdateTrainingDataBodySchema.parse(
req.body
);
const { collectionId, dataId, q, a, chunkIndex } = parseApiInput({
req,
bodySchema: UpdateTrainingDataBodySchema
}).body;
const { teamId } = await authDatasetCollection({
const { collection } = await authDatasetCollection({
req,
authToken: true,
authApiKey: true,
......@@ -23,13 +25,17 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
per: WritePermissionVal
});
const trainingMatch = {
teamId: collection.teamId,
datasetId: collection.datasetId,
collectionId: collection._id
};
// If dataId is not passed, all error data in this collection will be retried.
if (!dataId) {
await MongoDatasetTraining.updateMany(
{
teamId,
datasetId,
collectionId,
...trainingMatch,
errorMsg: { $exists: true, $ne: null }
},
{
......@@ -42,7 +48,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
}
// Single data retry logic
const data = await MongoDatasetTraining.findOne({ teamId, datasetId, _id: dataId });
const data = await MongoDatasetTraining.findOne({ ...trainingMatch, _id: dataId });
if (!data) {
return Promise.reject('data not found');
......@@ -52,8 +58,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
if (data.imageId && q) {
await MongoDatasetTraining.updateOne(
{
teamId,
datasetId,
...trainingMatch,
_id: dataId
},
{
......@@ -69,8 +74,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
} else {
await MongoDatasetTraining.updateOne(
{
teamId,
datasetId,
...trainingMatch,
_id: dataId
},
{
......
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { DatasetCollectionTypeEnum } from '@fastgpt/global/core/dataset/constants';
const {
mockAuthDatasetCollection,
mockCreateCollectionAndInsertData,
mockDelCollection,
mockMongoSessionRun,
mockCollectionTagsToTagLabel,
mockAddAuditLog
} = vi.hoisted(() => ({
mockAuthDatasetCollection: vi.fn(),
mockCreateCollectionAndInsertData: vi.fn(),
mockDelCollection: vi.fn(),
mockMongoSessionRun: vi.fn((fn: any) => fn('session')),
mockCollectionTagsToTagLabel: vi.fn(),
mockAddAuditLog: vi.fn()
}));
vi.mock('@/service/middleware/entry', () => ({
NextAPI: (handler: any) => handler
}));
vi.mock('@fastgpt/service/support/permission/dataset/auth', () => ({
authDatasetCollection: mockAuthDatasetCollection
}));
vi.mock('@fastgpt/service/core/dataset/collection/controller', () => ({
createCollectionAndInsertData: mockCreateCollectionAndInsertData,
delCollection: mockDelCollection
}));
vi.mock('@fastgpt/service/common/mongo/sessionRun', () => ({
mongoSessionRun: mockMongoSessionRun
}));
vi.mock('@fastgpt/service/core/dataset/collection/utils', () => ({
collectionTagsToTagLabel: mockCollectionTagsToTagLabel
}));
vi.mock('@fastgpt/service/support/user/audit/util', () => ({
addAuditLog: mockAddAuditLog,
getI18nDatasetType: vi.fn((type: string) => type)
}));
import handler from '@/pages/api/core/dataset/collection/create/reTrainingCollection';
const sourceDatasetId = '507f1f77bcf86cd799439011';
const sourceCollectionId = '507f1f77bcf86cd799439012';
const newCollectionId = '507f1f77bcf86cd799439013';
const foreignDatasetId = '507f1f77bcf86cd799439014';
const sourceCollection = {
_id: sourceCollectionId,
teamId: 'team-b',
tmbId: 'tmb-b',
datasetId: sourceDatasetId,
parentId: null,
name: 'source collection',
type: DatasetCollectionTypeEnum.file,
tags: ['tag-id'],
dataset: {
_id: sourceDatasetId,
teamId: 'team-b',
name: 'source dataset',
type: 'dataset',
vectorModel: 'text-embedding-3-small',
agentModel: 'gpt-4o-mini'
}
};
describe('reTrainingCollection', () => {
beforeEach(() => {
vi.clearAllMocks();
mockCollectionTagsToTagLabel.mockResolvedValue(['tag-label']);
mockCreateCollectionAndInsertData.mockResolvedValue({
collectionId: newCollectionId,
results: { insertLen: 0 }
});
mockAuthDatasetCollection.mockResolvedValue({
collection: sourceCollection,
teamId: 'team-b',
tmbId: 'tmb-b'
});
});
it('uses server-owned collection ownership fields when recreating a collection', async () => {
await handler({
body: {
collectionId: sourceCollectionId,
datasetId: foreignDatasetId,
chunkSize: 800
}
} as any);
expect(mockCreateCollectionAndInsertData).toHaveBeenCalledWith(
expect.objectContaining({
dataset: sourceCollection.dataset,
createCollectionParams: expect.objectContaining({
datasetId: sourceDatasetId,
teamId: 'team-b',
tmbId: 'tmb-b',
chunkSize: 800,
tags: ['tag-label']
})
})
);
});
it('ignores legacy datasetId in the request body', async () => {
await handler({
body: {
collectionId: sourceCollectionId,
datasetId: foreignDatasetId
}
} as any);
expect(mockCreateCollectionAndInsertData).toHaveBeenCalledWith(
expect.objectContaining({
dataset: sourceCollection.dataset,
createCollectionParams: expect.objectContaining({
datasetId: sourceDatasetId,
teamId: 'team-b',
tmbId: 'tmb-b'
})
})
);
});
});
......@@ -42,7 +42,6 @@ describe('delete training data test', () => {
const res = await Call<deleteTrainingDataBody, {}, deleteTrainingDataResponse>(handler, {
auth: root,
body: {
datasetId: dataset._id,
collectionId: collection._id,
dataId: trainingData._id
}
......@@ -57,4 +56,62 @@ describe('delete training data test', () => {
expect(res.code).toBe(200);
expect(deletedTrainingData).toBeNull();
});
it('should ignore legacy datasetId and only delete data from the authorized collection', async () => {
const root = await getRootUser();
const [dataset, foreignDataset] = await Promise.all([
MongoDataset.create({
name: 'test',
teamId: root.teamId,
tmbId: root.tmbId,
vectorModel: 'test',
agentModel: 'test'
}),
MongoDataset.create({
name: 'foreign',
teamId: root.teamId,
tmbId: root.tmbId,
vectorModel: 'test',
agentModel: 'test'
})
]);
const [collection, foreignCollection] = await Promise.all([
MongoDatasetCollection.create({
name: 'test',
type: DatasetCollectionTypeEnum.file,
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: dataset._id
}),
MongoDatasetCollection.create({
name: 'foreign',
type: DatasetCollectionTypeEnum.file,
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: foreignDataset._id
})
]);
const foreignTrainingData = await MongoDatasetTraining.create({
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: foreignDataset._id,
collectionId: foreignCollection._id,
billId: 'test',
mode: TrainingModeEnum.chunk
});
const res = await Call<deleteTrainingDataBody, {}, deleteTrainingDataResponse>(handler, {
auth: root,
body: {
datasetId: foreignDataset._id,
collectionId: collection._id,
dataId: foreignTrainingData._id
} as any
});
const existingTrainingData = await MongoDatasetTraining.findById(foreignTrainingData._id);
expect(res.code).toBe(200);
expect(existingTrainingData).toBeTruthy();
});
});
......@@ -44,7 +44,6 @@ describe('get training data detail test', () => {
const res = await Call<getTrainingDataDetailBody, {}, getTrainingDataDetailResponse>(handler, {
auth: root,
body: {
datasetId: dataset._id,
collectionId: collection._id,
dataId: trainingData._id
}
......@@ -58,4 +57,62 @@ describe('get training data detail test', () => {
expect(res.data?.q).toBe('test');
expect(res.data?.a).toBe('test');
});
it('should ignore legacy datasetId and only read data from the authorized collection', async () => {
const root = await getRootUser();
const [dataset, foreignDataset] = await Promise.all([
MongoDataset.create({
name: 'test',
teamId: root.teamId,
tmbId: root.tmbId,
vectorModel: 'test',
agentModel: 'test'
}),
MongoDataset.create({
name: 'foreign',
teamId: root.teamId,
tmbId: root.tmbId,
vectorModel: 'test',
agentModel: 'test'
})
]);
const [collection, foreignCollection] = await Promise.all([
MongoDatasetCollection.create({
name: 'test',
type: DatasetCollectionTypeEnum.file,
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: dataset._id
}),
MongoDatasetCollection.create({
name: 'foreign',
type: DatasetCollectionTypeEnum.file,
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: foreignDataset._id
})
]);
const foreignTrainingData = await MongoDatasetTraining.create({
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: foreignDataset._id,
collectionId: foreignCollection._id,
billId: 'test',
mode: TrainingModeEnum.chunk,
q: 'foreign',
a: 'foreign'
});
const res = await Call<getTrainingDataDetailBody, {}, getTrainingDataDetailResponse>(handler, {
auth: root,
body: {
datasetId: foreignDataset._id,
collectionId: collection._id,
dataId: foreignTrainingData._id
} as any
});
expect(res.code).toBe(200);
expect(res.data).toBeNull();
});
});
......@@ -42,7 +42,6 @@ describe('update training data test', () => {
const res = await Call<updateTrainingDataBody, {}, updateTrainingDataResponse>(handler, {
auth: root,
body: {
datasetId: dataset._id,
collectionId: collection._id,
dataId: trainingData._id,
q: 'test',
......@@ -62,4 +61,67 @@ describe('update training data test', () => {
expect(updatedTrainingData?.a).toBe('test');
expect(updatedTrainingData?.chunkIndex).toBe(1);
});
it('should ignore legacy datasetId and only update data from the authorized collection', async () => {
const root = await getRootUser();
const [dataset, foreignDataset] = await Promise.all([
MongoDataset.create({
name: 'test',
teamId: root.teamId,
tmbId: root.tmbId,
vectorModel: 'test',
agentModel: 'test'
}),
MongoDataset.create({
name: 'foreign',
teamId: root.teamId,
tmbId: root.tmbId,
vectorModel: 'test',
agentModel: 'test'
})
]);
const [collection, foreignCollection] = await Promise.all([
MongoDatasetCollection.create({
name: 'test',
type: DatasetCollectionTypeEnum.file,
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: dataset._id
}),
MongoDatasetCollection.create({
name: 'foreign',
type: DatasetCollectionTypeEnum.file,
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: foreignDataset._id
})
]);
const foreignTrainingData = await MongoDatasetTraining.create({
teamId: root.teamId,
tmbId: root.tmbId,
datasetId: foreignDataset._id,
collectionId: foreignCollection._id,
billId: 'test',
mode: TrainingModeEnum.chunk,
q: 'origin',
a: 'origin'
});
const res = await Call<updateTrainingDataBody, {}, updateTrainingDataResponse>(handler, {
auth: root,
body: {
datasetId: foreignDataset._id,
collectionId: collection._id,
dataId: foreignTrainingData._id,
q: 'changed',
a: 'changed'
} as any
});
const existingTrainingData = await MongoDatasetTraining.findById(foreignTrainingData._id);
expect(res.code).not.toBe(200);
expect(existingTrainingData?.q).toBe('origin');
expect(existingTrainingData?.a).toBe('origin');
});
});
import { describe, expect, it, vi } from 'vitest';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { handler } from '@/pages/api/core/dataset/training/updateTrainingData';
import { MongoDatasetTraining } from '@fastgpt/service/core/dataset/training/schema';
import { authDatasetCollection } from '@fastgpt/service/support/permission/dataset/auth';
......@@ -7,6 +7,7 @@ import { TrainingModeEnum } from '@fastgpt/global/core/dataset/constants';
const datasetId = '507f1f77bcf86cd799439011';
const collectionId = '507f1f77bcf86cd799439012';
const dataId = '507f1f77bcf86cd799439013';
const foreignDatasetId = '507f1f77bcf86cd799439014';
vi.mock('@fastgpt/service/core/dataset/training/schema', () => ({
MongoDatasetTraining: {
......@@ -21,14 +22,20 @@ vi.mock('@fastgpt/service/support/permission/dataset/auth', () => ({
}));
describe('updateTrainingData', () => {
it('should retry all error data when dataId is not provided', async () => {
beforeEach(() => {
vi.clearAllMocks();
vi.mocked(authDatasetCollection).mockResolvedValue({
teamId: 'team1'
collection: {
_id: collectionId,
teamId: 'team1',
datasetId
}
});
});
it('should retry all error data when dataId is not provided', async () => {
const req = {
body: {
datasetId,
collectionId
}
};
......@@ -51,17 +58,12 @@ describe('updateTrainingData', () => {
});
it('should update single training data with image', async () => {
vi.mocked(authDatasetCollection).mockResolvedValue({
teamId: 'team1'
});
vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue({
imageId: 'image1'
});
const req = {
body: {
datasetId,
collectionId,
dataId,
q: 'question',
......@@ -76,6 +78,7 @@ describe('updateTrainingData', () => {
{
teamId: 'team1',
datasetId,
collectionId,
_id: dataId
},
{
......@@ -91,15 +94,10 @@ describe('updateTrainingData', () => {
});
it('should update single training data without image', async () => {
vi.mocked(authDatasetCollection).mockResolvedValue({
teamId: 'team1'
});
vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue({});
const req = {
body: {
datasetId,
collectionId,
dataId,
q: 'question',
......@@ -114,6 +112,7 @@ describe('updateTrainingData', () => {
{
teamId: 'team1',
datasetId,
collectionId,
_id: dataId
},
{
......@@ -128,15 +127,10 @@ describe('updateTrainingData', () => {
});
it('should reject when data not found', async () => {
vi.mocked(authDatasetCollection).mockResolvedValue({
teamId: 'team1'
});
vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue(null);
const req = {
body: {
datasetId,
collectionId,
dataId
}
......@@ -144,4 +138,37 @@ describe('updateTrainingData', () => {
await expect(handler(req as any)).rejects.toBe('data not found');
});
it('should ignore legacy request datasetId and use the authorized collection datasetId', async () => {
vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue({});
const req = {
body: {
datasetId: foreignDatasetId,
collectionId,
dataId,
q: 'question'
}
};
await handler(req as any);
expect(MongoDatasetTraining.findOne).toHaveBeenCalledWith({
teamId: 'team1',
datasetId,
collectionId,
_id: dataId
});
expect(MongoDatasetTraining.updateOne).toHaveBeenCalledWith(
{
teamId: 'team1',
datasetId,
collectionId,
_id: dataId
},
expect.objectContaining({
q: 'question'
})
);
});
});
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment