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' ...@@ -14,4 +14,6 @@ description: 'FastGPT V4.15.0-beta4 Release Notes'
## 🐛 Bug Fixes ## 🐛 Bug Fixes
1. Fixed a potential unauthorized access risk in training APIs.
## 🛠️ Code Improvements ## 🛠️ Code Improvements
...@@ -52,6 +52,7 @@ fastgpt-plugin: ...@@ -52,6 +52,7 @@ fastgpt-plugin:
## 🐛 修复 ## 🐛 修复
1. 模型获取多模态文件链接异常。 1. 模型获取多模态文件链接异常。
2. 修复 training 接口存在的潜在越权风险。
## 🛠️ 代码优化 ## 🛠️ 代码优化
......
...@@ -64,7 +64,6 @@ export type CreateCollectionResponseType = z.infer<typeof CreateCollectionRespon ...@@ -64,7 +64,6 @@ export type CreateCollectionResponseType = z.infer<typeof CreateCollectionRespon
* Route: POST /core/dataset/collection/create/reTrainingCollection * Route: POST /core/dataset/collection/create/reTrainingCollection
* ============================================================================ */ * ============================================================================ */
export const ReTrainingCollectionBodySchema = DatasetCollectionStoreDataSchema.extend({ export const ReTrainingCollectionBodySchema = DatasetCollectionStoreDataSchema.extend({
datasetId: z.string().meta({ description: '数据集 ID' }),
collectionId: z.string().meta({ description: '需要重新训练的集合 ID' }) collectionId: z.string().meta({ description: '需要重新训练的集合 ID' })
}); });
export type ReTrainingCollectionBodyType = z.infer<typeof ReTrainingCollectionBodySchema>; export type ReTrainingCollectionBodyType = z.infer<typeof ReTrainingCollectionBodySchema>;
......
...@@ -9,10 +9,6 @@ import { PaginationSchema, PaginationResponseSchema } from '../../../api'; ...@@ -9,10 +9,6 @@ import { PaginationSchema, PaginationResponseSchema } from '../../../api';
* Route: PUT /api/core/dataset/training/updateTrainingData * Route: PUT /api/core/dataset/training/updateTrainingData
* ============================================================================ */ * ============================================================================ */
export const UpdateTrainingDataBodySchema = z.object({ export const UpdateTrainingDataBodySchema = z.object({
datasetId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a05',
description: '知识库 ID'
}),
collectionId: ObjectIdSchema.meta({ collectionId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a06', example: '68ad85a7463006c963799a06',
description: '集合 ID' description: '集合 ID'
...@@ -63,10 +59,6 @@ export type RebuildEmbeddingResponse = z.infer<typeof RebuildEmbeddingResponseSc ...@@ -63,10 +59,6 @@ export type RebuildEmbeddingResponse = z.infer<typeof RebuildEmbeddingResponseSc
* Route: POST /api/core/dataset/training/deleteTrainingData * Route: POST /api/core/dataset/training/deleteTrainingData
* ============================================================================ */ * ============================================================================ */
export const DeleteTrainingDataBodySchema = z.object({ export const DeleteTrainingDataBodySchema = z.object({
datasetId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a05',
description: '知识库 ID'
}),
collectionId: ObjectIdSchema.meta({ collectionId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a06', example: '68ad85a7463006c963799a06',
description: '集合 ID' description: '集合 ID'
...@@ -86,10 +78,6 @@ export type DeleteTrainingDataResponse = z.infer<typeof DeleteTrainingDataRespon ...@@ -86,10 +78,6 @@ export type DeleteTrainingDataResponse = z.infer<typeof DeleteTrainingDataRespon
* Route: POST /api/core/dataset/training/getTrainingDataDetail * Route: POST /api/core/dataset/training/getTrainingDataDetail
* ============================================================================ */ * ============================================================================ */
export const GetTrainingDataDetailBodySchema = z.object({ export const GetTrainingDataDetailBodySchema = z.object({
datasetId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a05',
description: '知识库 ID'
}),
collectionId: ObjectIdSchema.meta({ collectionId: ObjectIdSchema.meta({
example: '68ad85a7463006c963799a06', example: '68ad85a7463006c963799a06',
description: '集合 ID' description: '集合 ID'
......
...@@ -162,6 +162,11 @@ export async function authDatasetCollection({ ...@@ -162,6 +162,11 @@ export async function authDatasetCollection({
isRoot: isRootFromHeader isRoot: isRootFromHeader
}); });
// collection 与 dataset 必须属于同一团队;否则说明对象归属已经损坏,不能继续按 datasetId 授权。
if (String(collection.teamId) !== String(dataset.teamId)) {
return Promise.reject(DatasetErrEnum.unAuthDataset);
}
return { return {
userId, userId,
teamId, 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 = ({ ...@@ -287,11 +287,9 @@ const ProgressView = ({
}; };
const ErrorView = ({ const ErrorView = ({
datasetId,
collectionId, collectionId,
refreshTrainingDetail refreshTrainingDetail
}: { }: {
datasetId: string;
collectionId: string; collectionId: string;
refreshTrainingDetail: () => void; refreshTrainingDetail: () => void;
}) => { }) => {
...@@ -321,7 +319,7 @@ const ErrorView = ({ ...@@ -321,7 +319,7 @@ const ErrorView = ({
}); });
const { runAsync: getData, loading: getDataLoading } = useRequest( const { runAsync: getData, loading: getDataLoading } = useRequest(
(data: { datasetId: string; collectionId: string; dataId: string }) => { (data: { collectionId: string; dataId: string }) => {
return getTrainingDataDetail(data); return getTrainingDataDetail(data);
}, },
{ {
...@@ -332,7 +330,7 @@ const ErrorView = ({ ...@@ -332,7 +330,7 @@ const ErrorView = ({
} }
); );
const { runAsync: deleteData, loading: deleteLoading } = useRequest( const { runAsync: deleteData, loading: deleteLoading } = useRequest(
(data: { datasetId: string; collectionId: string; dataId: string }) => { (data: { collectionId: string; dataId: string }) => {
return deleteTrainingData(data); return deleteTrainingData(data);
}, },
{ {
...@@ -343,7 +341,7 @@ const ErrorView = ({ ...@@ -343,7 +341,7 @@ const ErrorView = ({
} }
); );
const { runAsync: updateData, loading: updateLoading } = useRequest( 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); return updateTrainingData(data);
}, },
{ {
...@@ -364,7 +362,6 @@ const ErrorView = ({ ...@@ -364,7 +362,6 @@ const ErrorView = ({
onCancel={() => setEditChunk(undefined)} onCancel={() => setEditChunk(undefined)}
onSave={(data) => { onSave={(data) => {
updateData({ updateData({
datasetId,
collectionId, collectionId,
dataId: editChunk._id, dataId: editChunk._id,
...data ...data
...@@ -407,7 +404,7 @@ const ErrorView = ({ ...@@ -407,7 +404,7 @@ const ErrorView = ({
color={'myGray.600'} color={'myGray.600'}
leftIcon={<MyIcon name={'common/confirm/restoreTip'} w={4} />} leftIcon={<MyIcon name={'common/confirm/restoreTip'} w={4} />}
fontSize={'mini'} fontSize={'mini'}
onClick={() => updateData({ datasetId, collectionId, dataId: item._id })} onClick={() => updateData({ collectionId, dataId: item._id })}
> >
{t('dataset:dataset.ReTrain')} {t('dataset:dataset.ReTrain')}
</Button> </Button>
...@@ -418,7 +415,7 @@ const ErrorView = ({ ...@@ -418,7 +415,7 @@ const ErrorView = ({
color={'myGray.600'} color={'myGray.600'}
leftIcon={<MyIcon name={'edit'} w={4} />} leftIcon={<MyIcon name={'edit'} w={4} />}
fontSize={'mini'} fontSize={'mini'}
onClick={() => getData({ datasetId, collectionId, dataId: item._id })} onClick={() => getData({ collectionId, dataId: item._id })}
> >
{t('dataset:dataset.Edit_Chunk')} {t('dataset:dataset.Edit_Chunk')}
</Button> </Button>
...@@ -430,7 +427,7 @@ const ErrorView = ({ ...@@ -430,7 +427,7 @@ const ErrorView = ({
leftIcon={<MyIcon name={'delete'} w={4} />} leftIcon={<MyIcon name={'delete'} w={4} />}
fontSize={'mini'} fontSize={'mini'}
onClick={() => { onClick={() => {
deleteData({ datasetId, collectionId, dataId: item._id }); deleteData({ collectionId, dataId: item._id });
}} }}
> >
{t('dataset:dataset.Delete_Chunk')} {t('dataset:dataset.Delete_Chunk')}
...@@ -509,12 +506,10 @@ const EditView = ({ ...@@ -509,12 +506,10 @@ const EditView = ({
}; };
const TrainingStates = ({ const TrainingStates = ({
datasetId,
collectionId, collectionId,
defaultTab = 'states', defaultTab = 'states',
onClose onClose
}: { }: {
datasetId: string;
collectionId: string; collectionId: string;
defaultTab?: 'states' | 'errors'; defaultTab?: 'states' | 'errors';
onClose: () => void; onClose: () => void;
...@@ -534,7 +529,7 @@ const TrainingStates = ({ ...@@ -534,7 +529,7 @@ const TrainingStates = ({
// All retry logic // All retry logic
const { runAsync: handleRetryAll, loading: retrying } = useRequest( const { runAsync: handleRetryAll, loading: retrying } = useRequest(
() => updateTrainingData({ datasetId, collectionId }), () => updateTrainingData({ collectionId }),
{ {
manual: true, manual: true,
onSuccess: () => { onSuccess: () => {
...@@ -581,7 +576,6 @@ const TrainingStates = ({ ...@@ -581,7 +576,6 @@ const TrainingStates = ({
{tab === 'states' && trainingDetail && <ProgressView trainingDetail={trainingDetail} />} {tab === 'states' && trainingDetail && <ProgressView trainingDetail={trainingDetail} />}
{tab === 'errors' && ( {tab === 'errors' && (
<ErrorView <ErrorView
datasetId={datasetId}
collectionId={collectionId} collectionId={collectionId}
refreshTrainingDetail={refreshTrainingDetail} refreshTrainingDetail={refreshTrainingDetail}
/> />
......
...@@ -485,7 +485,6 @@ const CollectionCard = () => { ...@@ -485,7 +485,6 @@ const CollectionCard = () => {
{!!trainingStatesCollection && ( {!!trainingStatesCollection && (
<TrainingStates <TrainingStates
datasetId={datasetDetail._id}
collectionId={trainingStatesCollection.collectionId} collectionId={trainingStatesCollection.collectionId}
onClose={() => setTrainingStatesCollection(undefined)} onClose={() => setTrainingStatesCollection(undefined)}
/> />
......
...@@ -464,7 +464,6 @@ const DataCard = () => { ...@@ -464,7 +464,6 @@ const DataCard = () => {
)} )}
{errorModalId && ( {errorModalId && (
<TrainingStates <TrainingStates
datasetId={datasetId}
defaultTab={'errors'} defaultTab={'errors'}
collectionId={errorModalId} collectionId={errorModalId}
onClose={() => { onClose={() => {
......
...@@ -127,8 +127,14 @@ const Upload = () => { ...@@ -127,8 +127,14 @@ const Upload = () => {
}; };
if (importSource === ImportDataSourceEnum.reTraining) { if (importSource === ImportDataSourceEnum.reTraining) {
const reTrainingParams: Omit<typeof commonParams, 'datasetId'> & {
datasetId?: string;
} = {
...commonParams
};
delete reTrainingParams.datasetId;
const res = await postReTrainingDatasetFileCollection({ const res = await postReTrainingDatasetFileCollection({
...commonParams, ...reTrainingParams,
collectionId collectionId
}); });
retrainNewCollectionId.current = res.collectionId; retrainNewCollectionId.current = res.collectionId;
......
...@@ -9,6 +9,7 @@ import { addAuditLog } from '@fastgpt/service/support/user/audit/util'; ...@@ -9,6 +9,7 @@ import { addAuditLog } from '@fastgpt/service/support/user/audit/util';
import { AuditEventEnum } from '@fastgpt/global/support/user/audit/constants'; import { AuditEventEnum } from '@fastgpt/global/support/user/audit/constants';
import { getI18nDatasetType } from '@fastgpt/service/support/user/audit/util'; import { getI18nDatasetType } from '@fastgpt/service/support/user/audit/util';
import { collectionTagsToTagLabel } from '@fastgpt/service/core/dataset/collection/utils'; import { collectionTagsToTagLabel } from '@fastgpt/service/core/dataset/collection/utils';
import { parseApiInput } from '@fastgpt/service/common/zod/requestParseError';
import { import {
ReTrainingCollectionBodySchema, ReTrainingCollectionBodySchema,
ReTrainingCollectionResponseSchema, ReTrainingCollectionResponseSchema,
...@@ -16,9 +17,10 @@ import { ...@@ -16,9 +17,10 @@ import {
} from '@fastgpt/global/openapi/core/dataset/collection/createApi'; } from '@fastgpt/global/openapi/core/dataset/collection/createApi';
async function handler(req: ApiRequestProps): Promise<ReTrainingCollectionResponseType> { async function handler(req: ApiRequestProps): Promise<ReTrainingCollectionResponseType> {
const { collectionId: inputCollectionId, ...data } = ReTrainingCollectionBodySchema.parse( const { collectionId: inputCollectionId, ...data } = parseApiInput({
req.body req,
); bodySchema: ReTrainingCollectionBodySchema
}).body;
const { collection, teamId, tmbId } = await authDatasetCollection({ const { collection, teamId, tmbId } = await authDatasetCollection({
req, req,
...@@ -41,6 +43,9 @@ async function handler(req: ApiRequestProps): Promise<ReTrainingCollectionRespon ...@@ -41,6 +43,9 @@ async function handler(req: ApiRequestProps): Promise<ReTrainingCollectionRespon
createCollectionParams: { createCollectionParams: {
...collection, ...collection,
...data, ...data,
datasetId: collection.datasetId,
teamId: collection.teamId,
tmbId: collection.tmbId,
parentId: collection.parentId ?? undefined, parentId: collection.parentId ?? undefined,
updateTime: new Date(), updateTime: new Date(),
tags: await collectionTagsToTagLabel({ tags: await collectionTagsToTagLabel({
......
...@@ -11,12 +11,12 @@ import { ...@@ -11,12 +11,12 @@ import {
} from '@fastgpt/global/openapi/core/dataset/training/api'; } from '@fastgpt/global/openapi/core/dataset/training/api';
async function handler(req: ApiRequestProps): Promise<DeleteTrainingDataResponse> { async function handler(req: ApiRequestProps): Promise<DeleteTrainingDataResponse> {
const { datasetId, collectionId, dataId } = parseApiInput({ const { collectionId, dataId } = parseApiInput({
req, req,
bodySchema: DeleteTrainingDataBodySchema bodySchema: DeleteTrainingDataBodySchema
}).body; }).body;
const { teamId } = await authDatasetCollection({ const { collection } = await authDatasetCollection({
req, req,
authToken: true, authToken: true,
authApiKey: true, authApiKey: true,
...@@ -25,8 +25,9 @@ async function handler(req: ApiRequestProps): Promise<DeleteTrainingDataResponse ...@@ -25,8 +25,9 @@ async function handler(req: ApiRequestProps): Promise<DeleteTrainingDataResponse
}); });
await MongoDatasetTraining.deleteOne({ await MongoDatasetTraining.deleteOne({
teamId, teamId: collection.teamId,
datasetId, datasetId: collection.datasetId,
collectionId: collection._id,
_id: dataId _id: dataId
}); });
......
...@@ -14,12 +14,12 @@ import { S3Buckets } from '@fastgpt/service/common/s3/config/constants'; ...@@ -14,12 +14,12 @@ import { S3Buckets } from '@fastgpt/service/common/s3/config/constants';
import { parseApiInput } from '@fastgpt/service/common/zod/requestParseError'; import { parseApiInput } from '@fastgpt/service/common/zod/requestParseError';
async function handler(req: ApiRequestProps): Promise<GetTrainingDataDetailResponse> { async function handler(req: ApiRequestProps): Promise<GetTrainingDataDetailResponse> {
const { datasetId, collectionId, dataId } = parseApiInput({ const { collectionId, dataId } = parseApiInput({
req, req,
bodySchema: GetTrainingDataDetailBodySchema bodySchema: GetTrainingDataDetailBodySchema
}).body; }).body;
const { teamId } = await authDatasetCollection({ const { collection } = await authDatasetCollection({
req, req,
authToken: true, authToken: true,
authApiKey: true, authApiKey: true,
...@@ -27,7 +27,12 @@ async function handler(req: ApiRequestProps): Promise<GetTrainingDataDetailRespo ...@@ -27,7 +27,12 @@ async function handler(req: ApiRequestProps): Promise<GetTrainingDataDetailRespo
per: ReadPermissionVal 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) { if (!data) {
return GetTrainingDataDetailResponseSchema.parse(null); return GetTrainingDataDetailResponseSchema.parse(null);
......
...@@ -9,13 +9,15 @@ import { ...@@ -9,13 +9,15 @@ import {
UpdateTrainingDataResponseSchema, UpdateTrainingDataResponseSchema,
type UpdateTrainingDataResponse type UpdateTrainingDataResponse
} from '@fastgpt/global/openapi/core/dataset/training/api'; } from '@fastgpt/global/openapi/core/dataset/training/api';
import { parseApiInput } from '@fastgpt/service/common/zod/requestParseError';
async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse> { async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse> {
const { datasetId, collectionId, dataId, q, a, chunkIndex } = UpdateTrainingDataBodySchema.parse( const { collectionId, dataId, q, a, chunkIndex } = parseApiInput({
req.body req,
); bodySchema: UpdateTrainingDataBodySchema
}).body;
const { teamId } = await authDatasetCollection({ const { collection } = await authDatasetCollection({
req, req,
authToken: true, authToken: true,
authApiKey: true, authApiKey: true,
...@@ -23,13 +25,17 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse ...@@ -23,13 +25,17 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
per: WritePermissionVal 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 is not passed, all error data in this collection will be retried.
if (!dataId) { if (!dataId) {
await MongoDatasetTraining.updateMany( await MongoDatasetTraining.updateMany(
{ {
teamId, ...trainingMatch,
datasetId,
collectionId,
errorMsg: { $exists: true, $ne: null } errorMsg: { $exists: true, $ne: null }
}, },
{ {
...@@ -42,7 +48,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse ...@@ -42,7 +48,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
} }
// Single data retry logic // Single data retry logic
const data = await MongoDatasetTraining.findOne({ teamId, datasetId, _id: dataId }); const data = await MongoDatasetTraining.findOne({ ...trainingMatch, _id: dataId });
if (!data) { if (!data) {
return Promise.reject('data not found'); return Promise.reject('data not found');
...@@ -52,8 +58,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse ...@@ -52,8 +58,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
if (data.imageId && q) { if (data.imageId && q) {
await MongoDatasetTraining.updateOne( await MongoDatasetTraining.updateOne(
{ {
teamId, ...trainingMatch,
datasetId,
_id: dataId _id: dataId
}, },
{ {
...@@ -69,8 +74,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse ...@@ -69,8 +74,7 @@ async function handler(req: ApiRequestProps): Promise<UpdateTrainingDataResponse
} else { } else {
await MongoDatasetTraining.updateOne( await MongoDatasetTraining.updateOne(
{ {
teamId, ...trainingMatch,
datasetId,
_id: dataId _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', () => { ...@@ -42,7 +42,6 @@ describe('delete training data test', () => {
const res = await Call<deleteTrainingDataBody, {}, deleteTrainingDataResponse>(handler, { const res = await Call<deleteTrainingDataBody, {}, deleteTrainingDataResponse>(handler, {
auth: root, auth: root,
body: { body: {
datasetId: dataset._id,
collectionId: collection._id, collectionId: collection._id,
dataId: trainingData._id dataId: trainingData._id
} }
...@@ -57,4 +56,62 @@ describe('delete training data test', () => { ...@@ -57,4 +56,62 @@ describe('delete training data test', () => {
expect(res.code).toBe(200); expect(res.code).toBe(200);
expect(deletedTrainingData).toBeNull(); 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', () => { ...@@ -44,7 +44,6 @@ describe('get training data detail test', () => {
const res = await Call<getTrainingDataDetailBody, {}, getTrainingDataDetailResponse>(handler, { const res = await Call<getTrainingDataDetailBody, {}, getTrainingDataDetailResponse>(handler, {
auth: root, auth: root,
body: { body: {
datasetId: dataset._id,
collectionId: collection._id, collectionId: collection._id,
dataId: trainingData._id dataId: trainingData._id
} }
...@@ -58,4 +57,62 @@ describe('get training data detail test', () => { ...@@ -58,4 +57,62 @@ describe('get training data detail test', () => {
expect(res.data?.q).toBe('test'); expect(res.data?.q).toBe('test');
expect(res.data?.a).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', () => { ...@@ -42,7 +42,6 @@ describe('update training data test', () => {
const res = await Call<updateTrainingDataBody, {}, updateTrainingDataResponse>(handler, { const res = await Call<updateTrainingDataBody, {}, updateTrainingDataResponse>(handler, {
auth: root, auth: root,
body: { body: {
datasetId: dataset._id,
collectionId: collection._id, collectionId: collection._id,
dataId: trainingData._id, dataId: trainingData._id,
q: 'test', q: 'test',
...@@ -62,4 +61,67 @@ describe('update training data test', () => { ...@@ -62,4 +61,67 @@ describe('update training data test', () => {
expect(updatedTrainingData?.a).toBe('test'); expect(updatedTrainingData?.a).toBe('test');
expect(updatedTrainingData?.chunkIndex).toBe(1); 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 { handler } from '@/pages/api/core/dataset/training/updateTrainingData';
import { MongoDatasetTraining } from '@fastgpt/service/core/dataset/training/schema'; import { MongoDatasetTraining } from '@fastgpt/service/core/dataset/training/schema';
import { authDatasetCollection } from '@fastgpt/service/support/permission/dataset/auth'; import { authDatasetCollection } from '@fastgpt/service/support/permission/dataset/auth';
...@@ -7,6 +7,7 @@ import { TrainingModeEnum } from '@fastgpt/global/core/dataset/constants'; ...@@ -7,6 +7,7 @@ import { TrainingModeEnum } from '@fastgpt/global/core/dataset/constants';
const datasetId = '507f1f77bcf86cd799439011'; const datasetId = '507f1f77bcf86cd799439011';
const collectionId = '507f1f77bcf86cd799439012'; const collectionId = '507f1f77bcf86cd799439012';
const dataId = '507f1f77bcf86cd799439013'; const dataId = '507f1f77bcf86cd799439013';
const foreignDatasetId = '507f1f77bcf86cd799439014';
vi.mock('@fastgpt/service/core/dataset/training/schema', () => ({ vi.mock('@fastgpt/service/core/dataset/training/schema', () => ({
MongoDatasetTraining: { MongoDatasetTraining: {
...@@ -21,14 +22,20 @@ vi.mock('@fastgpt/service/support/permission/dataset/auth', () => ({ ...@@ -21,14 +22,20 @@ vi.mock('@fastgpt/service/support/permission/dataset/auth', () => ({
})); }));
describe('updateTrainingData', () => { describe('updateTrainingData', () => {
it('should retry all error data when dataId is not provided', async () => { beforeEach(() => {
vi.clearAllMocks();
vi.mocked(authDatasetCollection).mockResolvedValue({ 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 = { const req = {
body: { body: {
datasetId,
collectionId collectionId
} }
}; };
...@@ -51,17 +58,12 @@ describe('updateTrainingData', () => { ...@@ -51,17 +58,12 @@ describe('updateTrainingData', () => {
}); });
it('should update single training data with image', async () => { it('should update single training data with image', async () => {
vi.mocked(authDatasetCollection).mockResolvedValue({
teamId: 'team1'
});
vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue({ vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue({
imageId: 'image1' imageId: 'image1'
}); });
const req = { const req = {
body: { body: {
datasetId,
collectionId, collectionId,
dataId, dataId,
q: 'question', q: 'question',
...@@ -76,6 +78,7 @@ describe('updateTrainingData', () => { ...@@ -76,6 +78,7 @@ describe('updateTrainingData', () => {
{ {
teamId: 'team1', teamId: 'team1',
datasetId, datasetId,
collectionId,
_id: dataId _id: dataId
}, },
{ {
...@@ -91,15 +94,10 @@ describe('updateTrainingData', () => { ...@@ -91,15 +94,10 @@ describe('updateTrainingData', () => {
}); });
it('should update single training data without image', async () => { it('should update single training data without image', async () => {
vi.mocked(authDatasetCollection).mockResolvedValue({
teamId: 'team1'
});
vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue({}); vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue({});
const req = { const req = {
body: { body: {
datasetId,
collectionId, collectionId,
dataId, dataId,
q: 'question', q: 'question',
...@@ -114,6 +112,7 @@ describe('updateTrainingData', () => { ...@@ -114,6 +112,7 @@ describe('updateTrainingData', () => {
{ {
teamId: 'team1', teamId: 'team1',
datasetId, datasetId,
collectionId,
_id: dataId _id: dataId
}, },
{ {
...@@ -128,15 +127,10 @@ describe('updateTrainingData', () => { ...@@ -128,15 +127,10 @@ describe('updateTrainingData', () => {
}); });
it('should reject when data not found', async () => { it('should reject when data not found', async () => {
vi.mocked(authDatasetCollection).mockResolvedValue({
teamId: 'team1'
});
vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue(null); vi.mocked(MongoDatasetTraining.findOne).mockResolvedValue(null);
const req = { const req = {
body: { body: {
datasetId,
collectionId, collectionId,
dataId dataId
} }
...@@ -144,4 +138,37 @@ describe('updateTrainingData', () => { ...@@ -144,4 +138,37 @@ describe('updateTrainingData', () => {
await expect(handler(req as any)).rejects.toBe('data not found'); 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