diff --git a/api/server/routes/categories.js b/api/server/routes/categories.js index 612bc378603..afd2f17c210 100644 --- a/api/server/routes/categories.js +++ b/api/server/routes/categories.js @@ -1,15 +1,31 @@ const express = require('express'); +const { PermissionTypes, Permissions } = require('librechat-data-provider'); +const { checkAccess, createGetPromptCategoriesHandler } = require('@librechat/api'); +const { requireJwtAuth, configMiddleware } = require('~/server/middleware'); +const { + getRoleByName, + getPromptGroupAccessContext, + getDistinctPromptGroupCategories, +} = require('~/models'); + const router = express.Router(); -const { requireJwtAuth } = require('~/server/middleware'); -const { getCategories } = require('~/models'); -router.get('/', requireJwtAuth, async (req, res) => { - try { - const categories = await getCategories(); - res.status(200).send(categories); - } catch (error) { - res.status(500).send({ message: 'Failed to retrieve categories', error: error.message }); - } -}); +router.get( + '/', + requireJwtAuth, + configMiddleware, + createGetPromptCategoriesHandler({ + getPromptGroupAccessContext, + getDistinctPromptGroupCategories, + canUsePrompts: (req) => + checkAccess({ + req, + user: req.user, + permissionType: PermissionTypes.PROMPTS, + permissions: [Permissions.USE], + getRoleByName, + }), + }), +); module.exports = router; diff --git a/api/server/routes/categories.test.js b/api/server/routes/categories.test.js new file mode 100644 index 00000000000..6654db995ae --- /dev/null +++ b/api/server/routes/categories.test.js @@ -0,0 +1,217 @@ +const express = require('express'); +const request = require('supertest'); +const mongoose = require('mongoose'); +const { MongoMemoryServer } = require('mongodb-memory-server'); +const { + SystemRoles, + ResourceType, + AccessRoleIds, + PrincipalType, + PermissionBits, + PermissionTypes, + Permissions, +} = require('librechat-data-provider'); +let mockBaseConfig = {}; + +const mockHrOverride = { + principalType: 'role', + principalId: 'HR', + priority: 10, + overrides: { + prompts: { + categories: { + enableDefaultCategories: false, + list: [{ value: 'benefits', label: 'Benefits' }], + }, + }, + }, +}; + +jest.mock('~/server/services/Config', () => ({ + getAppConfig: jest.fn(async ({ role } = {}) => { + const { mergeConfigOverrides } = require('@librechat/data-schemas'); + return role === 'HR' ? mergeConfigOverrides(mockBaseConfig, [mockHrOverride]) : mockBaseConfig; + }), +})); + +jest.mock('~/models', () => { + const mongoose = require('mongoose'); + const { createMethods } = require('@librechat/data-schemas'); + const methods = createMethods(mongoose, { + removeAllPermissions: async ({ resourceType, resourceId }) => { + await mongoose.models.AclEntry?.deleteMany({ resourceType, resourceId }); + }, + }); + return { + ...methods, + getPromptGroupAccessContext: jest.fn(methods.getPromptGroupAccessContext), + getDistinctPromptGroupCategories: jest.fn(methods.getDistinctPromptGroupCategories), + }; +}); + +jest.mock('~/server/middleware', () => ({ + requireJwtAuth: (req, res, next) => next(), + configMiddleware: jest.requireActual('~/server/middleware/config/app'), +})); + +const builtins = [ + { label: 'com_ui_idea', value: 'idea' }, + { label: 'com_ui_travel', value: 'travel' }, + { label: 'com_ui_teach_or_explain', value: 'teach_or_explain' }, + { label: 'com_ui_write', value: 'write' }, + { label: 'com_ui_shop', value: 'shop' }, + { label: 'com_ui_code', value: 'code' }, + { label: 'com_ui_misc', value: 'misc' }, + { label: 'com_ui_roleplay', value: 'roleplay' }, + { label: 'com_ui_finance', value: 'finance' }, +]; + +let app; +let mongoServer; +let models; +let users; +let currentUser; + +beforeAll(async () => { + mongoServer = await MongoMemoryServer.create(); + await mongoose.connect(mongoServer.getUri()); + + const { AccessRole, User, Role } = require('~/db/models'); + await Role.create({ + name: 'NO_PROMPTS', + permissions: { [PermissionTypes.PROMPTS]: { [Permissions.USE]: false } }, + }); + await AccessRole.create({ + accessRoleId: AccessRoleIds.PROMPTGROUP_OWNER, + name: 'Owner', + resourceType: ResourceType.PROMPTGROUP, + permBits: + PermissionBits.VIEW | PermissionBits.EDIT | PermissionBits.DELETE | PermissionBits.SHARE, + }); + users = { + a: await User.create({ name: 'A', email: 'a@example.com', role: SystemRoles.USER }), + b: await User.create({ name: 'B', email: 'b@example.com', role: SystemRoles.USER }), + hr: await User.create({ name: 'HR', email: 'hr@example.com', role: 'HR' }), + noPrompts: await User.create({ name: 'NP', email: 'np@example.com', role: 'NO_PROMPTS' }), + }; + models = require('~/models'); + + app = express(); + app.use((req, res, next) => { + req.user = { + id: currentUser._id.toString(), + _id: currentUser._id, + role: currentUser.role, + }; + next(); + }); + app.use('/api/categories', require('./categories')); +}); + +beforeEach(() => { + currentUser = users.a; + mockBaseConfig = {}; + jest.clearAllMocks(); +}); + +afterAll(async () => { + await mongoose.disconnect(); + await mongoServer.stop(); +}); + +describe('GET /api/categories', () => { + it('defaults unchanged when no prompts config is set', async () => { + const res = await request(app).get('/api/categories'); + + expect(res.status).toBe(200); + expect(res.body).toEqual(builtins); + }); + + it('applies the role override per requesting role', async () => { + currentUser = users.hr; + const hr = await request(app).get('/api/categories'); + expect(hr.status).toBe(200); + expect(hr.body).toEqual([{ value: 'benefits', label: 'Benefits' }]); + + currentUser = users.a; + const user = await request(app).get('/api/categories'); + expect(user.body).toEqual(builtins); + }); + + it('gives no error leak when the custom category read fails', async () => { + mockBaseConfig = { prompts: { categories: { allowCustom: true } } }; + models.getDistinctPromptGroupCategories.mockRejectedValueOnce(new Error('secret db detail')); + + const res = await request(app).get('/api/categories'); + + expect(res.status).toBe(500); + expect(res.body).toEqual({ message: 'Failed to retrieve categories' }); + expect(res.text).not.toContain('secret'); + }); + + it('keeps custom categories access scoped to the requesting user', async () => { + mockBaseConfig = { prompts: { categories: { allowCustom: true } } }; + const { grantPermission } = require('~/server/services/PermissionService'); + const { group } = await models.createPromptGroup({ + prompt: { prompt: 'secret text', type: 'text' }, + group: { name: 'private group', category: 'A-Private' }, + author: users.a._id.toString(), + authorName: users.a.name, + }); + await grantPermission({ + principalType: PrincipalType.USER, + principalId: users.a._id, + resourceType: ResourceType.PROMPTGROUP, + resourceId: group._id, + accessRoleId: AccessRoleIds.PROMPTGROUP_OWNER, + grantedBy: users.a._id, + }); + + currentUser = users.b; + const forB = await request(app).get('/api/categories'); + expect(forB.body.map((c) => c.value)).not.toContain('A-Private'); + + currentUser = users.a; + const forA = await request(app).get('/api/categories'); + expect(forA.body).toContainEqual({ value: 'A-Private', label: 'A-Private', custom: true }); + }); + + it('serves configured categories only to a user without prompt-use permission', async () => { + mockBaseConfig = { + prompts: { + categories: { allowCustom: true, enableDefaultCategories: false, list: [{ value: 'hr' }] }, + }, + }; + const { grantPermission } = require('~/server/services/PermissionService'); + const { group } = await models.createPromptGroup({ + prompt: { prompt: 'text', type: 'text' }, + group: { name: 'shared group', category: 'Stored' }, + author: users.noPrompts._id.toString(), + authorName: users.noPrompts.name, + }); + await grantPermission({ + principalType: PrincipalType.USER, + principalId: users.noPrompts._id, + resourceType: ResourceType.PROMPTGROUP, + resourceId: group._id, + accessRoleId: AccessRoleIds.PROMPTGROUP_OWNER, + grantedBy: users.noPrompts._id, + }); + + currentUser = users.noPrompts; + const denied = await request(app).get('/api/categories'); + expect(denied.status).toBe(200); + expect(denied.body).toEqual([{ value: 'hr', label: 'hr' }]); + expect(models.getDistinctPromptGroupCategories).not.toHaveBeenCalled(); + }); + + it('no reads when off: custom categories disabled', async () => { + mockBaseConfig = { prompts: { categories: { allowCustom: false } } }; + + const res = await request(app).get('/api/categories'); + + expect(res.status).toBe(200); + expect(models.getPromptGroupAccessContext).not.toHaveBeenCalled(); + expect(models.getDistinctPromptGroupCategories).not.toHaveBeenCalled(); + }); +}); diff --git a/api/server/routes/config.js b/api/server/routes/config.js index fd0adb57ae9..b1a11ed4728 100644 --- a/api/server/routes/config.js +++ b/api/server/routes/config.js @@ -22,6 +22,7 @@ const { resolveCodeEnvironmentMoveCapabilities, resolveCodeWorkspaceInheritanceCapability, resolveCodeEnvironmentTransitionVersion, + getPromptCategoriesStartupConfig, loadConversationListLimits, } = require('@librechat/api'); const { @@ -315,6 +316,7 @@ router.get('/', async function (req, res) { appConfig, endpoint: EModelEndpoint.agents, }), + promptCategories: getPromptCategoriesStartupConfig(appConfig), turnstile: appConfig?.turnstileConfig, modelSpecs: sanitizeModelSpecs(excludeHiddenModelSpecs(appConfig?.modelSpecs)), balance: balanceConfig, diff --git a/api/server/routes/prompts.test.js b/api/server/routes/prompts.test.js index 57b89b57c67..2bf0205ead9 100644 --- a/api/server/routes/prompts.test.js +++ b/api/server/routes/prompts.test.js @@ -313,6 +313,31 @@ describe('Prompt Routes - ACL Permissions', () => { await expect(PromptGroup.countDocuments()).resolves.toBe(0); }); + it('should reject a reserved category on create', async () => { + const response = await request(app) + .post('/api/prompts') + .send({ + prompt: { prompt: 'Category prompt', type: 'text' }, + group: { name: 'Category Group', category: 'sys__x' }, + }); + + expect(response.status).toBe(400); + await expect(PromptGroup.countDocuments()).resolves.toBe(0); + }); + + it('should store a valid custom category trimmed on create', async () => { + const response = await request(app) + .post('/api/prompts') + .send({ + prompt: { prompt: 'Category prompt', type: 'text' }, + group: { name: 'Category Group', category: ' Onboarding ' }, + }); + + expect(response.status).toBe(200); + const stored = await PromptGroup.findOne({ name: 'Category Group' }).lean(); + expect(stored.category).toBe('Onboarding'); + }); + it('should create a prompt and grant owner permissions', async () => { const promptData = { prompt: { @@ -943,6 +968,16 @@ describe('Prompt Routes - ACL Permissions', () => { await AclEntry.deleteMany({}); }); + it('should reject an over-long category on update and keep the stored category', async () => { + await request(app) + .patch(`/api/prompts/groups/${testGroup._id}`) + .send({ category: 'a'.repeat(101) }) + .expect(400); + + const stored = await PromptGroup.findById(testGroup._id).lean(); + expect(stored.category).toBe('security-test'); + }); + it('should allow updating allowed fields (name, category, oneliner)', async () => { const updateData = { name: 'Updated Group Name', diff --git a/client/src/components/Prompts/fields/CategorySelector.tsx b/client/src/components/Prompts/fields/CategorySelector.tsx index b0d163d21bc..f02e96282d8 100644 --- a/client/src/components/Prompts/fields/CategorySelector.tsx +++ b/client/src/components/Prompts/fields/CategorySelector.tsx @@ -1,12 +1,18 @@ -import React, { useId, useMemo, useState } from 'react'; +import React, { useCallback, useEffect, useId, useMemo, useRef, useState } from 'react'; import * as Ariakit from '@ariakit/react'; import { useTranslation } from 'react-i18next'; -import { DropdownPopup } from '@librechat/client'; -import { LocalStorageKeys } from 'librechat-data-provider'; import { useFormContext, Controller } from 'react-hook-form'; +import { Button, DropdownPopup, Input } from '@librechat/client'; +import { + LocalStorageKeys, + SYSTEM_CATEGORY_PREFIX, + promptCategoryValueSchema, +} from 'librechat-data-provider'; import type { MenuItemProps } from '@librechat/client'; import type { ReactNode } from 'react'; +import { useGetStartupConfig } from '~/data-provider'; import { usePromptGroupsContext } from '~/Providers'; +import { CategoryIcon } from '~/components/Prompts'; import { useCategories } from '~/hooks'; import { cn } from '~/utils'; @@ -29,6 +35,12 @@ const CategorySelector: React.FC = ({ const { t } = useTranslation(); const formContext = useFormContext(); const [isOpen, setIsOpen] = useState(false); + const [isCreating, setIsCreating] = useState(false); + const [draft, setDraft] = useState(''); + const triggerRef = useRef(null); + const inputRef = useRef(null); + const { data: startupConfig } = useGetStartupConfig(); + const allowCustom = startupConfig?.promptCategories?.allowCustom === true; const { hasAccess } = usePromptGroupsContext() ?? {}; const { categories, emptyCategory } = useCategories({ hasAccess }); @@ -38,13 +50,17 @@ const CategorySelector: React.FC = ({ const watchedCategory = watch ? watch('category') : currentCategory; - const categoryOption = useMemo( - () => - (categories ?? []).find( - (category) => category.value === (watchedCategory ?? currentCategory), - ) ?? emptyCategory, - [watchedCategory, categories, currentCategory, emptyCategory], - ); + const categoryOption = useMemo(() => { + const selected = watchedCategory ?? currentCategory; + const found = (categories ?? []).find((category) => category.value === selected); + if (found) { + return found; + } + if (selected && allowCustom) { + return { value: selected, label: selected, icon: }; + } + return emptyCategory; + }, [watchedCategory, categories, currentCategory, emptyCategory, allowCustom]); const displayCategory = useMemo(() => { if (!categoryOption.value && !('icon' in categoryOption)) { @@ -57,29 +73,101 @@ const CategorySelector: React.FC = ({ return categoryOption; }, [categoryOption, t]); + const selectCategory = useCallback( + (value: string) => { + if (formContext && setValue) { + setValue('category', value, { shouldDirty: false }); + } + localStorage.setItem(LocalStorageKeys.LAST_PROMPT_CATEGORY, value); + onValueChange?.(value); + setIsOpen(false); + }, + [formContext, setValue, onValueChange], + ); + const menuItems: MenuItemProps[] = useMemo(() => { if (!categories) return []; - return categories.map((category) => ({ + const items: MenuItemProps[] = categories.map((category) => ({ id: `${menuId}-item-${category.value}`, label: category.label, icon: 'icon' in category ? category.icon : undefined, - onClick: () => { - const value = category.value || ''; - if (formContext && setValue) { - setValue('category', value, { shouldDirty: false }); - } - localStorage.setItem(LocalStorageKeys.LAST_PROMPT_CATEGORY, value); - onValueChange?.(value); - setIsOpen(false); - }, + onClick: () => selectCategory(category.value || ''), })); - }, [categories, formContext, menuId, setValue, onValueChange]); + + if (allowCustom) { + items.push({ + id: `${menuId}-new-category`, + label: t('com_ui_new_category'), + onClick: () => { + setIsOpen(false); + setIsCreating(true); + }, + }); + } + return items; + }, [categories, allowCustom, menuId, selectCategory, t]); + + useEffect(() => { + if (isCreating) { + inputRef.current?.focus(); + } + }, [isCreating]); + + const trimmed = draft.trim(); + const existing = useMemo(() => { + const needle = trimmed.toLowerCase(); + return (categories ?? []).find( + (category) => + category.value && + (category.value.toLowerCase() === needle || category.label.toLowerCase() === needle), + ); + }, [categories, trimmed]); + const parsed = promptCategoryValueSchema.safeParse(draft); + const invalidReason = (() => { + if (existing || parsed.success || !trimmed) { + return null; + } + if (trimmed.startsWith(SYSTEM_CATEGORY_PREFIX)) { + return t('com_ui_category_reserved_prefix'); + } + return parsed.error.issues.some((issue) => issue.code === 'too_big') + ? t('com_ui_category_too_long') + : t('com_ui_category_invalid_chars'); + })(); + const canSubmit = existing != null || parsed.success; + const reasonId = `${menuId}-new-category-reason`; + + const closeCreate = () => { + setIsCreating(false); + setDraft(''); + }; + + const submitDraft = () => { + if (!canSubmit) { + return; + } + selectCategory(existing ? existing.value : trimmed); + closeCreate(); + }; + + const handleDraftKeyDown = (e: React.KeyboardEvent) => { + if (e.key === 'Escape') { + e.preventDefault(); + closeCreate(); + triggerRef.current?.focus(); + } + if (e.key === 'Enter' && !(e.nativeEvent.isComposing || e.keyCode === 229)) { + e.preventDefault(); + submitDraft(); + } + }; const trigger = ( = ({ ); - return formContext ? ( - ( - - )} - /> - ) : ( + const popup = ( = ({ portal={portal} /> ); + + const createForm = isCreating && ( +
+ setDraft(e.target.value)} + onKeyDown={handleDraftKeyDown} + aria-label={t('com_ui_category_name')} + aria-describedby={invalidReason ? reasonId : undefined} + aria-invalid={invalidReason ? true : undefined} + /> + + {invalidReason && ( + + {invalidReason} + + )} +
+ ); + + return ( + <> + {formContext ? popup} /> : popup} + {createForm} + + ); }; export default CategorySelector; diff --git a/client/src/components/Prompts/fields/__tests__/CategorySelector.spec.tsx b/client/src/components/Prompts/fields/__tests__/CategorySelector.spec.tsx new file mode 100644 index 00000000000..20fff3ea026 --- /dev/null +++ b/client/src/components/Prompts/fields/__tests__/CategorySelector.spec.tsx @@ -0,0 +1,198 @@ +import React from 'react'; +import { RecoilRoot } from 'recoil'; +import '@testing-library/jest-dom/extend-expect'; +import { FormProvider, useForm } from 'react-hook-form'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import { render, screen, fireEvent, waitFor } from '@testing-library/react'; +import type { TCategory } from 'librechat-data-provider'; +import { useGetCategories, useUpdatePromptGroup } from '~/data-provider'; +import CategorySelector from '../CategorySelector'; + +const mockGetCategories = jest.fn(); +const mockGetStartupConfig = jest.fn(); +const mockUpdatePromptGroup = jest.fn(); + +jest.mock('librechat-data-provider', () => { + const actual = jest.requireActual('librechat-data-provider'); + return { + ...actual, + dataService: { + ...actual.dataService, + getCategories: () => mockGetCategories(), + getStartupConfig: () => mockGetStartupConfig(), + updatePromptGroup: (vars: unknown) => mockUpdatePromptGroup(vars), + }, + }; +}); + +const categories: TCategory[] = [ + { value: 'hr', label: 'Human Resources' }, + { value: 'idea', label: 'com_ui_idea' }, +]; + +let formValues: { category?: string } = {}; + +const FormProbe = ({ initial = '' }: { initial?: string }) => { + const methods = useForm({ defaultValues: { category: initial } }); + formValues = methods.watch(); + return ( + + + + ); +}; + +const renderSelector = (allowCustom?: boolean, initial?: string) => { + mockGetStartupConfig.mockResolvedValue( + allowCustom === undefined ? {} : { promptCategories: { allowCustom } }, + ); + const queryClient = new QueryClient(); + const utils = render( + + + + + , + ); + return { ...utils, queryClient }; +}; + +const openMenu = async () => { + await waitFor(() => expect(mockGetCategories).toHaveBeenCalled()); + fireEvent.click(screen.getByRole('button', { name: /category/i })); +}; + +const openForm = async () => { + await openMenu(); + fireEvent.click(await screen.findByRole('menuitem', { name: /New category/ })); + return screen.findByRole('textbox', { name: 'Category name' }); +}; + +describe('CategorySelector custom categories', () => { + beforeEach(() => { + formValues = {}; + mockGetCategories.mockResolvedValue(categories); + }); + + afterEach(() => { + mockGetCategories.mockReset(); + mockGetStartupConfig.mockReset(); + mockUpdatePromptGroup.mockReset(); + localStorage.clear(); + }); + + it('creates a category from the keyboard and shows it in the trigger', async () => { + renderSelector(true); + const input = await openForm(); + await waitFor(() => expect(input).toHaveFocus()); + + fireEvent.change(input, { target: { value: 'Onboarding' } }); + expect(screen.getByRole('button', { name: 'Create “Onboarding”' })).toBeEnabled(); + fireEvent.keyDown(input, { key: 'Enter' }); + + await waitFor(() => expect(formValues.category).toBe('Onboarding')); + expect(screen.getByRole('button', { name: /category/i })).toHaveTextContent('Onboarding'); + expect(screen.queryByRole('textbox', { name: 'Category name' })).not.toBeInTheDocument(); + }); + + it('closes the form on Escape and returns focus to the trigger', async () => { + renderSelector(true); + const input = await openForm(); + + fireEvent.keyDown(input, { key: 'Escape' }); + + expect(screen.queryByRole('textbox', { name: 'Category name' })).not.toBeInTheDocument(); + expect(screen.getByRole('button', { name: /category/i })).toHaveFocus(); + }); + + it.each([' hr ', 'human resources'])('reuses the existing category for %p', async (typed) => { + renderSelector(true); + const input = await openForm(); + + fireEvent.change(input, { target: { value: typed } }); + fireEvent.click(screen.getByRole('button', { name: 'Use “Human Resources”' })); + + await waitFor(() => expect(formValues.category).toBe('hr')); + }); + + it.each([ + ['sys__x', 'reserved'], + ['a'.repeat(101), 'too long'], + ])('disables creation for %p and links the reason', async (typed, reason) => { + renderSelector(true); + const input = await openForm(); + + fireEvent.change(input, { target: { value: typed } }); + + const button = screen.getByRole('button', { name: /^Create/ }); + expect(button).toBeDisabled(); + const describedBy = button.getAttribute('aria-describedby') ?? ''; + const message = document.getElementById(describedBy); + expect(message).toHaveTextContent(new RegExp(reason === 'reserved' ? 'sys__' : '100')); + }); + + it('offers no new-category item when custom categories are not allowed', async () => { + renderSelector(false); + await openMenu(); + + await screen.findByRole('menuitem', { name: /Human Resources/ }); + expect(screen.queryByRole('menuitem', { name: /New category/ })).not.toBeInTheDocument(); + }); +}); + +describe('CategorySelector selected value missing from the list', () => { + beforeEach(() => { + mockGetCategories.mockResolvedValue(categories); + }); + + afterEach(() => { + mockGetCategories.mockReset(); + mockGetStartupConfig.mockReset(); + }); + + it('shows the placeholder when custom categories are not allowed', async () => { + renderSelector(false, 'travel'); + await waitFor(() => expect(mockGetStartupConfig).toHaveBeenCalled()); + + const trigger = screen.getByRole('button', { name: /category/i }); + expect(trigger).toHaveTextContent('Category'); + expect(trigger).not.toHaveTextContent('travel'); + }); + + it('shows the value itself when custom categories are allowed', async () => { + renderSelector(true, 'Onboarding'); + + await waitFor(() => + expect(screen.getByRole('button', { name: /category/i })).toHaveTextContent('Onboarding'), + ); + }); +}); + +describe('categories invalidation after saving a group', () => { + it('refetches categories after a group update', async () => { + mockGetCategories.mockResolvedValue(categories); + mockUpdatePromptGroup.mockResolvedValue({ _id: 'g1' }); + const Harness = () => { + useGetCategories(); + const { mutate } = useUpdatePromptGroup(); + return ( +