From 289ea7da8dfcbd4a37b68b3576f586af8bf505e7 Mon Sep 17 00:00:00 2001 From: ManishMadan2882 Date: Tue, 15 Sep 2026 04:18:04 +0530 Subject: [PATCH 001/130] feat(widget): align search bar with widget theme, add voice search --- docs/content/Extensions/search-widget.mdx | 4 +- extensions/react-widget/README.md | 2 + extensions/react-widget/src/App.tsx | 2 +- .../src/components/ComposerControls.tsx | 51 +- .../src/components/DocsGPTWidget.tsx | 201 ++------ .../react-widget/src/components/SearchBar.tsx | 446 +++++++++--------- .../react-widget/src/components/tokens.ts | 66 +++ .../react-widget/src/hooks/useAttachments.ts | 5 +- .../react-widget/src/hooks/useDictation.ts | 65 +++ .../react-widget/src/hooks/useVoiceInput.ts | 54 +-- .../src/requests/attachmentsApi.ts | 16 +- extensions/react-widget/src/types/index.ts | 34 +- 12 files changed, 506 insertions(+), 440 deletions(-) create mode 100644 extensions/react-widget/src/hooks/useDictation.ts diff --git a/docs/content/Extensions/search-widget.mdx b/docs/content/Extensions/search-widget.mdx index 93a15350..879fba57 100644 --- a/docs/content/Extensions/search-widget.mdx +++ b/docs/content/Extensions/search-widget.mdx @@ -98,9 +98,11 @@ The DocsGPT Search Bar Widget offers a range of customizable properties that all |-----------------|-----------|-------------------------------------|--------------------------------------------------------------------------------------------------| | **`apiKey`** | `string` | `"your-api-key"` | API key for authentication with your DocsGPT API. Leave empty if no authentication is required. | | **`apiHost`** | `string` | `"https://gptcloud.arc53.com"` | **Required.** The URL of your DocsGPT API backend. This endpoint handles vector similarity search queries. | -| **`theme`** | `"dark" \| "light"` | `"dark"` | Color theme of the search bar. Options: `"dark"` or `"light"`. Defaults to `"dark"`. | +| **`theme`** | `"dark" \| "light"` | `"dark"` | Color theme of the search bar and of the chat it opens. Options: `"dark"` or `"light"`. Defaults to `"dark"`. | | **`placeholder`** | `string` | `"Search or Ask AI..."` | Placeholder text displayed in the search input field. | | **`width`** | `string` | `"256px"` | Width of the search bar. Accepts any valid CSS width value (e.g., `"300px"`, `"100%"`, `"20rem"`). | +| **`allowedFileExtensions`** | `string[]` | _unset_ | Passed to the chat opened from "Ask the AI". File extensions its composer accepts, e.g. `['.pdf', '.md']`; attachments stay off while unset. See the [chat widget page](/Extensions/chat-widget). | +| **`showMicButton`** | `boolean` | `false` | Adds a microphone to the search field for dictating a query, and passes the same option to the chat opened from "Ask the AI". Uses the browser's Web Speech API. See the [chat widget page](/Extensions/chat-widget). | --- diff --git a/extensions/react-widget/README.md b/extensions/react-widget/README.md index d93ed110..858235d9 100644 --- a/extensions/react-widget/README.md +++ b/extensions/react-widget/README.md @@ -172,6 +172,8 @@ import { SearchBar } from "docsgpt-react"; | **`theme`** | `"dark" \| "light"` | `"dark"` | The theme of the search bar. Accepts `"dark"` or `"light"`. | | **`placeholder`** | `string` | `"Search or Ask AI..."` | Placeholder text displayed in the search input field. | | **`width`** | `string` | `"256px"` | Width of the search bar. Accepts any valid CSS width value (e.g., `"300px"`, `"100%"`, `"20rem"`). | +| **`allowedFileExtensions`** | `string[]` | _unset_ | Passed to the chat opened from "Ask the AI". File extensions its composer accepts, e.g. `['.pdf', '.md']`; attachments stay off while unset. | +| **`showMicButton`** | `boolean` | `false` | Adds a microphone to the search field for dictating a query, and passes the same option to the chat opened from "Ask the AI". Uses the browser's Web Speech API. | Feel free to reach out if you need help customizing or extending the `SearchBar`! diff --git a/extensions/react-widget/src/App.tsx b/extensions/react-widget/src/App.tsx index b30249e3..0c1f60ba 100644 --- a/extensions/react-widget/src/App.tsx +++ b/extensions/react-widget/src/App.tsx @@ -4,7 +4,7 @@ import { SearchBar } from './components/SearchBar'; export const App = () => { return (
- + ` width: 14px; @@ -238,7 +238,6 @@ const ChipRemove = styled.button` } `; -/** Chip status, for the hover title. */ const statusLabel = (attachment: Attachment): string => { if (attachment.status === 'uploading') return `Uploading ${attachment.progress}%`; @@ -254,8 +253,7 @@ export const AttachmentChips = ({ attachments: Attachment[]; onRemove: (id: string) => void; }) => { - // A touch user cannot see a tooltip, and a phone picker's unsupported - // file lands here. Say why in the open. + // Tooltips are unreachable on touch, so failure reasons are shown inline. const failures = attachments.filter( (attachment) => attachment.status === 'failed' && attachment.error, ); @@ -322,12 +320,19 @@ export const ControlGroup = styled.div` gap: 6px; `; -const ControlButton = styled.button<{ $recording?: boolean }>` +const ControlButton = styled.button<{ + $recording?: boolean; + $iconOnly?: boolean; +}>` display: inline-flex; align-items: center; gap: 5px; - height: 28px; - padding: 0 10px; + height: ${(props) => (props.$iconOnly ? '32px' : '28px')}; + padding: ${(props) => (props.$iconOnly ? '0' : '0 10px')}; + ${(props) => + props.$iconOnly + ? 'width: 32px; flex-shrink: 0; justify-content: center;' + : ''} border-radius: ${radii.full}; border: 1px solid ${(props) => @@ -371,9 +376,7 @@ export const AttachButton = ({ onClick: () => void; disabled?: boolean; }) => ( - // A button opening a hidden input, not a label wrapping one: a - // `display: none` input takes no focus, so the label form is unreachable - // from the keyboard. + // A label wrapping a display:none input cannot be reached by keyboard. void; + variant?: 'pill' | 'icon'; }) => ( @@ -422,7 +428,7 @@ export const MicButton = ({ ) : ( ); @@ -472,22 +478,21 @@ export const SentAttachments = ({ ); -// One bar per animation frame, so the row holds about a second of history. +// One bar per frame: roughly a second of history at 60fps. const WAVEFORM_BARS = 48; const BAR_GAP_RATIO = 0.4; // Speech peaks well below full scale; the floor keeps a quiet line visible. const LEVEL_GAIN = 2.8; const MIN_BAR_RATIO = 0.06; -const WaveformRow = styled.div` +const WaveformRow = styled.div<{ $minHeight: string }>` display: flex; flex: 1; min-width: 0; align-items: center; gap: 10px; padding: 0 10px; - min-height: ${(props) => - props.theme.dimensions!.size === 'large' ? '60px' : '40px'}; + min-height: ${(props) => props.$minHeight}; `; const WaveformCanvas = styled.canvas` @@ -505,16 +510,18 @@ const ListeningLabel = styled.span` `; /** - * Live microphone level, standing in for the input while dictation runs. - * Canvas inside a rAF loop, because sixty React updates a second would - * re-render the whole composer. A null analyser draws the resting line. + * Live microphone level, drawn to canvas in a rAF loop so frames do not + * re-render React. A null analyser draws a flat line. */ export const VoiceWaveform = ({ analyserRef, label, + minHeight = '40px', }: { analyserRef: React.RefObject; label: string; + /** Min height of the input it replaces, to avoid a layout shift. */ + minHeight?: string; }) => { const canvasRef = React.useRef(null); const theme = useTheme(); @@ -543,7 +550,7 @@ export const VoiceWaveform = ({ if (!samples || samples.length !== analyser.fftSize) samples = new Uint8Array(analyser.fftSize); analyser.getByteTimeDomainData(samples); - // Time-domain bytes ride on 128; RMS of the deviation is the level. + // Time-domain bytes are centred on 128. let sumSquares = 0; for (let index = 0; index < samples.length; index += 1) { const deviation = (samples[index] - 128) / 128; @@ -554,7 +561,7 @@ export const VoiceWaveform = ({ levels.push(Math.min(1, level * LEVEL_GAIN)); levels.shift(); - // Re-read each frame: the panel resizes and the row reflows. + // The panel can resize, so dimensions are read each frame. const ratio = window.devicePixelRatio || 1; const width = canvas.clientWidth; const height = canvas.clientHeight; @@ -592,7 +599,7 @@ export const VoiceWaveform = ({ }, [analyserRef, barColor]); return ( - + diff --git a/extensions/react-widget/src/components/DocsGPTWidget.tsx b/extensions/react-widget/src/components/DocsGPTWidget.tsx index c6a0669f..58aee967 100644 --- a/extensions/react-widget/src/components/DocsGPTWidget.tsx +++ b/extensions/react-widget/src/components/DocsGPTWidget.tsx @@ -24,7 +24,7 @@ import { normalizeExtensions, useAttachments, } from '../hooks/useAttachments'; -import { useVoiceInput, voiceInputSupported } from '../hooks/useVoiceInput'; +import { useDictation } from '../hooks/useDictation'; import { AttachButton, AttachmentChips, @@ -35,12 +35,11 @@ import { DropOverlay, DropTarget, MicButton, - type MicButtonState, SentAttachments, VoiceWaveform, } from './ComposerControls'; import { DEFAULT_AVATAR } from './defaultAvatar'; -import { radii } from './tokens'; +import { radii, themes } from './tokens'; import { ThemeProvider } from 'styled-components'; import MarkdownIt from 'markdown-it'; import { @@ -181,69 +180,6 @@ const ArrowDownIcon = (props: React.SVGProps) => ( ); -const themes = { - dark: { - bg: '#222327', - text: '#fff', - primary: { - text: '#FAFAFA', - bg: '#222327', - }, - secondary: { - text: '#A1A1AA', - bg: '#33343A', - }, - shimmer: { - base: '#A1A1AA', - highlight: '#FAFAFA', - }, - accent: { - base: '#8860DB', - hover: '#9B7BE4', - strong: '#6D42C5', - contrast: '#FFFFFF', - soft: 'rgba(136, 96, 219, 0.18)', - link: '#A78BFA', - }, - hairline: 'rgba(255, 255, 255, 0.08)', - danger: { - text: '#F87171', - soft: 'rgba(248, 113, 113, 0.10)', - border: 'rgba(248, 113, 113, 0.32)', - }, - }, - light: { - bg: '#fff', - text: '#000', - primary: { - text: '#222327', - bg: '#fff', - }, - secondary: { - text: '#71717A', - bg: '#F4F4F5', - }, - shimmer: { - base: '#71717A', - highlight: '#D4D4D8', - }, - accent: { - base: '#8860DB', - hover: '#7A4FD0', - strong: '#6D42C5', - contrast: '#FFFFFF', - soft: 'rgba(136, 96, 219, 0.12)', - link: '#6D42C5', - }, - hairline: 'rgba(0, 0, 0, 0.08)', - danger: { - text: '#B91C1C', - soft: 'rgba(185, 28, 28, 0.06)', - border: 'rgba(185, 28, 28, 0.24)', - }, - }, -}; - const sizesConfig = { small: { size: 'small', width: '320px', height: '400px' }, medium: { size: 'medium', width: '400px', height: '80vh' }, @@ -656,28 +592,20 @@ const ActionButton = styled.button<{ ${(props) => props.$active && css` - color: ${ - props.$tone === 'danger' - ? props.theme.danger!.text - : props.theme.accent!.base - }; - background-color: ${ - props.$tone === 'danger' - ? props.theme.danger!.soft - : props.theme.accent!.soft - }; + color: ${props.$tone === 'danger' + ? props.theme.danger!.text + : props.theme.accent!.base}; + background-color: ${props.$tone === 'danger' + ? props.theme.danger!.soft + : props.theme.accent!.soft}; &:hover { - color: ${ - props.$tone === 'danger' - ? props.theme.danger!.text - : props.theme.accent!.base - }; - background-color: ${ - props.$tone === 'danger' - ? props.theme.danger!.soft - : props.theme.accent!.soft - }; + color: ${props.$tone === 'danger' + ? props.theme.danger!.text + : props.theme.accent!.base}; + background-color: ${props.$tone === 'danger' + ? props.theme.danger!.soft + : props.theme.accent!.soft}; } svg { @@ -1188,7 +1116,6 @@ const HeroDescription = styled.p` padding: 0px; `; const Hyperlink = styled.a` - /* Inherits the tagline colour; the underline is what marks it as a link. */ color: inherit; text-decoration: underline; /* Keeps descenders clear of the rule at 11px. */ @@ -1370,7 +1297,7 @@ export const WidgetCore = ({ heroTitle = 'Welcome to DocsGPT !', heroDescription = 'This chatbot is built with DocsGPT and utilises GenAI, please review important information using sources.', size = 'medium', - theme = 'light', + theme = 'dark', collectFeedback = true, isOpen = false, showSources = true, @@ -1400,12 +1327,10 @@ export const WidgetCore = ({ const endMessageRef = React.useRef(null); const promptRef = React.useRef(null); const attachmentInputRef = React.useRef(null); - // The draft as it stood when dictation began; recognised words extend it. - const voiceBaseRef = React.useRef(''); // dragenter/dragleave fire per child crossed, hence a depth count. const dragDepthRef = React.useRef(0); - // The list doubles as the on switch: no accepted types, no attachments. + // An empty list disables attachments. const acceptedExtensions = React.useMemo( () => normalizeExtensions(allowedFileExtensions), [allowedFileExtensions], @@ -1422,8 +1347,6 @@ export const WidgetCore = ({ completed: completedAttachments, } = useAttachments({ apiKey, apiHost, acceptedExtensions }); - // One place for the height arithmetic: typing, transcripts and sends all - // go through here. const resizePrompt = React.useCallback(() => { const el = promptRef.current; if (!el) return; @@ -1436,42 +1359,31 @@ export const WidgetCore = ({ )}px`; }, [size]); - // The textarea is controlled by `prompt`, so its value is the live draft. - const handleVoiceStart = React.useCallback(() => { - voiceBaseRef.current = promptRef.current?.value ?? ''; - }, []); + const getPromptDraft = React.useCallback( + () => promptRef.current?.value ?? '', + [], + ); - const applyTranscript = React.useCallback( - (text: string) => { - // The base is fixed at the start, so revised interim words overwrite - // only themselves and a typed half-question survives. - const base = voiceBaseRef.current; - setPrompt(base.trim() ? `${base.replace(/\s+$/, '')}\n${text}` : text); + const applyDraft = React.useCallback( + (value: string) => { + setPrompt(value); // The height can only be recomputed once React has committed. window.requestAnimationFrame(resizePrompt); }, [resizePrompt], ); - const handleVoiceEnd = React.useCallback(() => { - // Once at the end, not on every interim word. + const focusPrompt = React.useCallback(() => { window.requestAnimationFrame(() => promptRef.current?.focus()); }, []); - const { - recordingState, - error: voiceError, - toggle: toggleVoiceInput, - clearError: clearVoiceError, - analyserRef: voiceAnalyserRef, - } = useVoiceInput({ - onStart: handleVoiceStart, - onTranscript: applyTranscript, - onEnd: handleVoiceEnd, + const dictation = useDictation({ + enabled: showMicButton, + getDraft: getPromptDraft, + onDraftChange: applyDraft, + separator: '\n', + onEnd: focusPrompt, }); - - // Firefox has no SpeechRecognition, nor does an insecure origin. - const canUseVoice = showMicButton && voiceInputSupported(); const md = new MarkdownIt(); //Custom markdown for the table md.renderer.rules.table_open = () => @@ -1738,8 +1650,7 @@ export const WidgetCore = ({ if (status === 'loading') return; const prompt = queries[index]?.prompt; if (!prompt) return; - // The composer list was cleared on send, so the row's ids are the only - // record of what the question carried. + // The composer list is cleared on send, so retry reuses the row's ids. const attached = queries[index]?.attachments; setQueries((prev: Query[]) => { const updated = [...prev]; @@ -1753,8 +1664,8 @@ export const WidgetCore = ({ ); }; - // Pending and failed attachments both hold the send, so neither is - // silently dropped. Pending reads first: it clears on its own. + // Pending and failed attachments both block sending. The pending note + // takes precedence because it clears on its own. const pendingNote = pendingCount > 0 ? `Waiting for ${pendingCount} file${pendingCount === 1 ? '' : 's'} to finish\u2026` @@ -1764,8 +1675,7 @@ export const WidgetCore = ({ ? 'Remove the file that could not be attached, then send.' : null; const sendBlockedReason = pendingNote ?? failedNote; - const isDictating = - recordingState === 'recording' || recordingState === 'transcribing'; + const isDictating = dictation.isDictating; const canSubmit = prompt.trim().length > 0 && !sendBlockedReason && @@ -1774,7 +1684,7 @@ export const WidgetCore = ({ const submitPrompt = async () => { if (!canSubmit) return; - // Before the value clears, or the composer sits tall and empty a frame. + // Reset first, so the empty composer does not render tall for a frame. if (promptRef.current) promptRef.current.style.height = 'auto'; await appendQuery(prompt); }; @@ -1802,7 +1712,7 @@ export const WidgetCore = ({ ) => { const value = event.target.value; // A stale voice error would hide the attachment notice below. - if (voiceError) clearVoiceError(); + if (dictation.error) dictation.clearError(); setPrompt(value); resizePrompt(); if (value.includes('\n')) { @@ -1829,17 +1739,16 @@ export const WidgetCore = ({ if (file) files.push(file); } if (files.length === 0) return; - // Or the file also lands in the textarea as binary noise. + // Keeps the file from also being pasted as text. e.preventDefault(); addFiles(files); }; - // Otherwise the overlay flashes on every drag crossing the page. + // Ignore drags that carry no files, such as selected text. const isFileDrag = (e: React.DragEvent) => Array.from(e.dataTransfer?.types ?? []).includes('Files'); - // A file dropped mid-dictation would queue behind the waveform, unseen - // and unsendable until recording stops. + // Files cannot be attached while dictating. const acceptsFiles = attachmentsEnabled && !isDictating; const handleDragEnter = (e: React.DragEvent) => { @@ -1891,16 +1800,11 @@ export const WidgetCore = ({ }; // Neither feature enabled keeps the original single-row composer. - const hasComposerControls = attachmentsEnabled || canUseVoice; + const hasComposerControls = attachmentsEnabled || dictation.available; - const micButtonState: MicButtonState = - recordingState === 'recording' || recordingState === 'transcribing' - ? recordingState - : 'idle'; - - // A failure to act on outranks a wait that clears itself. - const composerNote = voiceError - ? { text: voiceError, tone: 'danger' as const } + // Errors take precedence over the pending-upload note. + const composerNote = dictation.error + ? { text: dictation.error, tone: 'danger' as const } : sendBlockedReason ? { text: sendBlockedReason, @@ -2196,16 +2100,16 @@ export const WidgetCore = ({ {isDictating && ( )} - {/* Hidden, not unmounted: the ref stays live for reading - the draft and refocusing afterwards. */} + {/* Kept mounted so promptRef stays valid while dictating. */} )} - {canUseVoice && ( + {dictation.available && ( )} diff --git a/extensions/react-widget/src/components/SearchBar.tsx b/extensions/react-widget/src/components/SearchBar.tsx index f3e6b6ea..f8f75b27 100644 --- a/extensions/react-widget/src/components/SearchBar.tsx +++ b/extensions/react-widget/src/components/SearchBar.tsx @@ -1,7 +1,10 @@ import React from 'react'; -import styled, { ThemeProvider, createGlobalStyle } from 'styled-components'; +import styled, { ThemeProvider, keyframes } from 'styled-components'; import { WidgetCore } from './DocsGPTWidget'; import { DEFAULT_AVATAR } from './defaultAvatar'; +import { radii, themes } from './tokens'; +import { MicButton, VoiceWaveform } from './ComposerControls'; +import { useDictation } from '../hooks/useDictation'; import { SearchBarProps } from '@/types'; import { getSearchResults } from '../requests/searchAPI'; import { Result } from '@/types'; @@ -15,135 +18,107 @@ import { ListBulletIcon, QuoteIcon, } from '@radix-ui/react-icons'; -const themes = { - dark: { - name: 'dark', - bg: '#202124', - text: '#EDEDED', - primary: { - text: '#FAFAFA', - bg: '#111111', - }, - secondary: { - text: '#A1A1AA', - bg: '#38383b', - }, - }, - light: { - name: 'light', - bg: '#EAEAEA', - text: '#171717', - primary: { - text: '#222327', - bg: '#fff', - }, - secondary: { - text: '#A1A1AA', - bg: '#F6F6F6', - }, - }, -}; -const GlobalStyle = createGlobalStyle` - .highlight { - color: ${(props) => (props.theme.name === 'dark' ? '#4B9EFF' : '#0066CC')}; - font-weight: 500; +const spin = keyframes` + to { + transform: rotate(360deg); } `; -const loadGeistFont = () => { - const link = document.createElement('link'); - link.href = - 'https://fonts.googleapis.com/css2?family=Geist:wght@100..900&display=swap'; - link.rel = 'stylesheet'; - document.head.appendChild(link); -}; - const Main = styled.div` all: initial; - font-family: 'Geist', sans-serif; + font-family: + -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; `; + const SearchButton = styled.button<{ $inputWidth: string }>` - padding: 6px 6px; - font-family: inherit; + box-sizing: border-box; width: ${({ $inputWidth }) => $inputWidth}; - border-radius: 8px; - display: inline; - color: ${(props) => props.theme.secondary.text}; - outline: none; - border: none; - background-color: ${(props) => props.theme.secondary.bg}; - -webkit-appearance: none; - -moz-appearance: none; - appearance: none; - transition: background-color 128ms linear; + height: 36px; + padding: 0 72px 0 12px; + font-family: inherit; + font-size: 14px; text-align: left; + color: ${(props) => props.theme.secondary.text}; + background-color: ${(props) => props.theme.secondary.bg}; + border: 1px solid ${(props) => props.theme.hairline}; + border-radius: ${radii.sm}; + outline: none; cursor: pointer; + -webkit-appearance: none; + appearance: none; + transition: + color 0.15s ease, + border-color 0.15s ease; + + &:hover { + color: ${(props) => props.theme.primary.text}; + } + + &:focus-visible { + border-color: ${(props) => props.theme.accent!.base}; + box-shadow: 0 0 0 3px ${(props) => props.theme.accent!.soft}; + } `; const Container = styled.div` position: relative; display: inline-block; `; + const SearchOverlay = styled.div` position: fixed; top: 0; left: 0; width: 100%; height: 100%; - background-color: #0000001a; - backdrop-filter: blur(8px); - -webkit-backdrop-filter: blur(8px); z-index: 99; + background-color: rgba(0, 0, 0, 0.5); `; const SearchResults = styled.div` position: fixed; + top: 50%; + left: 50%; + z-index: 100; + transform: translate(-50%, -50%); + box-sizing: border-box; display: flex; flex-direction: column; - background-color: ${(props) => - props.theme.name === 'dark' - ? 'rgba(0, 0, 0, 0.15)' - : 'rgba(255, 255, 255, 0.4)'}; - border: 1px solid rgba(255, 255, 255, 0.18); - border-radius: 15px; - padding: 8px 0px 8px 0px; width: 792px; max-width: 90vw; height: 396px; - z-index: 100; - left: 50%; - top: 50%; - transform: translate(-50%, -50%); + padding: 8px 0; + overflow: hidden; color: ${(props) => props.theme.primary.text}; - - box-shadow: 0 8px 32px 0 rgba(31, 38, 135, 0.37); - backdrop-filter: blur(82px); - -webkit-backdrop-filter: blur(82px); - border-radius: 10px; - - box-sizing: border-box; + background-color: ${(props) => props.theme.primary.bg}; + border: 1px solid ${(props) => props.theme.hairline}; + border-radius: ${radii.panel}; + box-shadow: + 0 12px 44px rgba(0, 0, 0, 0.18), + 0 2px 8px rgba(0, 0, 0, 0.1); @media only screen and (max-width: 768px) { - height: 80vh; width: 90vw; + height: 80vh; } `; const SearchResultsScroll = styled.div` flex: 1; - overflow-y: auto; + padding: 0 16px; overflow-x: hidden; + overflow-y: auto; scrollbar-gutter: stable; scrollbar-width: thin; - scrollbar-color: #383838 transparent; - padding: 0 16px; + scrollbar-color: ${(props) => props.theme.hairline} transparent; `; const IconTitleWrapper = styled.div` display: flex; align-items: center; gap: 8px; + color: ${(props) => props.theme.secondary.text}; .element-icon { margin: 4px; @@ -151,208 +126,199 @@ const IconTitleWrapper = styled.div` `; const Title = styled.h3` - font-size: 15px; - font-weight: 400; - color: ${(props) => props.theme.primary.text}; margin: 0; + font-size: 15px; + font-weight: 500; + color: ${(props) => props.theme.primary.text}; overflow-wrap: break-word; - white-space: normal; - overflow: hidden; - text-overflow: ellipsis; `; + const ContentWrapper = styled.div` display: flex; flex-direction: column; - gap: 12px; + gap: 8px; `; const ResultWrapper = styled.div` + box-sizing: border-box; display: flex; align-items: flex-start; width: 100%; - box-sizing: border-box; - padding: 8px 16px; - cursor: pointer; - background-color: transparent; - font-family: 'Geist', sans-serif; - border-radius: 8px; - - word-wrap: break-word; + padding: 10px 12px; + overflow: hidden; overflow-wrap: break-word; word-break: break-word; - white-space: normal; - overflow: hidden; - text-overflow: ellipsis; + border-radius: ${radii.sm}; + cursor: pointer; + transition: background-color 0.15s ease; &:hover { - backdrop-filter: blur(8px); - -webkit-backdrop-filter: blur(8px); + background-color: ${(props) => props.theme.secondary.bg}; } `; const Content = styled.div` display: flex; - margin-left: 8px; flex-direction: column; - gap: 8px; - padding: 4px 0px 0px 12px; - font-size: 15px; - color: ${(props) => props.theme.primary.text}; - line-height: 1.6; - border-left: 2px solid ${(props) => props.theme.primary.text}CC; + gap: 6px; + margin-left: 7px; + padding: 2px 0 0 14px; overflow: hidden; + font-size: 14px; + line-height: 1.6; + color: ${(props) => props.theme.secondary.text}; + border-left: 2px solid ${(props) => props.theme.hairline}; + + /* Scoped so the host page's own .highlight elements are untouched. */ + .highlight { + padding: 1px 2px; + border-radius: 4px; + font-weight: 500; + color: ${(props) => props.theme.primary.text}; + background-color: ${(props) => props.theme.accent!.mark}; + } + + /* Snippet HTML can contain raw links. */ + a { + color: ${(props) => props.theme.accent!.link}; + text-decoration: underline; + text-underline-offset: 2px; + } + + a:hover { + color: ${(props) => props.theme.accent!.base}; + } `; + const ContentSegment = styled.div` display: flex; align-items: flex-start; gap: 8px; padding-right: 16px; - overflow-wrap: break-word; - white-space: normal; overflow: hidden; - text-overflow: ellipsis; + overflow-wrap: break-word; `; const Toolkit = styled.kbd` position: absolute; - right: 4px; top: 50%; - transform: translateY(-50%); - background-color: ${(props) => props.theme.primary.bg}; - color: ${(props) => props.theme.secondary.text}; - font-weight: 600; - font-size: 10px; - padding: 3px 6px; - border: 1px solid ${(props) => props.theme.secondary.text}; - border-radius: 4px; - display: flex; - align-items: center; - justify-content: center; + right: 8px; z-index: 1; + transform: translateY(-50%); + padding: 2px 6px; + font-family: inherit; + font-size: 11px; + font-weight: 500; + line-height: 1.6; + white-space: nowrap; + color: ${(props) => props.theme.secondary.text}; + background-color: ${(props) => props.theme.primary.bg}; + border: 1px solid ${(props) => props.theme.hairline}; + border-radius: ${radii.sm}; pointer-events: none; `; -const Loader = styled.div` - margin: 2rem auto; - border: 4px solid - ${(props) => - props.theme.name === 'dark' - ? 'rgba(255, 255, 255, 0.2)' - : 'rgba(0, 0, 0, 0.1)'}; - border-top: 4px solid - ${(props) => - props.theme.name === 'dark' ? '#FFFFFF' : props.theme.primary.bg}; - border-radius: 50%; - width: 12px; - height: 12px; - animation: spin 1s linear infinite; - @keyframes spin { - 0% { - transform: rotate(0deg); - } - 100% { - transform: rotate(360deg); - } - } +const Loader = styled.div` + width: 16px; + height: 16px; + margin: 2rem auto; + border: 2px solid ${(props) => props.theme.hairline}; + border-top-color: ${(props) => props.theme.accent!.base}; + border-radius: 50%; + animation: ${spin} 0.8s linear infinite; `; const NoResults = styled.div` margin-top: 2rem; - text-align: center; font-size: 14px; - color: ${(props) => (props.theme.name === 'dark' ? '#E0E0E0' : '#505050')}; - font-weight: 500; + text-align: center; + color: ${(props) => props.theme.secondary.text}; `; + const AskAIButton = styled.button` + box-sizing: border-box; display: flex; align-items: center; - justify-content: flex-start; gap: 12px; width: calc(100% - 32px); - margin: 0 16px 16px 16px; - box-sizing: border-box; - height: 50px; - padding: 8px 24px; - border: none; - border-radius: 8px; - color: ${(props) => props.theme.text}; + margin: 0 16px 12px; + padding: 10px 12px; + font-family: inherit; + font-size: 15px; + font-weight: 500; + text-align: left; + color: ${(props) => props.theme.primary.text}; + background-color: ${(props) => props.theme.secondary.bg}; + border: 1px solid ${(props) => props.theme.hairline}; + border-radius: ${radii.md}; cursor: pointer; - font-size: 16px; - backdrop-filter: blur(16px); - -webkit-backdrop-filter: blur(16px); - background-color: ${(props) => - props.theme.name === 'dark' - ? 'rgba(255, 255, 255, 0.05)' - : 'rgba(0, 0, 0, 0.03)'}; + transition: + background-color 0.15s ease, + border-color 0.15s ease; - &:hover { - backdrop-filter: blur(20px); - -webkit-backdrop-filter: blur(20px); - background-color: ${(props) => - props.theme.name === 'dark' - ? 'rgba(255, 255, 255, 0.1)' - : 'rgba(0, 0, 0, 0.06)'}; + &:hover:not(:disabled) { + background-color: ${(props) => props.theme.accent!.soft}; + border-color: ${(props) => props.theme.accent!.soft}; + } + + &:disabled { + opacity: 0.5; + cursor: default; + } + + &:focus-visible { + outline: 2px solid ${(props) => props.theme.accent!.base}; + outline-offset: 2px; } `; const SearchHeader = styled.div` display: flex; align-items: center; - gap: 8px; + gap: 12px; margin-bottom: 12px; - padding-bottom: 12px; - border-bottom: 1px solid - ${(props) => - props.theme.name === 'dark' ? '#FFFFFF24' : 'rgba(0, 0, 0, 0.14)'}; + padding: 4px 16px 12px; + border-bottom: 1px solid ${(props) => props.theme.hairline}; `; -const TextField = styled.input` - width: calc(100% - 32px); - margin: 0 16px; - padding: 12px 16px; - border: none; +const TextField = styled.input<{ $hidden?: boolean }>` + ${(props) => (props.$hidden ? 'display: none;' : '')} + flex: 1; + min-width: 0; + padding: 8px 0; + font-family: inherit; + font-size: 18px; + color: ${(props) => props.theme.primary.text}; background-color: transparent; - color: ${(props) => props.theme.text}; - font-size: 20px; - font-weight: 400; + border: none; outline: none; - &:focus { - border-color: none; - } - &::placeholder { - color: ${(props) => - props.theme.name === 'dark' - ? 'rgba(255, 255, 255, 0.6)' - : 'rgba(0, 0, 0, 0.5)'} !important; - opacity: 100%; /* Force opacity to ensure placeholder is visible */ - font-weight: 500; + color: ${(props) => props.theme.secondary.text}; + opacity: 1; } `; const EscapeInstruction = styled.kbd` - display: flex; - align-items: center; - justify-content: center; - margin: 12px 16px 0; - padding: 4px 8px; - border-radius: 4px; - background-color: transparent; - border: 1px solid - ${(props) => - props.theme.name === 'dark' - ? 'rgba(237, 237, 237, 0.6)' - : 'rgba(23, 23, 23, 0.6)'}; - color: ${(props) => (props.theme.name === 'dark' ? '#EDEDED' : '#171717')}; - font-size: 12px; - font-family: 'Geist', sans-serif; + flex-shrink: 0; + padding: 2px 8px; + font-family: inherit; + font-size: 11px; + font-weight: 500; + line-height: 1.6; white-space: nowrap; + color: ${(props) => props.theme.secondary.text}; + background-color: ${(props) => props.theme.secondary.bg}; + border: 1px solid ${(props) => props.theme.hairline}; + border-radius: ${radii.sm}; cursor: pointer; - width: fit-content; - -webkit-appearance: none; - -moz-appearance: none; - appearance: none; +`; + +const SearchNote = styled.div` + margin: -4px 16px 8px; + font-size: 12px; + line-height: 1.5; + color: ${(props) => props.theme.danger!.text}; `; export const SearchBar = ({ @@ -362,6 +328,8 @@ export const SearchBar = ({ placeholder = 'Search or Ask AI...', width = '256px', buttonText = 'Search here', + allowedFileExtensions, + showMicButton, }: SearchBarProps) => { const [input, setInput] = React.useState(''); const [loading, setLoading] = React.useState(false); @@ -374,6 +342,21 @@ export const SearchBar = ({ null, ); const abortControllerRef = React.useRef(null); + const getSearchDraft = React.useCallback( + () => inputRef.current?.value ?? '', + [], + ); + // Deferred until the input is unhidden; hidden fields can't take focus. + const focusSearchInput = React.useCallback(() => { + window.requestAnimationFrame(() => inputRef.current?.focus()); + }, []); + const dictation = useDictation({ + enabled: Boolean(showMicButton), + getDraft: getSearchDraft, + onDraftChange: setInput, + separator: ' ', + onEnd: focusSearchInput, + }); const browserOS = getOS(); const isTouch = 'ontouchstart' in window; @@ -383,7 +366,6 @@ export const SearchBar = ({ }; React.useEffect(() => { - loadGeistFont(); const handleClickOutside = (event: MouseEvent) => { if ( containerRef.current && @@ -446,6 +428,11 @@ export const SearchBar = ({ }; }, [input]); + // Stop recording if the palette closes mid-dictation. + React.useEffect(() => { + if (!isResultVisible && dictation.isDictating) dictation.stop(); + }, [isResultVisible, dictation.isDictating, dictation.stop]); + const handleKeyDown = (event: React.KeyboardEvent) => { if (event.key === 'Enter') { event.preventDefault(); @@ -464,9 +451,8 @@ export const SearchBar = ({ }; return ( - +
- setIsResultVisible(true)} @@ -479,19 +465,47 @@ export const SearchBar = ({ setIsResultVisible(false)} /> + {dictation.isDictating && ( + + )} setInput(e.target.value)} + onChange={(e) => { + if (dictation.error) dictation.clearError(); + setInput(e.target.value); + }} onKeyDown={(e) => handleKeyDown(e)} placeholder={placeholder} autoFocus /> + {dictation.available && ( + + )} setIsResultVisible(false)}> Esc - + {dictation.error && ( + {dictation.error} + )} + Ask the AI @@ -600,6 +614,8 @@ export const SearchBar = ({ isOpen={isWidgetOpen} handleClose={handleClose} size={'large'} + allowedFileExtensions={allowedFileExtensions} + showMicButton={showMicButton} />
diff --git a/extensions/react-widget/src/components/tokens.ts b/extensions/react-widget/src/components/tokens.ts index 097cf678..b9ae788a 100644 --- a/extensions/react-widget/src/components/tokens.ts +++ b/extensions/react-widget/src/components/tokens.ts @@ -1,3 +1,69 @@ +/** Colour tokens shared by the chat widget and the search bar. */ +export const themes = { + dark: { + bg: '#222327', + text: '#fff', + primary: { + text: '#FAFAFA', + bg: '#222327', + }, + secondary: { + text: '#A1A1AA', + bg: '#33343A', + }, + shimmer: { + base: '#A1A1AA', + highlight: '#FAFAFA', + }, + accent: { + base: '#8860DB', + hover: '#9B7BE4', + strong: '#6D42C5', + contrast: '#FFFFFF', + soft: 'rgba(136, 96, 219, 0.18)', + mark: 'rgba(136, 96, 219, 0.18)', + link: '#A78BFA', + }, + hairline: 'rgba(255, 255, 255, 0.08)', + danger: { + text: '#F87171', + soft: 'rgba(248, 113, 113, 0.10)', + border: 'rgba(248, 113, 113, 0.32)', + }, + }, + light: { + bg: '#fff', + text: '#000', + primary: { + text: '#222327', + bg: '#fff', + }, + secondary: { + text: '#71717A', + bg: '#F4F4F5', + }, + shimmer: { + base: '#71717A', + highlight: '#D4D4D8', + }, + accent: { + base: '#8860DB', + hover: '#7A4FD0', + strong: '#6D42C5', + contrast: '#FFFFFF', + soft: 'rgba(136, 96, 219, 0.12)', + mark: 'rgba(136, 96, 219, 0.26)', + link: '#6D42C5', + }, + hairline: 'rgba(0, 0, 0, 0.08)', + danger: { + text: '#B91C1C', + soft: 'rgba(185, 28, 28, 0.06)', + border: 'rgba(185, 28, 28, 0.24)', + }, + }, +}; + /** Corner radii shared by the widget's chrome. */ export const radii = { sm: '8px', diff --git a/extensions/react-widget/src/hooks/useAttachments.ts b/extensions/react-widget/src/hooks/useAttachments.ts index 24f59da9..c81e9f88 100644 --- a/extensions/react-widget/src/hooks/useAttachments.ts +++ b/extensions/react-widget/src/hooks/useAttachments.ts @@ -97,7 +97,7 @@ export const useAttachments = ({ if (controller.signal.aborted) return; if (outcome.state === 'completed') { - // No id means nothing to send, whatever the task status said. + // Without an attachment id there is nothing to send. if (outcome.attachmentId) patch(id, { status: 'completed', @@ -129,8 +129,7 @@ export const useAttachments = ({ const id = generateId(); const extension = fileExtension(file.name); - // A rejected file still gets a chip: mobile pickers ignore - // `accept`, and a silent drop looks broken. + // Mobile pickers ignore accept, so rejected files still get a chip. if (!acceptedExtensions.includes(extension)) { setAttachments((prev) => [ ...prev, diff --git a/extensions/react-widget/src/hooks/useDictation.ts b/extensions/react-widget/src/hooks/useDictation.ts new file mode 100644 index 00000000..34a06711 --- /dev/null +++ b/extensions/react-widget/src/hooks/useDictation.ts @@ -0,0 +1,65 @@ +import React from 'react'; + +import type { MicButtonState } from '../components/ComposerControls'; +import { useVoiceInput, voiceInputSupported } from './useVoiceInput'; + +interface UseDictationOptions { + enabled: boolean; + /** Read when dictation starts; the transcript is appended to it. */ + getDraft: () => string; + onDraftChange: (value: string) => void; + /** '\n' for a textarea, ' ' for a single-line input. */ + separator: string; + onEnd?: () => void; +} + +/** Microphone state plus merging the live transcript into an input's draft. */ +export const useDictation = ({ + enabled, + getDraft, + onDraftChange, + separator, + onEnd, +}: UseDictationOptions) => { + const baseRef = React.useRef(''); + + const handleStart = React.useCallback(() => { + baseRef.current = getDraft(); + }, [getDraft]); + + const handleTranscript = React.useCallback( + (text: string) => { + // The base is captured once, so interim revisions only replace the + // transcript. + const base = baseRef.current; + onDraftChange( + base.trim() ? `${base.replace(/\s+$/, '')}${separator}${text}` : text, + ); + }, + [onDraftChange, separator], + ); + + const voice = useVoiceInput({ + onStart: handleStart, + onTranscript: handleTranscript, + onEnd, + }); + + const state: MicButtonState = + voice.recordingState === 'recording' || + voice.recordingState === 'transcribing' + ? voice.recordingState + : 'idle'; + + return { + // Firefox has no SpeechRecognition, nor does an insecure origin. + available: enabled && voiceInputSupported(), + state, + isDictating: state !== 'idle', + error: voice.error, + clearError: voice.clearError, + toggle: voice.toggle, + stop: voice.stop, + analyserRef: voice.analyserRef, + }; +}; diff --git a/extensions/react-widget/src/hooks/useVoiceInput.ts b/extensions/react-widget/src/hooks/useVoiceInput.ts index 516f48b3..0185c8c3 100644 --- a/extensions/react-widget/src/hooks/useVoiceInput.ts +++ b/extensions/react-widget/src/hooks/useVoiceInput.ts @@ -39,13 +39,12 @@ const recognitionClass = (): SpeechRecognitionConstructor | undefined => { export const voiceInputSupported = (): boolean => { if (typeof window === 'undefined') return false; - // Gated on a secure origin, where the failure looks identical to the API - // being absent. + // The API is only exposed on secure origins. if (!window.isSecureContext) return false; return recognitionClass() !== undefined; }; -/** Message worth showing, or null for outcomes that are not errors. */ +/** User-facing message, or null for outcomes that are not errors. */ const errorMessage = (code: string): string | null => { switch (code) { // The user pressed stop. @@ -74,7 +73,7 @@ interface UseVoiceInputOptions { onEnd?: () => void; } -/** The composer's microphone; one button toggles it. */ +/** Browser speech recognition with interim results. */ export const useVoiceInput = ({ onStart, onTranscript, @@ -84,12 +83,12 @@ export const useVoiceInput = ({ React.useState('idle'); const [error, setError] = React.useState(null); const recognitionRef = React.useRef(null); - // The Web Speech API exposes no audio, so the waveform needs a second - // capture of its own. Kept in a ref so the draw loop re-renders nothing. + // The Web Speech API exposes no audio, so the waveform uses a separate + // capture. Refs, so drawing does not re-render. const analyserRef = React.useRef(null); const audioContextRef = React.useRef(null); const levelStreamRef = React.useRef(null); - // Finals accumulate; the interim tail is rebuilt each event, never kept. + // Final results accumulate; interim text is rebuilt on each event. const finalTranscriptRef = React.useRef(''); const isMountedRef = React.useRef(true); @@ -101,23 +100,21 @@ export const useVoiceInput = ({ audioContextRef.current = null; }, []); - /** Best-effort: a browser that refuses this still transcribes fine. */ + /** Best effort: recognition still works if this capture is refused. */ const startLevelMeter = React.useCallback(async () => { const audioWindow = window as LegacyAudioWindow; const AudioContextClass = audioWindow.AudioContext ?? audioWindow.webkitAudioContext; if (!AudioContextClass || !navigator.mediaDevices?.getUserMedia) return; - // Never `await context.resume()`: without user activation that promise - // never settles at all, stranding everything below it, so no analyser is - // built. Constructed before the first await so it stays under the click's - // activation, which is insurance for browsers that do not treat a live - // capture as activation the way Chrome does. + // Never `await context.resume()`: without user activation the promise + // never settles and the analyser below is never built. Constructed before + // the first await so it is created within the click's activation. const context = new AudioContextClass(); try { const stream = await navigator.mediaDevices.getUserMedia({ audio: true }); - // Recording may have ended while permission was granted. + // Recording may have ended while permission was pending. if (!isMountedRef.current || !recognitionRef.current) { stream.getTracks().forEach((track) => track.stop()); void context.close().catch(() => undefined); @@ -127,8 +124,7 @@ export const useVoiceInput = ({ const analyser = context.createAnalyser(); analyser.fftSize = LEVEL_FFT_SIZE; analyser.smoothingTimeConstant = LEVEL_SMOOTHING; - // Not connected to the destination: that would play the microphone - // back through the page's speakers. + // Left unconnected to the destination, which would play the mic back. context.createMediaStreamSource(stream).connect(analyser); levelStreamRef.current = stream; @@ -143,8 +139,7 @@ export const useVoiceInput = ({ isMountedRef.current = true; return () => { isMountedRef.current = false; - // `abort`, not `stop`: nothing left to deliver, and it frees the - // microphone at once. + // abort() releases the microphone immediately. recognitionRef.current?.abort(); recognitionRef.current = null; stopLevelMeter(); @@ -175,8 +170,8 @@ export const useVoiceInput = ({ recognition.onresult = (event: SpeechRecognitionEvent) => { let interim = ''; - // From `resultIndex`, not 0: the API returns the whole list every - // time, and earlier results are already banked. + // Start at resultIndex: the list also contains results already + // accumulated. for ( let index = event.resultIndex; index < event.results.length; @@ -199,8 +194,7 @@ export const useVoiceInput = ({ recognition.onend = () => { recognitionRef.current = null; - // Not on stop(), so the bars stay live while the last utterance - // finalises. + // Here, so the waveform stays live until the final result arrives. stopLevelMeter(); if (!isMountedRef.current) return; // An error already set its own state; do not overwrite it. @@ -224,22 +218,26 @@ export const useVoiceInput = ({ void startLevelMeter(); }, [onEnd, onStart, onTranscript, startLevelMeter, stopLevelMeter]); + const stop = React.useCallback(() => { + if (recordingState !== 'recording') return; + // stop() lets the last utterance deliver its final result. + setRecordingState('transcribing'); + recognitionRef.current?.stop(); + }, [recordingState]); + const toggle = React.useCallback(() => { if (recordingState === 'transcribing') return; if (recordingState === 'recording') { - // `stop`, not `abort`: the last utterance still has a final result - // to deliver. - setRecordingState('transcribing'); - recognitionRef.current?.stop(); + stop(); return; } start(); - }, [recordingState, start]); + }, [recordingState, start, stop]); const clearError = React.useCallback(() => { setError(null); setRecordingState((previous) => (previous === 'error' ? 'idle' : previous)); }, []); - return { recordingState, error, toggle, clearError, analyserRef }; + return { recordingState, error, toggle, stop, clearError, analyserRef }; }; diff --git a/extensions/react-widget/src/requests/attachmentsApi.ts b/extensions/react-widget/src/requests/attachmentsApi.ts index 13117d11..829fda48 100644 --- a/extensions/react-widget/src/requests/attachmentsApi.ts +++ b/extensions/react-widget/src/requests/attachmentsApi.ts @@ -1,8 +1,8 @@ /** - * `/api/store_attachment` returns a Celery task id; the parsed attachment - * only exists once that task lands, hence the polling below. The id `/stream` - * expects is the task result's `attachment_id`, not the one the upload - * response carries. Auth is the agent key as `api_key` in the multipart body. + * `/api/store_attachment` returns a Celery task id; the attachment row only + * exists once that task finishes, hence the polling below. Sending its id to + * `/stream` any earlier finds no row, and the attachment is silently dropped. + * Auth is the agent key as `api_key` in the multipart body. */ const UPLOAD_TIMEOUT_MS = 2 * 60 * 1000; @@ -105,11 +105,11 @@ export function uploadAttachment({ } /** - * A 503 or transport failure reads as `pending`: the worker fleet is - * unreachable, which says nothing about this file. + * A 503 or transport failure is treated as `pending`: those mean the worker + * fleet is unreachable, and the task may still complete. * - * A failed task's `result` is the raw exception, which has no business on a - * third-party page, so no reason is returned and the caller supplies one. + * A failed task's `result` is the raw exception, so no reason is returned; + * the caller supplies one. */ export async function fetchTaskOutcome( taskId: string, diff --git a/extensions/react-widget/src/types/index.ts b/extensions/react-widget/src/types/index.ts index 5a07ab94..9a81af79 100644 --- a/extensions/react-widget/src/types/index.ts +++ b/extensions/react-widget/src/types/index.ts @@ -12,8 +12,6 @@ declare module 'styled-components' { text: string; bg: string; }; - /** Present only in SearchBar theme */ - name?: string; /** Gradient stops for the swept status text. */ shimmer?: { base: string; @@ -25,6 +23,8 @@ declare module 'styled-components' { strong: string; contrast: string; soft: string; + /** Background behind a search keyword match. */ + mark: string; link: string; }; hairline?: string; @@ -51,7 +51,10 @@ export type Status = 'idle' | 'loading' | 'failed'; export type FEEDBACK = 'LIKE' | 'DISLIKE'; export type AttachmentStatus = - 'uploading' | 'processing' | 'completed' | 'failed'; + | 'uploading' + | 'processing' + | 'completed' + | 'failed'; export interface Attachment { /** Client-side key for the chip; never sent to the server. */ @@ -62,7 +65,7 @@ export interface Attachment { progress: number; /** Server-side id; what `/stream` expects in its `attachments` array. */ attachmentId?: string; - /** Why it failed, when there is something worth showing. */ + /** User-facing failure reason, if any. */ error?: string; } @@ -120,17 +123,15 @@ export interface WidgetProps { defaultOpen?: boolean; /** * File extensions the composer accepts, e.g. `['.pdf', '.md', '.png']`. - * Attachments stay off until this is set, since every uploaded file is - * parsed and billed against the key owner's token budget. A leading dot is - * optional and matching is case-insensitive. + * Attachments are disabled while unset. A leading dot is optional and + * matching is case-insensitive. */ allowedFileExtensions?: string[]; /** - * Show the microphone that dictates into the input via the browser's Web - * Speech API. Off by default: outside Chromium builds with on-device - * recognition the browser forwards audio to its vendor's speech service. - * The button hides itself where the API is missing or the origin is - * insecure. + * Show a microphone that dictates into the input via the browser's Web + * Speech API. Hidden where the API is unavailable or the origin is + * insecure. Outside Chromium builds with on-device recognition, audio is + * sent to the browser vendor's speech service. */ showMicButton?: boolean; } @@ -141,7 +142,14 @@ export interface WidgetCoreProps extends WidgetProps { prefilledQuery?: string; } -export interface SearchBarProps { +/** + * Both props are forwarded to the chat opened from "Ask the AI"; + * `showMicButton` also adds a microphone to the search field. + */ +export interface SearchBarProps extends Pick< + WidgetProps, + 'allowedFileExtensions' | 'showMicButton' +> { apiHost?: string; apiKey?: string; theme?: THEME; From 02f2fa198af7ae85d13ede9d708ac902c1143c1a Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 19:26:24 +0100 Subject: [PATCH 002/130] fix: disabled STT also covers live finish; clarify model cache location --- .../content/Deploying/Development-Environment.mdx | 2 +- docsgpt/api/user/attachments/routes.py | 2 ++ tests/api/user/attachments/test_routes.py | 15 +++++++++++++++ 3 files changed, 18 insertions(+), 1 deletion(-) diff --git a/docs/content/Deploying/Development-Environment.mdx b/docs/content/Deploying/Development-Environment.mdx index f2cfe269..6c01879f 100644 --- a/docs/content/Deploying/Development-Environment.mdx +++ b/docs/content/Deploying/Development-Environment.mdx @@ -78,7 +78,7 @@ To run the DocsGPT backend locally, you'll need to set up a Python environment a 3. **Embedding Model (no action needed):** - The embedding model is downloaded automatically the first time you ingest a document, and cached under `models/` in the repository root for subsequent runs. Set `EMBEDDINGS_CACHE_DIR` to use another directory. + The embedding model is downloaded automatically the first time you ingest a document, and cached for subsequent runs under `models/` in the data home: the repository root, unless `DOCSGPT_HOME` points elsewhere. Set `EMBEDDINGS_CACHE_DIR` to use another directory. For an offline or air-gapped machine, fetch it ahead of time instead: diff --git a/docsgpt/api/user/attachments/routes.py b/docsgpt/api/user/attachments/routes.py index 44cb8a7e..7db27b99 100644 --- a/docsgpt/api/user/attachments/routes.py +++ b/docsgpt/api/user/attachments/routes.py @@ -688,6 +688,8 @@ class LiveSpeechToTextFinish(Resource): jsonify({"success": False, "message": "Authentication required"}), 401, ) + if not STTCreator.is_enabled(settings.STT_PROVIDER): + return _feature_disabled(_STT_DISABLED_MESSAGE) redis_client = _require_live_stt_redis() if hasattr(redis_client, "status_code"): diff --git a/tests/api/user/attachments/test_routes.py b/tests/api/user/attachments/test_routes.py index 8f2f2201..e7979a6d 100644 --- a/tests/api/user/attachments/test_routes.py +++ b/tests/api/user/attachments/test_routes.py @@ -1929,6 +1929,21 @@ class TestSpeechToTextDisabled: assert _get_response_json(response) == self.DISABLED mock_create_stt.assert_not_called() + def test_live_stt_finish_returns_404_without_touching_redis(self, flask_app): + from docsgpt.api.user.attachments import routes + + app = Flask(__name__) + with patch.object(routes.settings, "STT_PROVIDER", "none"), patch.object( + routes, "_require_live_stt_redis" + ) as require_redis, app.test_request_context( + "/api/stt/live/finish", method="POST", json={"session_id": "abc"} + ): + request.decoded_token = {"sub": "test_user"} + response = routes.LiveSpeechToTextFinish().post() + assert _get_response_status(response) == 404 + assert _get_response_json(response) == self.DISABLED + require_redis.assert_not_called() + # ===================================================================== # Coverage gap tests (lines 136, 256, 330, 337, 443, 457, 560, 590) From 22f759b3bf38a2f19bd84bc19c3ce81796b4d8ae Mon Sep 17 00:00:00 2001 From: Pavel Date: Wed, 16 Sep 2026 00:14:13 +0400 Subject: [PATCH 003/130] Guardrail docs --- docs/content/Agents/_meta.js | 4 + docs/content/Agents/guardrails.mdx | 384 +++++++++++++++++++++++++++++ 2 files changed, 388 insertions(+) create mode 100644 docs/content/Agents/guardrails.mdx diff --git a/docs/content/Agents/_meta.js b/docs/content/Agents/_meta.js index 238c1488..b1bccc42 100644 --- a/docs/content/Agents/_meta.js +++ b/docs/content/Agents/_meta.js @@ -3,6 +3,10 @@ export default { "title": "🤖 Agent Basics", "href": "/Agents/basics" }, + "guardrails": { + "title": "🛡️ Guardrails", + "href": "/Agents/guardrails" + }, "api": { "title": "🔌 Agent API", "href": "/Agents/api" diff --git a/docs/content/Agents/guardrails.mdx b/docs/content/Agents/guardrails.mdx new file mode 100644 index 00000000..eb645725 --- /dev/null +++ b/docs/content/Agents/guardrails.mdx @@ -0,0 +1,384 @@ +--- +title: Agent Guardrails +description: Scan user input, retrieved sources, tool results and answers with built-in checks — PII, secrets, banned terms, link policy, prompt injection, groundedness and an LLM judge — and choose to flag, redact or block. Includes the audit journal, instance floor and API reference. +--- + +import { Callout, Tabs } from 'nextra/components'; + +# Agent Guardrails 🛡️ + +Guardrails are per-agent content controls that run **inside** an agent turn, at the points where text changes hands: when the user's question arrives, when retrieved documents are about to reach the model, when a tool returns, and when the answer is on its way out. Each control pairs a **check** (a detector) with a **stage** (where it runs) and an **action** (what happens on a match). + +They apply everywhere the agent runs — the chat UI, the [Agent API](/Agents/api), the [OpenAI-compatible API](/Agents/openai-compatible), [webhooks](/Agents/webhooks), scheduled runs and the embeddable widget — and every decision is written to an audit journal you can review from the agent's logs page. + + +Guardrails are a defence-in-depth layer, not a replacement for a well-written prompt or for tool approvals. The pattern and heuristic checks are deterministic and fast; the LLM judge is semantic but costs a model call. Start in **monitor mode**, read the journal, then promote to enforcement. + + +## How a turn is scanned + +An agent turn passes through four intervention points. A control attached to a stage sees the text at that stage and can leave it alone, mask parts of it, or stop the turn. + +| Stage | What is scanned | `redact` does | `block` does | +| --- | --- | --- | --- | +| `input` | The user's question, before anything is sent to the model | The model and the stored conversation both receive the masked question | The turn ends immediately with the block message; nothing reaches the model | +| `retrieval` | The retrieved document chunks, formatted as they will appear in the prompt | Masked in the prompt **and** in the sources shown to the user and stored with the conversation | The model is told the sources were withheld by policy; the user sees `[Withheld by a content policy.]` in place of each source | +| `tool_result` | Each tool's result string, before it fans out to the model, the UI and the journal | The masked result is what the model and the user see | The result is replaced with a note telling the model it could not be used and must not speculate about its contents | +| `output` | The answer as it streams from the model | Masked **before** the text leaves the server | The stream stops with the block message; any tokens already delivered are retracted from the client and the stored message | + +A few details worth knowing: + +- **Input redaction is what gets stored.** If a PII control redacts an email address in the question, the conversation history holds the redacted version. The raw text never lands in the database. +- **Retrieval scanning covers custom prompts too.** If your prompt template interpolates documents itself (see [Customising prompts](/Guides/Customising-prompts)), the rendered documents are still scanned and the verdict is patched back into the prompt. +- **Structured output is scanned whole.** When an agent has a JSON schema, redacting mid-token would produce invalid JSON, so the complete document is buffered and scanned once. +- **Output blocks after streaming has begun are retractions.** Tokens on the wire cannot be recalled, so the server tells the client to clear the partial answer, replaces the persisted message with the block message, and clears any reasoning trace. On a reload the user sees only the block message. + +### Streaming without leaks + +Output controls run *before* a token is released. Deterministic checks hold a small lookback window (sized from the longest match any active check can produce, up to 8 KB for a PEM private key) and re-scan `held + new` on every chunk, so a card number or API key split across two stream deltas is still caught. Redacted spans are never cut in half: the release point is pulled back so a match is either fully masked or fully held. + +Remote checks (the LLM judge) cannot afford a call per token, so the guard accumulates text to a sentence boundary (around 400 characters) and evaluates whole segments. A stream that never produces a sentence boundary is force-released past a 16 KB ceiling so it cannot stall forever. + +The **groundedness** check only makes sense over a finished answer, so it is deferred to the end of the stream. A `block` from it is therefore always a retraction. + +## Actions + +| Action | Effect | Available for | +| --- | --- | --- | +| `flag` | Record the decision in the journal and the turn's activity log. Nothing about the answer changes. | Every check | +| `redact` | Replace each matched span with a mask and continue. | Checks that report spans: `pii`, `secrets`, `denylist`, `url` | +| `block` | Stop the turn and return the agent's block message. | Every check | + +Within one stage the **most restrictive outcome wins**: if two controls match and one says block, the stage blocks. If several redact, all of their spans are masked, and overlapping spans are unioned so a short match can never leave part of a longer one in the clear. + +Redaction masks are check-specific: PII uses the entity label (`[EMAIL]`, `[CREDIT_CARD]`), secrets use `[REDACTED]`, banned terms use `***`, and disallowed links use ``. + +## Enforcement modes + +| Mode | Behaviour | +| --- | --- | +| `monitor_only` (default) | Every control runs, but every action is downgraded to `flag`. Nothing is changed or blocked; the journal shows what *would* have happened. Streamed answers pass through untouched and are scanned once at the end. | +| `scan_all` | Actions are enforced as configured. | + +Monitor mode is the supported rollout path. Turn on the checks you want, run real traffic for a few days, look at the **Guardrail activity** panel for the control that is over-triggering, tune its settings, then switch to `scan_all`. + +## Built-in checks + +| Key | Label | Stages | Redacts | Remote | Typical latency | +| --- | --- | --- | --- | --- | --- | +| `pii` | Personal information | input, retrieval, tool_result, output | Yes | No | ~2 ms | +| `secrets` | Credentials and secrets | input, retrieval, tool_result, output | Yes | No | ~2 ms | +| `denylist` | Banned terms | input, retrieval, tool_result, output | Yes | No | ~1 ms | +| `url` | Link policy | input, retrieval, tool_result, output | Yes | No | ~2 ms | +| `injection` | Prompt injection (heuristic) | input, retrieval, tool_result | No | No | ~3 ms | +| `groundedness` | Grounding in sources | output | No | No | ~5 ms | +| `policy` | Custom policy (LLM judge) | input, retrieval, tool_result, output | No | Yes | ~1 s | + +The live catalog for your instance, including which checks the operator has allowed, is served by `GET /api/guardrails/catalog`. + +### `pii` — Personal information + +Pattern matching for structured identifiers. Reliable for the formats below; it does **not** find names or free-text addresses. + +| Setting | Default | Notes | +| --- | --- | --- | +| `entities` | `["EMAIL", "PHONE", "US_SSN", "CREDIT_CARD"]` | Non-empty subset of `EMAIL`, `PHONE`, `US_SSN`, `CREDIT_CARD`, `IPV4`, `IBAN` | + +Card numbers must be 13–19 digits and pass a Luhn check before they count, which keeps order numbers and long IDs from matching. Each match is reported under its entity name, so the journal tells you *which* kind of PII appeared. + +### `secrets` — Credentials and secrets + +No settings. Detects by known formats: AWS access keys, GitHub tokens, OpenAI and Anthropic keys, Slack tokens, Google API keys, JWTs, PEM private-key blocks (the whole armored block, not just the header) and generic `password=` / `api_key:` style assignments where only the value is masked. + +### `denylist` — Banned terms + +| Setting | Default | Notes | +| --- | --- | --- | +| `terms` | — (required) | 1–500 terms, each ≤ 128 characters | +| `match` | `"word"` | `"word"` matches whole words only; `"substring"` matches anywhere | +| `case_sensitive` | `false` | | + +Useful for competitor names, internal codenames, or phrases you never want an agent to repeat. + +### `url` — Link policy + +| Setting | Default | Notes | +| --- | --- | --- | +| `allow_hosts` | `[]` | Up to 200 hosts. When non-empty, any link whose host is not in the list (or a subdomain of one) is disallowed | +| `block_hosts` | `[]` | Up to 200 hosts. Links to these hosts (or their subdomains) are always disallowed | + +At least one of the two lists is required. Hosts are matched against the parsed URL authority, so `https://allowed.com@evil.tld/` resolves to `evil.tld`. A URL that cannot be parsed is treated as disallowed. + +### `injection` — Prompt injection (heuristic) + +| Setting | Default | Notes | +| --- | --- | --- | +| `min_hits` | `1` | 1–10. Number of injection-like phrases needed before the check triggers | + +Matches the phrasings that appear in real indirect-injection payloads: instruction overrides ("ignore previous instructions"), role hijacks ("you are now…"), system-prompt exfiltration ("reveal your instructions"), fake conversation turns (`system:` at the start of a line) and tool coercion ("you must immediately call the tool…"). It is most valuable at the `retrieval` and `tool_result` stages, where text an attacker may have planted arrives with the user's authority. + + +This check catches unobfuscated payloads only. A motivated attacker can evade it trivially. Pair it with the `policy` judge if you need semantic coverage. + + +### `groundedness` — Grounding in sources + +Output-only. Measures the lexical overlap between the answer and the retrieved sources using 4-word shingles and flags answers that fall below a threshold. + +| Setting | Default | Notes | +| --- | --- | --- | +| `min_overlap` | `0.3` | 0–1. Fraction of the answer's shingles that must appear in the sources | +| `min_words` | `25` | 1–1000. Shorter answers are skipped | +| `require_retrieval` | `true` | When `true`, an answer produced with **no** retrieved sources triggers with category `NO_SOURCES` | + +Lexical overlap is a proxy for support, not entailment. Keep this on `flag` until you have tuned the threshold against real traffic. When the sources contain no comparable text the check reports *not evaluated* rather than a verdict. + +### `policy` — Custom policy (LLM judge) + +Write a policy in plain language — a topic to stay off, a tone to hold, a rule to enforce — and a judge model decides whether the content breaks it. + +| Setting | Default | Notes | +| --- | --- | --- | +| `policy` | — (required) | 10–2500 characters of policy text | +| `confidence_threshold` | `0.7` | 0–1. The judge must both report a violation **and** be at least this confident | +| `max_chars` | `8000` | 200–100000. Only the first `max_chars` of the content are sent to the judge | +| `model` | `null` | Optional model id override for this control | + +The judge is the instance's own model provider, so a self-hosted deployment gets a semantic guardrail with no extra vendor account. The model used is, in order: the control's `model`, then the instance-wide `GUARDRAILS_JUDGE_MODEL`, then the model the agent is answering with. Judge calls are tagged `guardrail` in token usage so their cost shows up separately from the agent's own generation. + +The content is passed to the judge as untrusted data inside a delimited envelope, with fences and the envelope's own tags neutralised, and the judge is instructed to ignore any directions it finds inside. If the judge times out, errors, or returns something unparsable, the control reports *not evaluated* and the fail-open policy below decides what happens. + +## When a check cannot run + +A timeout, a provider error, or a missing judge model is **not** a clean pass. The control reports `not_evaluated`, the journal records it, and the agent's failure policy applies: + +| Setting | Default | Meaning | +| --- | --- | --- | +| `fail_open` | `true` | Let the turn continue when a check could not run. Set to `false` to block the turn instead whenever a `block` **or** `redact` control could not run — fail-closed exists so that unscanned text never reaches the user, and a broken PII detector would otherwise release exactly what it was there to remove | +| `timeout_ms` | `2000` | 100–60000. Deadline for the remote checks at one stage. Local pattern checks run inline and are not subject to it | + +Remote controls at one stage run in parallel under a single deadline. At most 8 remote controls run per stage; any beyond that are reported as *not evaluated*. + +## Configuring guardrails in the UI + +Open the agent in the builder and expand the **Guardrails** section. + +1. **Enable guardrails.** Nothing runs until this is on. +2. **Enforcement mode.** Leave it on *Monitor only* while you calibrate; switch to *Enforce everywhere* when the journal looks right. +3. **Checks.** Each check is a card with one chip per supported stage. Turning a chip on adds a control with the `flag` action; use the action selector on the chip to promote it to `redact` or `block`, and **Configure** to edit its settings. The card shows the approximate latency the check adds. +4. **Blocked-response message.** Up to 500 characters, shown to the user whenever a control blocks. Defaults to *"Sorry, I can't help with that request."* +5. **Continue if a check fails** and **Check timeout (ms)** map to `fail_open` and `timeout_ms`. + +Checks that cannot run without settings (`denylist`, `url`, `policy`, and `pii` with no entities selected) are marked *Not configured* and block saving until they are filled in, so a half-configured control can never be published as if it were protecting you. + + +Guardrails are the **agent owner's** policy. Team members with edit access can see the configuration but cannot change it — an editor who could clear a control would silently strip protection from everyone using the agent. Controls required by the instance floor (below) appear locked and cannot be removed. + + +Guardrails also apply to a **draft** agent in the builder preview, which is the natural place to try a control before publishing. + +## Configuring guardrails via the API + +Guardrails live under `guardrails` in the agent's `config` field. Pass `config` as a JSON string when creating or updating an agent through `POST /api/create_agent` or `PUT /api/update_agent/` (the same multipart form the builder uses). `update_agent` replaces the whole `config`; send the complete object each time. + +```json +{ + "guardrails": { + "enabled": true, + "mode": "scan_all", + "fail_open": true, + "timeout_ms": 2000, + "block_message": "Sorry, I can't help with that request.", + "controls": [ + { "check": "secrets", "stage": "output", "action": "redact" }, + { "check": "pii", "stage": "input", "action": "redact", + "settings": { "entities": ["EMAIL", "PHONE", "CREDIT_CARD"] } }, + { "check": "injection", "stage": "retrieval", "action": "block", + "settings": { "min_hits": 1 } }, + { "check": "denylist", "stage": "output", "action": "redact", + "settings": { "terms": ["Project Nimbus", "Acme Corp"], "match": "word" } }, + { "check": "url", "stage": "output", "action": "redact", + "settings": { "allow_hosts": ["docs.example.com", "example.com"] } }, + { "check": "policy", "stage": "output", "action": "block", + "settings": { + "policy": "Never give legal, medical or investment advice. Never quote pricing that is not in the retrieved sources.", + "confidence_threshold": 0.8 + } }, + { "check": "groundedness", "stage": "output", "action": "flag", + "settings": { "min_overlap": 0.3, "min_words": 25 } } + ] + } +} +``` + + + + ```bash + curl -X PUT http://localhost:7091/api/update_agent/ \ + -H "Authorization: Bearer " \ + -F 'config={"guardrails":{"enabled":true,"mode":"monitor_only","controls":[{"check":"secrets","stage":"output","action":"redact"}]}}' + ``` + + + ```python + import json, requests + + config = {"guardrails": { + "enabled": True, + "mode": "monitor_only", + "controls": [{"check": "secrets", "stage": "output", "action": "redact"}], + }} + requests.put( + "http://localhost:7091/api/update_agent/", + headers={"Authorization": "Bearer "}, + data={"config": json.dumps(config)}, + ).raise_for_status() + ``` + + + +Each control accepts `check`, `stage`, `action` (default `flag`), `enabled` (default `true`) and `settings`. Omitted top-level fields take the defaults shown above; `enabled` defaults to `false`. + +Writes are validated **strictly**. The request is rejected with HTTP 400 and the message *"Invalid config: one or more guardrail controls failed validation."* when: + +- a `check` is unknown, or is not allowed by the instance's `GUARDRAILS_CHECKS_ENABLED`; +- a check is attached to a stage it does not support (for example `groundedness` at `input`); +- `redact` is requested on a check that reports no spans (`injection`, `groundedness`, `policy`); +- a control's settings are out of range, or a required setting is missing; +- the same `(check, stage)` pair appears twice, or there are more than 50 controls; +- `mode`, `timeout_ms` or `block_message` is out of bounds. + +Reads are **lenient**: a stored control that has stopped validating (its check was disallowed by the operator, or removed in an upgrade) is dropped on its own and logged, and the remaining controls keep running. Agent export files carry `config`; on import an invalid guardrails block is dropped with a warning rather than failing the import. + +### What a blocked turn looks like to a client + +On the streaming endpoints, a block produces two final events. The first tells the client to retract anything it has rendered; the second carries the operator's block message as a user-facing error: + +```json +{"type": "guardrail", "guardrail": {"stage": "output", "categories": ["AWS_ACCESS_KEY"], "checks": ["secrets"]}, "retract": true} +{"type": "error", "error": "Sorry, I can't help with that request."} +``` + +The DocsGPT chat UI and the React widget handle both. For webhook and scheduled runs, the run is recorded with the block message and none of the blocked text is stored. + +## The audit journal + +Every control that **triggers**, and every control that **could not run**, writes a row to the `guardrail_events` table — in both enforcement modes, and for `flag` actions as well as `redact` and `block`. A streamed answer that re-matches the same span on every chunk produces one row, not one per chunk. Rows also flow into the turn's activity log under the `guardrail` component, so a decision is visible next to the tool calls and retrieval it belongs to. + +### Guardrail activity panel + +Open an agent's **Logs** page. Below the usage logs, the **Guardrail activity** panel shows, for a trailing window of 7, 30 or 90 days: + +- four totals — **Blocked**, **Redacted**, **Flagged** and **Not evaluated** — because "we refused to answer", "we masked something", "we noticed something" and "a check silently stopped working" are four different problems; +- a per-check breakdown, so you can see which control is doing the firing; +- the most recent 100 decisions, filterable by check and outcome, each with its stage, category (`EMAIL`, `INSTRUCTION_OVERRIDE`, `UNGROUNDED`, …) and the detector's one-line detail. + +### Journal API + +| Endpoint | Purpose | +| --- | --- | +| `GET /api/guardrails/catalog` | Available checks (with stages, redaction support, latency hint and remote flag), stages, modes, allowed actions per stage, PII entity names, the default block message, and which `(check, stage)` pairs the instance floor claims. | +| `GET /api/guardrails/events?agent_id=&limit=100&offset=0` | Decisions for one agent, newest first. `limit` is capped at 500. | +| `GET /api/guardrails/summary?days=30&agent_id=` | Totals and a `breakdown` grouped by check, stage, action, outcome and category. `agent_id` is optional; `days` is capped at 365. | + +All three require a user token. Event rows are scoped to the **requesting user**: on a shared agent you see the decisions made on your own conversations, not other members'. Responses never include the agent's API key or the matched text. + +### What is stored, and for how long + +By default the journal records *that* something matched — check, stage, action, outcome, category, a score where the check produces one, a match count and a short detail string — but **not the text**. Pre-redaction text is exactly what a PII control exists to keep out of storage. Set `GUARDRAILS_STORE_SCANNED_TEXT=true` to persist a sample of the first matched value (up to 200 characters) alongside each row for forensic review; it is stored in the database but never returned by the API. + +Rows older than `GUARDRAILS_EVENTS_RETENTION_DAYS` (default 30) are purged by a daily Celery beat task. The `message_id` link is set to `NULL` rather than cascading when a conversation is deleted, so the compliance trail outlives the conversation it came from. + +## Instance settings + +Operators control guardrails deployment-wide with these settings: + +| Setting | Default | Purpose | +| --- | --- | --- | +| `GUARDRAILS_ENABLED` | `true` | Master switch. `false` disables every stage on every agent; the builder shows a notice explaining that nothing configured will run. | +| `GUARDRAILS_CHECKS_ENABLED` | `[]` | Allowlist of check keys. Empty means every registered check. A disallowed check cannot be saved through the API, and existing controls that use it are dropped on read. | +| `GUARDRAILS_FLOOR` | `{}` | A guardrails config fragment every agent inherits and cannot weaken. See below. | +| `GUARDRAILS_JUDGE_MODEL` | unset | Model id for `policy` controls that do not set their own `model`. Unset reuses the agent's model. | +| `GUARDRAILS_STORE_SCANNED_TEXT` | `false` | Persist a sample of matched text with each journal row. | +| `GUARDRAILS_EVENTS_RETENTION_DAYS` | `30` | Journal retention. Minimum 1. | + +List and dict settings are read from the environment as JSON, for example: + +```bash +GUARDRAILS_CHECKS_ENABLED='["pii", "secrets", "denylist", "url", "injection", "groundedness"]' +``` + + +Leaving `policy` out of `GUARDRAILS_CHECKS_ENABLED` is how an air-gapped or privacy-sensitive deployment guarantees that no user text is sent to a judge model, regardless of what agent owners configure. + + +### The instance floor + +`GUARDRAILS_FLOOR` lets an operator impose a minimum policy on every agent. It uses the same shape as an agent's `guardrails` object and **must include `"enabled": true`** — without it the floor parses but applies to nothing, and a warning is logged. + +```bash +GUARDRAILS_FLOOR='{ + "enabled": true, + "mode": "scan_all", + "fail_open": false, + "controls": [ + {"check": "secrets", "stage": "output", "action": "redact"}, + {"check": "secrets", "stage": "tool_result", "action": "redact"}, + {"check": "injection", "stage": "retrieval", "action": "block"} + ] +}' +``` + +The floor is merged into each agent's own configuration at run time. An agent may tighten, never loosen: + +- Guardrails are forced on for every agent, even one whose owner never enabled them. +- If the floor's mode is `scan_all`, the merged mode is `scan_all`. Otherwise the agent's mode stands. +- If the floor sets `fail_open: false`, the agent is fail-closed. The merged `timeout_ms` is the larger of the two. +- Floor controls the agent does not define are added. +- Where both define the same `(check, stage)`, the floor's **settings** are authoritative and the **stricter action** wins (`block` > `redact` > `flag`). The two settings dicts are deliberately not merged: adding to `denylist.terms` tightens, but adding to `url.allow_hosts` loosens, so an agent that could edit floor settings could always find a loosening edit. An agent that needs different settings attaches its own control at a stage the floor does not claim. + +In the builder, floor controls appear as active and locked. The catalog exposes only which `(check, stage)` pairs the floor claims and their action — the floor's settings (banned-term lists, policy prompts) stay server-side so they cannot be read and evaded by any authenticated user. + +The floor also applies to the individual AI Agent nodes inside a workflow, which do not otherwise carry per-agent controls (see below). + +## Scope and limitations + +- **Workflow agents** run the `input` stage with their own controls, but the AI Agent nodes inside the workflow run only the instance floor, not the parent agent's controls. Aggregate output guarding across a workflow is not yet wired. +- **Pattern checks match formats, not meaning.** `pii` does not find names; `injection` misses obfuscated payloads; `groundedness` measures word overlap, not truth. Use `policy` where you need a semantic judgement, and keep an eye on the *Not evaluated* count when you do. +- **Redaction is best-effort against structured identifiers.** A value that does not match a known format passes through. Treat `redact` as a safety net, not as a data-loss-prevention guarantee. +- **Latency.** Local checks add low single-digit milliseconds. A `policy` control adds a model round trip per scanned segment, and on the `output` stage it holds the stream until each sentence boundary has been judged. + +## Extending: writing your own check + +Checks are plain Python classes registered with a small registry, mirroring how chunkers and retrievers are pluggable. Subclass `GuardrailCheck` from `docsgpt.guardrails`, declare the stages you support, implement `scan`, and register it: + +```python +from docsgpt.guardrails import GuardrailCheck, GuardrailCreator, Stage +from docsgpt.guardrails.types import CheckOutcome, Span + + +class TicketIdCheck(GuardrailCheck): + name = "ticket_id" + label = "Internal ticket ids" + description = "Masks references to internal tracker tickets." + supported_stages = {Stage.OUTPUT, Stage.TOOL_RESULT} + supports_redaction = True # scan() reports spans + latency_hint_ms = 1 + max_match_chars = 16 # longest match; sizes the streaming window + + def scan(self, text, stage, context): + import re + spans = [ + Span(m.start(), m.end(), "TICKET_ID", replacement="[TICKET]") + for m in re.finditer(r"\bOPS-\d{3,6}\b", text) + ] + if not spans: + return CheckOutcome.clean() + return CheckOutcome.hit(categories=["TICKET_ID"], spans=spans, + detail=f"{len(spans)} ticket id(s)") + + +GuardrailCreator.register(TicketIdCheck.name, TicketIdCheck) +``` + +Override `validate_settings` to strictly validate and normalise per-control settings on write, set `remote = True` for anything that makes a network call (so it runs under the stage deadline), and set `requires_complete_text = True` if the verdict is only meaningful over a finished answer. `max_match_chars` must cover the longest span the check can report, or the streaming guard may release the tail of a match before scanning it. Once registered, the check appears in the catalog and the builder automatically. From 95fd4bdefe075149347eab6cf458d52f1bcf850a Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 22:41:53 +0100 Subject: [PATCH 004/130] feat: serve the UI from the backend image, one-port standalone stack The backend image builds the web UI with scripts/build_frontend.sh and serves it through docsgpt/ui.py, so the standalone Compose file drops the frontend container. UI and API share port 7091, published on 127.0.0.1 unless DOCSGPT_BIND says otherwise. POSTGRES_PASSWORD is configurable, and an optional https profile puts Caddy in front of a public domain. docker-image-verify.yml starts the standalone stack on the image it built and checks the API, the UI, /config.js and a client-side route on one port. --- .dockerignore | 19 ++++- .github/workflows/docker-image-verify.yml | 48 ++++++++++- deployment/docker-compose-hub.yaml | 4 +- deployment/docker-compose-standalone.yaml | 89 +++++++++++++-------- docs/content/Deploying/Air-Gapped.mdx | 5 +- docs/content/Deploying/Docker-Deploying.mdx | 86 ++++++++++++++++++-- docs/content/changelog.mdx | 14 ++++ docs/content/upgrading.mdx | 4 + docsgpt/Dockerfile | 18 ++++- 9 files changed, 236 insertions(+), 51 deletions(-) diff --git a/.dockerignore b/.dockerignore index fd316071..50a5237f 100644 --- a/.dockerignore +++ b/.dockerignore @@ -1,11 +1,14 @@ # Build context for docsgpt/Dockerfile is the repository root, so the image can # carry the `application` import alias next to the `docsgpt` package. Allow only -# what the image needs; everything else (frontend, docs, tests, venvs) stays out. +# what the image needs; everything else (docs, tests, venvs) stays out. * !docsgpt/ # Only the alias file: an upgraded checkout may still hold gitignored # application/{inputs,indexes,vectors,.env} from the old layout. !application/__init__.py +# The web UI's source and the script that builds it (the `ui` stage). +!frontend/ +!scripts/build_frontend.sh # Inside the package: caches, local runtime data and secrets never ship. **/__pycache__/ @@ -23,5 +26,17 @@ docsgpt/*.pkl docsgpt/.env docsgpt/.env.* docsgpt/Dockerfile -# The backend image serves no UI (the frontend image does); a local UI build stays out. +# A local UI build stays out: the `ui` stage builds the one the image serves. docsgpt/static/ + +# Inside the frontend: installed packages and local builds are redone in the +# `ui` stage, and local VITE_* overrides must not be baked into a published image. +frontend/node_modules/ +frontend/dist/ +frontend/.env.local +frontend/.env.*.local +frontend/*.log +# The frontend image's own build files; changing them must not rebuild this UI. +frontend/Dockerfile +frontend/.dockerignore +frontend/docker/ diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index 911581fe..fec67707 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -3,6 +3,8 @@ name: Verify the Docker image works offline # Builds the backend image and runs its offline check with networking off, so # a change that reintroduces a first-request download (a tokenizer, tiktoken's # encoding, an embedding model) fails here instead of in an air-gapped install. +# Then starts the standalone Compose stack on the image and checks that the one +# published port serves both the API and the web UI. on: workflow_dispatch: @@ -17,6 +19,10 @@ on: - 'docsgpt/vectorstore/model_registry.py' - 'docsgpt/parser/tokenization.py' - 'docsgpt/vectorstore/embeddings_local.py' + - 'docsgpt/ui.py' + - 'frontend/**' + - 'scripts/build_frontend.sh' + - 'deployment/docker-compose-standalone.yaml' - '.github/workflows/docker-image-verify.yml' permissions: @@ -45,7 +51,8 @@ jobs: context: . platforms: linux/amd64 load: true - tags: docsgpt:verify${{ matrix.variant }} + # The name the standalone Compose file runs, under a tag no registry has. + tags: arc53/docsgpt:verify${{ matrix.variant }} build-args: | EXTRAS=${{ matrix.variant == '-docling' && 'docling' || '' }} INSTALL_TESSERACT=${{ matrix.variant == '-docling' && 'true' || 'false' }} @@ -54,14 +61,49 @@ jobs: - name: Image size env: - IMAGE: docsgpt:verify${{ matrix.variant }} + IMAGE: arc53/docsgpt:verify${{ matrix.variant }} run: | docker image inspect "$IMAGE" --format '{{.Size}}' | awk '{printf "uncompressed: %.2f GB\n", $1/1e9}' docker history "$IMAGE" --format '{{.Size}}\t{{.CreatedBy}}' | head -20 - name: Offline verification (no network) env: - IMAGE: docsgpt:verify${{ matrix.variant }} + IMAGE: arc53/docsgpt:verify${{ matrix.variant }} run: | docker run --rm --network none "$IMAGE" \ python -m docsgpt.scripts.verify_offline + + - name: The standalone stack serves the API and the UI on one port + env: + DOCSGPT_IMAGE_TAG: verify + DOCSGPT_IMAGE_VARIANT: ${{ matrix.variant }} + run: | + set -euo pipefail + # --pull missing keeps the image built above; postgres and redis are pulled. + docker compose -f deployment/docker-compose-standalone.yaml up -d --pull missing backend + base=http://127.0.0.1:7091 + for _ in $(seq 1 90); do + if curl -fsS "$base/api/health" >/dev/null 2>&1; then break; fi + sleep 2 + done + curl -fsS "$base/api/health" + echo + curl -fsS "$base/" | grep -q 'src="/config.js"' + curl -fsS "$base/config.js" | grep -q 'window.__DOCSGPT_ENV__' + # A client-side route falls back to the UI's index.html. + curl -fsS "$base/settings" | grep -q 'src="/config.js"' + echo "API and UI served on $base" + + - name: Stack logs + if: failure() + env: + DOCSGPT_IMAGE_TAG: verify + DOCSGPT_IMAGE_VARIANT: ${{ matrix.variant }} + run: docker compose -f deployment/docker-compose-standalone.yaml logs --no-color + + - name: Stop the stack + if: always() + env: + DOCSGPT_IMAGE_TAG: verify + DOCSGPT_IMAGE_VARIANT: ${{ matrix.variant }} + run: docker compose -f deployment/docker-compose-standalone.yaml down -v diff --git a/deployment/docker-compose-hub.yaml b/deployment/docker-compose-hub.yaml index 0f2a8a80..ce81ec95 100644 --- a/deployment/docker-compose-hub.yaml +++ b/deployment/docker-compose-hub.yaml @@ -2,8 +2,8 @@ # DOCSGPT_IMAGE_TAG develop (default, follows main) or a release, e.g. 0.20.0 # DOCSGPT_IMAGE_VARIANT empty (default, slim) or -docling: docling parser engine, # its models, and tesseract baked in (OCR-ready) -# Set them in ../.env or the shell. deployment/docker-compose-standalone.yaml is -# the same stack without a git checkout. +# Set them in ../.env or the shell. deployment/docker-compose-standalone.yaml runs +# the same images without a git checkout, with the backend serving the UI on one port. name: docsgpt-oss services: diff --git a/deployment/docker-compose-standalone.yaml b/deployment/docker-compose-standalone.yaml index 06d2ce7a..ca69253b 100644 --- a/deployment/docker-compose-standalone.yaml +++ b/deployment/docker-compose-standalone.yaml @@ -3,52 +3,43 @@ # curl -fsSLO https://raw.githubusercontent.com/arc53/DocsGPT/main/deployment/docker-compose-standalone.yaml # printf 'LLM_PROVIDER=docsgpt\nVITE_API_STREAMING=true\nINTERNAL_KEY=%s\n' "$(openssl rand -hex 16)" > .env # docker compose -f docker-compose-standalone.yaml up -d +# open http://localhost:7091 # # INTERNAL_KEY is the shared secret the worker uses to hand finished indexes to # the API; without it every ingest fails with a 401 (setup.sh generates one). -# open http://localhost:5173 # -# Every release also attaches this file as an asset. Settings come from .env -# next to this file (any DocsGPT setting; the compose-internal service URLs -# below take precedence). Data lives in named volumes, so `docker compose -# down` keeps it and `docker compose down -v` removes it. +# The backend image serves the web UI and the API on one port. Every release +# also attaches this file as an asset. Settings come from .env next to this +# file (any DocsGPT setting, VITE_* included; the compose-internal service URLs +# below take precedence). Data lives in named volumes, so `docker compose down` +# keeps it and `docker compose down -v` removes it. # # DOCSGPT_IMAGE_TAG release to run, e.g. 0.20.0 (default: latest release); # develop follows the main branch # DOCSGPT_IMAGE_VARIANT empty (slim, default) or -docling: docling parser # engine, its models, and tesseract baked in (OCR-ready) +# DOCSGPT_BIND interface the port is published on: 127.0.0.1 (default, +# this machine only) or 0.0.0.0 (every interface; set +# AUTH_TYPE, see the DocsGPT settings guide) +# DOCSGPT_PORT host port for the UI and API (default: 7091) +# POSTGRES_PASSWORD database password (default: docsgpt). Read when the +# postgres volume is first created; changing it later +# does not change the existing database's password. +# Use URL-safe characters (e.g. openssl rand -hex 24). +# DOCSGPT_DOMAIN public domain for the `https` profile (below) # EMBEDDINGS_NAME defaults to granite here (this stack always starts on # fresh volumes, so there is no older index to keep # compatible); the code default stays mpnet for upgrades. +# +# HTTPS for a public domain: point the domain's DNS at this machine, open ports +# 80 and 443, then +# DOCSGPT_DOMAIN=docs.example.com docker compose -f docker-compose-standalone.yaml --profile https up -d +# Caddy obtains and renews the certificate and proxies to the backend. Putting +# DOCSGPT_DOMAIN and COMPOSE_PROFILES=https in .env instead makes every later +# `up`, `down` and `logs` include Caddy without the flag. name: docsgpt services: - frontend: - image: arc53/docsgpt-fe:${DOCSGPT_IMAGE_TAG:-latest} - env_file: - - path: .env - required: false - environment: - # Every VITE_* the app reads. A bare name is passed through only when it is - # set in the shell or the --env-file, so an unset one does not reach the - # container as an empty string and override the image's own default. - - VITE_API_HOST=${VITE_API_HOST:-http://localhost:7091} - - VITE_API_STREAMING=${VITE_API_STREAMING:-true} - - VITE_BASE_URL - - VITE_GOOGLE_CLIENT_ID - - VITE_GOOGLE_PICKER_API_KEY - - VITE_SHARE_POINT_CLIENT_ID - - VITE_CONFLUENCE_CLIENT_ID - - VITE_NOTIFICATION_TEXT - - VITE_NOTIFICATION_LINK - - VITE_ENABLE_VOICE_INPUT - - VITE_DISABLE_SOURCE_FE - - VITE_USE_V - ports: - - "5173:5173" - depends_on: - - backend - backend: image: arc53/docsgpt:${DOCSGPT_IMAGE_TAG:-latest}${DOCSGPT_IMAGE_VARIANT:-} # Same as docker-compose-hub.yaml: the data volumes are written by root so @@ -62,10 +53,10 @@ services: - CELERY_BROKER_URL=redis://redis:6379/0 - CELERY_RESULT_BACKEND=redis://redis:6379/1 - CACHE_REDIS_URL=redis://redis:6379/2 - - POSTGRES_URI=postgresql://docsgpt:docsgpt@postgres:5432/docsgpt + - POSTGRES_URI=postgresql://docsgpt:${POSTGRES_PASSWORD:-docsgpt}@postgres:5432/docsgpt - EMBEDDINGS_NAME=${EMBEDDINGS_NAME:-ibm-granite/granite-embedding-311m-multilingual-r2} ports: - - "7091:7091" + - "${DOCSGPT_BIND:-127.0.0.1}:${DOCSGPT_PORT:-7091}:7091" volumes: - indexes:/app/indexes - inputs:/app/inputs @@ -90,7 +81,7 @@ services: - CELERY_BROKER_URL=redis://redis:6379/0 - CELERY_RESULT_BACKEND=redis://redis:6379/1 - CACHE_REDIS_URL=redis://redis:6379/2 - - POSTGRES_URI=postgresql://docsgpt:docsgpt@postgres:5432/docsgpt + - POSTGRES_URI=postgresql://docsgpt:${POSTGRES_PASSWORD:-docsgpt}@postgres:5432/docsgpt - API_URL=http://backend:7091 - EMBEDDINGS_NAME=${EMBEDDINGS_NAME:-ibm-granite/granite-embedding-311m-multilingual-r2} volumes: @@ -112,7 +103,7 @@ services: image: postgres:16-alpine environment: - POSTGRES_USER=docsgpt - - POSTGRES_PASSWORD=docsgpt + - POSTGRES_PASSWORD=${POSTGRES_PASSWORD:-docsgpt} - POSTGRES_DB=docsgpt volumes: - postgres_data:/var/lib/postgresql/data @@ -123,8 +114,36 @@ services: retries: 10 restart: unless-stopped + caddy: + image: caddy:2-alpine + profiles: [https] + environment: + - DOCSGPT_DOMAIN=${DOCSGPT_DOMAIN:-} + # The domain is checked here rather than with ${DOCSGPT_DOMAIN:?}: compose + # interpolates every service, so a required variable would break the stack + # for everyone who does not use this profile. + entrypoint: ["/bin/sh", "-c"] + command: + - >- + if [ -z "$$DOCSGPT_DOMAIN" ]; then + echo "caddy: set DOCSGPT_DOMAIN to the public domain" >&2; exit 1; + fi; + exec caddy reverse-proxy --from "$$DOCSGPT_DOMAIN" --to backend:7091 + ports: + - "80:80" + - "443:443" + - "443:443/udp" + volumes: + - caddy_data:/data + - caddy_config:/config + depends_on: + - backend + restart: unless-stopped + volumes: indexes: inputs: vectors: postgres_data: + caddy_data: + caddy_config: diff --git a/docs/content/Deploying/Air-Gapped.mdx b/docs/content/Deploying/Air-Gapped.mdx index a564109d..7c7cfd8f 100644 --- a/docs/content/Deploying/Air-Gapped.mdx +++ b/docs/content/Deploying/Air-Gapped.mdx @@ -39,13 +39,14 @@ On a machine with internet access, pull the images and save them to one file: ```bash TAG=latest # or a release, e.g. 0.19.0 docker pull arc53/docsgpt:$TAG -docker pull arc53/docsgpt-fe:$TAG docker pull redis:6-alpine docker pull postgres:16-alpine docker save -o docsgpt-images.tar \ - arc53/docsgpt:$TAG arc53/docsgpt-fe:$TAG redis:6-alpine postgres:16-alpine + arc53/docsgpt:$TAG redis:6-alpine postgres:16-alpine ``` +The backend image serves the web UI too, so the standalone stack needs no frontend image. Add `arc53/docsgpt-fe:$TAG` only if you run the checkout Compose files or Kubernetes, which use it. + Copy `docsgpt-images.tar` and the [standalone Compose file](/Deploying/Docker-Deploying#quickest-setup-pre-built-images-no-checkout) into the air-gapped network, then load the images (or push them to your internal registry): ```bash diff --git a/docs/content/Deploying/Docker-Deploying.mdx b/docs/content/Deploying/Docker-Deploying.mdx index 00932c35..b9e24c90 100644 --- a/docs/content/Deploying/Docker-Deploying.mdx +++ b/docs/content/Deploying/Docker-Deploying.mdx @@ -21,10 +21,13 @@ Docker is the recommended method for deploying DocsGPT, providing a consistent a Every release publishes ready-to-run images to Docker Hub (`arc53/docsgpt`, `arc53/docsgpt-fe`) and GitHub Container Registry (`ghcr.io/arc53/docsgpt`, -`ghcr.io/arc53/docsgpt-fe`) for `linux/amd64` and `linux/arm64`. The images -contain everything the default configuration needs (embedding models, -tokenizers, tiktoken's encoding), so a fresh container makes no downloads on -first use. You do not need the source tree to run them: +`ghcr.io/arc53/docsgpt-fe`) for `linux/amd64` and `linux/arm64`. +`arc53/docsgpt` runs the API, serves the web UI and runs the worker; +`arc53/docsgpt-fe` is the separate frontend image the checkout Compose files +and Kubernetes use. The images contain everything the default configuration +needs (embedding models, tokenizers, tiktoken's encoding), so a fresh +container makes no downloads on first use. You do not need the source tree to +run them: 1. **Download the standalone Compose file** (also attached to every [release](https://github.com/arc53/DocsGPT/releases)): @@ -52,7 +55,9 @@ first use. You do not need the source tree to run them: docker compose -f docker-compose-standalone.yaml up -d ``` - Then open [http://localhost:5173/](http://localhost:5173/). Data lives in + Then open [http://localhost:7091/](http://localhost:7091/). The web UI and + the API share that port, which is published on `127.0.0.1`: only this + machine can reach it until you change `DOCSGPT_BIND` (below). Data lives in named Docker volumes; `docker compose -f docker-compose-standalone.yaml down` keeps it and `down -v` removes it. @@ -66,6 +71,77 @@ the shell, e.g. `DOCSGPT_IMAGE_TAG=0.20.0 DOCSGPT_IMAGE_VARIANT=-docling`. The same two variables drive `deployment/docker-compose-hub.yaml` in a checkout. +### Opening it from other machines + +Publish the port on every interface and turn on authentication in `.env`: + +```bash +DOCSGPT_BIND=0.0.0.0 +AUTH_TYPE=simple_jwt +JWT_SECRET_KEY= +``` + +Then run `docker compose -f docker-compose-standalone.yaml up -d` again. The UI +takes its API address from the page it was loaded from, so +`http://:7091/` works without further settings. Without +`AUTH_TYPE`, anyone who can reach the port can use DocsGPT. + +With `simple_jwt` the UI asks for a token, which the backend prints when it +starts: `docker compose -f docker-compose-standalone.yaml logs backend | grep "Simple JWT"`. +The token is signed with `JWT_SECRET_KEY`. Without that setting each container +generates its own secret, and a re-created container (after `pull` or a +settings change) gets a new one and so a new token. Over plain HTTP the token +travels as readable text; outside a trusted network, use HTTPS as below. +`DOCSGPT_PORT` changes the host port (default `7091`). See +[Authentication Settings](/Deploying/DocsGPT-Settings#authentication-settings) for the other modes. + +### HTTPS with your own domain + +The Compose file has an optional Caddy service that obtains and renews a +Let's Encrypt certificate and proxies to the backend. + +1. Point the domain's DNS records at the machine and open ports 80 and 443. +2. Add to `.env`: + + ```bash + COMPOSE_PROFILES=https + DOCSGPT_DOMAIN=docs.example.com + AUTH_TYPE=simple_jwt + JWT_SECRET_KEY= + ``` + +3. Run `docker compose -f docker-compose-standalone.yaml up -d` and open + `https://docs.example.com/`. + +`COMPOSE_PROFILES=https` in `.env` makes every later `up`, `down` and `logs` +include Caddy. Leave `DOCSGPT_BIND` at its default: Caddy reaches the backend +over the Compose network. + +### Database password + +The Postgres password defaults to `docsgpt`; the database is only reachable +inside the Compose network. To use your own, set `POSTGRES_PASSWORD` in `.env` +before the first start, with URL-safe characters (e.g. `openssl rand -hex 24`). +Postgres reads it only when its volume is created, so changing it later does +not change the existing database's password. + +### Upgrading from an earlier standalone file + +Before this change the standalone file ran a separate frontend container on +port 5173 and published both ports on every interface. After downloading the +new file: + +```bash +docker compose -f docker-compose-standalone.yaml pull +docker compose -f docker-compose-standalone.yaml up -d --remove-orphans +``` + +`--remove-orphans` removes the old frontend container. Open port 7091 instead +of 5173. Your data volumes are unchanged. If you opened DocsGPT from other +machines, follow [Opening it from other machines](#opening-it-from-other-machines), +and remove `VITE_API_HOST` from `.env` if it points at `localhost`: the UI +would otherwise keep calling the visitor's own machine. + ## Using the Source Checkout With a clone of the repository, `deployment/docker-compose-hub.yaml` runs the diff --git a/docs/content/changelog.mdx b/docs/content/changelog.mdx index 9defbbd1..0aa9324a 100644 --- a/docs/content/changelog.mdx +++ b/docs/content/changelog.mdx @@ -11,6 +11,20 @@ The notable changes in each release. Every release on GitHub also carries [auto-generated notes](https://github.com/arc53/DocsGPT/releases) listing every merged pull request, and [Upgrading](/upgrading) covers the steps an existing deployment has to take. +## Unreleased + +### The standalone Docker stack runs on one port + +The `arc53/docsgpt` image now serves the web UI next to the API, the way `docsgpt api` does from +the Python package. `docker-compose-standalone.yaml` no longer runs a frontend container: the UI +and the API share port 7091, published on `127.0.0.1` by default. The UI takes its API address +from the page it was loaded from, so opening the stack from another machine works without +setting `VITE_API_HOST`. New Compose settings: `DOCSGPT_BIND` and `DOCSGPT_PORT` for where the +port is published, `POSTGRES_PASSWORD`, and an `https` profile that puts Caddy with an automatic +certificate in front of a public domain. The `arc53/docsgpt-fe` image is still published for the +checkout Compose files and Kubernetes. See +[Upgrading from an earlier standalone file](/Deploying/Docker-Deploying#upgrading-from-an-earlier-standalone-file). + ## 0.20.0 ### DocsGPT installs from PyPI diff --git a/docs/content/upgrading.mdx b/docs/content/upgrading.mdx index e00029a5..7df9a0e3 100644 --- a/docs/content/upgrading.mdx +++ b/docs/content/upgrading.mdx @@ -11,6 +11,10 @@ import { Callout } from 'nextra/components' **Upgrading from 0.16.x?** User data moved from MongoDB to Postgres in 0.17.0. Follow the [Postgres Migration guide](/Deploying/Postgres-Migration) before running `docker compose pull` or `git pull` — existing deployments will not start cleanly without it. +## Standalone Compose file: one port + +The standalone Compose file (`docker-compose-standalone.yaml`) no longer runs a frontend container. The backend image serves the web UI on port 7091, and the port is published on `127.0.0.1` unless you set `DOCSGPT_BIND`. After downloading the new file, start it with `--remove-orphans` to remove the old frontend container, then open port 7091 instead of 5173. If you opened DocsGPT from other machines, see [Upgrading from an earlier standalone file](/Deploying/Docker-Deploying#upgrading-from-an-earlier-standalone-file). The checkout Compose files and the Kubernetes manifests are unchanged. + ## Embedding models DocsGPT now runs embeddings through [FastEmbed](https://github.com/qdrant/fastembed) (ONNX Runtime) instead of SentenceTransformer. The models are the same and the vectors are identical, so **your existing index needs no action** — `all-mpnet-base-v2` keeps working exactly as before. diff --git a/docsgpt/Dockerfile b/docsgpt/Dockerfile index 13fcac37..b0e9ae8f 100644 --- a/docsgpt/Dockerfile +++ b/docsgpt/Dockerfile @@ -1,4 +1,4 @@ -# DocsGPT backend image. +# DocsGPT image: the API, the web UI it serves, and the Celery worker. # # Build args: # EXTRAS comma-separated optional extras to bake in, matching the @@ -15,6 +15,19 @@ # docling's layout/table/OCR models. `python -m docsgpt.scripts.verify_offline` # under `docker run --network none` proves it. +# The web UI, built by the same script the Python package build runs. The API +# serves it from docsgpt/static (docsgpt/ui.py). The output is static files, so +# this stage runs on the build machine's platform whatever the target is. +FROM --platform=$BUILDPLATFORM node:22-bookworm-slim AS ui + +WORKDIR /src +COPY frontend/package.json frontend/package-lock.json frontend/ +RUN cd frontend && npm ci --include=dev --no-audit --no-fund +COPY frontend/ frontend/ +COPY scripts/build_frontend.sh scripts/ +RUN mkdir docsgpt && bash scripts/build_frontend.sh + + FROM ubuntu:24.04 AS builder ENV DEBIAN_FRONTEND=noninteractive @@ -87,7 +100,7 @@ RUN if [ "$INSTALL_TESSERACT" = "true" ]; then \ LABEL org.opencontainers.image.source="https://github.com/arc53/DocsGPT" \ org.opencontainers.image.title="DocsGPT" \ - org.opencontainers.image.description="DocsGPT backend: API and Celery worker" \ + org.opencontainers.image.description="DocsGPT: API, web UI and Celery worker" \ org.opencontainers.image.licenses="MIT" WORKDIR /app @@ -136,6 +149,7 @@ RUN if python -c "import docling" 2>/dev/null; then \ fi COPY --chown=appuser:appuser docsgpt /app/docsgpt +COPY --from=ui --chown=appuser:appuser /src/docsgpt/static /app/docsgpt/static # One-release alias so `-A application.app.celery` style entry points keep working. COPY --chown=appuser:appuser application/__init__.py /app/application/__init__.py From ac06527a8d8f27f026036d1d6ef0b0276e5496ad Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 22:45:54 +0100 Subject: [PATCH 005/130] feat: installed package keeps its data in ~/.docsgpt/server Outside a checkout the data home was the working directory, so running `docsgpt api` from another folder silently used different settings and data. It is now ~/.docsgpt/server (/opt/docsgpt for root on Linux); DOCSGPT_HOME and a checkout still take precedence. The API and worker commands create the home and point out a .env left in the working directory. --- docsgpt/cli.py | 20 +++++++++++++--- docsgpt/core/paths.py | 22 ++++++++++++----- tests/core/test_paths.py | 24 +++++++++++++++++-- tests/test_cli.py | 51 ++++++++++++++++++++++++++++++++++++++++ 4 files changed, 106 insertions(+), 11 deletions(-) diff --git a/docsgpt/cli.py b/docsgpt/cli.py index e9d4d670..36dc5de8 100644 --- a/docsgpt/cli.py +++ b/docsgpt/cli.py @@ -19,10 +19,24 @@ DEFAULT_PORT = 7091 def _announce_home() -> None: - """Say where runtime data and the env file come from; the API and the worker must agree.""" - from docsgpt.core.paths import env_file, home_dir + """Create the data home and say where data and the env file come from; the API and the worker must agree.""" + from pathlib import Path - print(f"docsgpt: data home {home_dir()} (env file {env_file()})", file=sys.stderr) + from docsgpt.core import paths + + home = paths.home_dir() + home.mkdir(parents=True, exist_ok=True) + env = paths.env_file() + print(f"docsgpt: data home {home} (env file {env})", file=sys.stderr) + # Up to 0.20 an installed package used the working directory as its home. + chosen = os.environ.get(paths.HOME_ENV) or os.environ.get(paths.ENV_FILE_ENV) or paths.checkout_root() + stray = Path.cwd() / ".env" + if not chosen and stray.is_file() and stray.resolve() != env.resolve(): + print( + f"docsgpt: {stray} is not used; settings come from {env}. " + f"Move the file there, or set DOCSGPT_HOME={Path.cwd()} to keep using this directory.", + file=sys.stderr, + ) def _gunicorn_options(host: str, port: int, workers: int) -> dict: diff --git a/docsgpt/core/paths.py b/docsgpt/core/paths.py index ff7b46bd..b8e2ec77 100644 --- a/docsgpt/core/paths.py +++ b/docsgpt/core/paths.py @@ -5,15 +5,18 @@ ships with it (prompts, model catalogs, migrations). A checkout, when the package is imported from one, is the directory holding ``pyproject.toml``. The data home is where runtime data lives: ``.env``, ``inputs``, ``indexes``. -The home is ``DOCSGPT_HOME`` when set, else the checkout, else the current -directory. That keeps a source checkout and the Docker image (which runs from -``/app`` with the package beside it) behaving as before, and gives a -``pip install docsgpt`` user a home that is not ``site-packages``. +The home is ``DOCSGPT_HOME`` when set, else the checkout, else the default +home (``~/.docsgpt/server``, or ``/opt/docsgpt`` for root on Linux). That keeps +a source checkout and the Docker image (which pins ``DOCSGPT_HOME=/app``) +behaving as before, and gives an installed package one home no matter which +directory the command runs from. ``docsgpt up`` keeps its stack in the same +place. """ from __future__ import annotations import os +import sys from pathlib import Path HOME_ENV = "DOCSGPT_HOME" @@ -31,12 +34,19 @@ def checkout_root() -> Path | None: return root if (root / "pyproject.toml").is_file() else None +def default_home() -> Path: + """The home outside a checkout: ``/opt/docsgpt`` for root on Linux, else ``~/.docsgpt/server``.""" + if sys.platform.startswith("linux") and hasattr(os, "geteuid") and os.geteuid() == 0: + return Path("/opt/docsgpt") + return Path.home() / ".docsgpt" / "server" + + def home_dir() -> Path: - """Directory for runtime data: ``DOCSGPT_HOME``, else the checkout, else cwd.""" + """Directory for runtime data: ``DOCSGPT_HOME``, else the checkout, else the default home.""" configured = os.environ.get(HOME_ENV) if configured: return Path(configured).expanduser().resolve() - return checkout_root() or Path.cwd() + return checkout_root() or default_home() def env_file() -> Path: diff --git a/tests/core/test_paths.py b/tests/core/test_paths.py index a1c80acd..73abec86 100644 --- a/tests/core/test_paths.py +++ b/tests/core/test_paths.py @@ -20,11 +20,31 @@ class TestHomeDir: monkeypatch.setenv(paths.HOME_ENV, str(tmp_path)) assert paths.home_dir() == tmp_path.resolve() - def test_an_installed_package_falls_back_to_cwd(self, monkeypatch, tmp_path): + def test_an_installed_package_uses_the_default_home_not_cwd(self, monkeypatch, tmp_path): + """`docsgpt api` from any directory finds the same .env and data.""" monkeypatch.delenv(paths.HOME_ENV, raising=False) monkeypatch.setattr(paths, "checkout_root", lambda: None) + monkeypatch.setattr(paths, "default_home", lambda: tmp_path / "home") monkeypatch.chdir(tmp_path) - assert paths.home_dir() == Path.cwd() + assert paths.home_dir() == tmp_path / "home" + + +class TestDefaultHome: + def test_a_user_gets_a_folder_in_their_home(self, monkeypatch, tmp_path): + monkeypatch.setattr(paths.sys, "platform", "darwin") + monkeypatch.setattr(paths.Path, "home", classmethod(lambda cls: tmp_path)) + assert paths.default_home() == tmp_path / ".docsgpt" / "server" + + def test_root_on_linux_gets_opt(self, monkeypatch): + monkeypatch.setattr(paths.sys, "platform", "linux") + monkeypatch.setattr(paths.os, "geteuid", lambda: 0, raising=False) + assert paths.default_home() == Path("/opt/docsgpt") + + def test_a_normal_user_on_linux_stays_in_their_home(self, monkeypatch, tmp_path): + monkeypatch.setattr(paths.sys, "platform", "linux") + monkeypatch.setattr(paths.os, "geteuid", lambda: 1000, raising=False) + monkeypatch.setattr(paths.Path, "home", classmethod(lambda cls: tmp_path)) + assert paths.default_home() == tmp_path / ".docsgpt" / "server" class TestEnvFile: diff --git a/tests/test_cli.py b/tests/test_cli.py index c8f934ce..452f5598 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -34,6 +34,57 @@ class TestTopLevel: subprocess.run([sys.executable, "-c", code], cwd=Path(__file__).resolve().parents[1], check=True) +class TestHome: + @staticmethod + def _installed(monkeypatch, tmp_path): + """An installed package (no checkout) with the default home under tmp_path.""" + from docsgpt.core import paths + + monkeypatch.delenv(paths.HOME_ENV, raising=False) + monkeypatch.delenv(paths.ENV_FILE_ENV, raising=False) + monkeypatch.setattr(paths, "checkout_root", lambda: None) + monkeypatch.setattr(paths, "default_home", lambda: tmp_path / "home") + return tmp_path / "home" + + def test_the_home_is_created_and_announced(self, monkeypatch, tmp_path, capsys): + home = self._installed(monkeypatch, tmp_path) + monkeypatch.chdir(tmp_path) + cli._announce_home() + assert home.is_dir() + assert f"data home {home}" in capsys.readouterr().err + + def test_an_env_file_left_in_the_working_directory_is_pointed_out(self, monkeypatch, tmp_path, capsys): + """Up to 0.20 an installed package read .env from the working directory.""" + home = self._installed(monkeypatch, tmp_path) + work = tmp_path / "work" + work.mkdir() + (work / ".env").write_text("LLM_PROVIDER=openai\n") + monkeypatch.chdir(work) + cli._announce_home() + err = capsys.readouterr().err + assert f"{work / '.env'} is not used" in err + assert f"DOCSGPT_HOME={work}" in err + assert str(home) in err + + def test_no_warning_when_the_home_is_chosen_explicitly(self, monkeypatch, tmp_path, capsys): + from docsgpt.core import paths + + self._installed(monkeypatch, tmp_path) + (tmp_path / ".env").write_text("LLM_PROVIDER=openai\n") + monkeypatch.setenv(paths.HOME_ENV, str(tmp_path / "elsewhere")) + monkeypatch.chdir(tmp_path) + cli._announce_home() + assert "is not used" not in capsys.readouterr().err + + def test_no_warning_when_the_working_directory_is_the_home(self, monkeypatch, tmp_path, capsys): + home = self._installed(monkeypatch, tmp_path) + home.mkdir() + (home / ".env").write_text("LLM_PROVIDER=openai\n") + monkeypatch.chdir(home) + cli._announce_home() + assert "is not used" not in capsys.readouterr().err + + class TestApi: def test_gunicorn_runs_with_the_image_settings_and_leaves_argv_alone(self, monkeypatch, capsys): application = MagicMock() From aa0e4ea2808f7d88eff3ce3eb65b52df714a8661 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 22:57:00 +0100 Subject: [PATCH 006/130] feat: docsgpt up runs and manages DocsGPT on Docker `docsgpt up` copies the standalone Compose file shipped with this package version into the stack directory (~/.docsgpt/server by default), writes its .env and starts the stack on the images of the same version. A first run asks who should reach DocsGPT (this computer, the network with a token, or a domain with HTTPS) and which model provider to use; flags answer the same questions for scripts. Re-running keeps secrets and settings and moves the image tag, and the database password is only generated for a new database. Also: down, status, logs, token, open, env, upgrade (uv tool installs upgrade themselves and run `up` again) and uninstall (keeps settings and data unless --purge). The commands import no Flask, Celery or settings. The wheel carries deployment/docker-compose-standalone.yaml as docsgpt/deploy/docker-compose.yaml; the sdist includes the source file. --- .github/workflows/package-build.yml | 2 + docsgpt/cli.py | 86 ++++++- docsgpt/deploy/__init__.py | 5 + docsgpt/deploy/commands.py | 381 ++++++++++++++++++++++++++++ docsgpt/deploy/docker.py | 143 +++++++++++ docsgpt/deploy/envfile.py | 74 ++++++ docsgpt/deploy/stack.py | 184 ++++++++++++++ pyproject.toml | 8 +- tests/deploy/__init__.py | 0 tests/deploy/test_commands.py | 286 +++++++++++++++++++++ tests/deploy/test_docker.py | 133 ++++++++++ tests/deploy/test_envfile.py | 72 ++++++ tests/deploy/test_stack.py | 196 ++++++++++++++ 13 files changed, 1567 insertions(+), 3 deletions(-) create mode 100644 docsgpt/deploy/__init__.py create mode 100644 docsgpt/deploy/commands.py create mode 100644 docsgpt/deploy/docker.py create mode 100644 docsgpt/deploy/envfile.py create mode 100644 docsgpt/deploy/stack.py create mode 100644 tests/deploy/__init__.py create mode 100644 tests/deploy/test_commands.py create mode 100644 tests/deploy/test_docker.py create mode 100644 tests/deploy/test_envfile.py create mode 100644 tests/deploy/test_stack.py diff --git a/.github/workflows/package-build.yml b/.github/workflows/package-build.yml index 0be880db..23798bfe 100644 --- a/.github/workflows/package-build.yml +++ b/.github/workflows/package-build.yml @@ -65,6 +65,8 @@ jobs: names = set(zipfile.ZipFile(glob.glob("dist/*.whl")[0]).namelist()) for required in ( "docsgpt/cli.py", + "docsgpt/deploy/commands.py", + "docsgpt/deploy/docker-compose.yaml", "docsgpt/alembic.ini", "docsgpt/alembic/env.py", "docsgpt/alembic/script.py.mako", diff --git a/docsgpt/cli.py b/docsgpt/cli.py index 36dc5de8..873cd099 100644 --- a/docsgpt/cli.py +++ b/docsgpt/cli.py @@ -1,4 +1,5 @@ -"""The ``docsgpt`` command: run the API, the worker and the maintenance scripts. +"""The ``docsgpt`` command: run the API, the worker and the maintenance scripts, +or run and manage DocsGPT on Docker (``docsgpt up``). Every subcommand imports what it needs when it runs, so ``docsgpt --help`` stays instant and does not touch the database. @@ -153,6 +154,17 @@ def _migrate(args: argparse.Namespace) -> int: return 0 +def _deploy(name: str): + """A subcommand handler that imports ``docsgpt.deploy.commands`` only when it runs.""" + + def handler(args: argparse.Namespace, context=None) -> int: + from docsgpt.deploy import commands + + return getattr(commands, name)(args, context) + + return handler + + # Maintenance scripts keep their own argument parsers; the command hands # everything after the script name to them untouched (argparse would try to # interpret the options itself). @@ -169,11 +181,70 @@ def _run_script(module: str, argv: list[str]) -> int: return int(importlib.import_module(f"docsgpt.scripts.{module}").main(argv) or 0) +def _add_deploy_commands(commands) -> None: + """``docsgpt up`` and the commands that manage the Docker stack it runs.""" + from docsgpt.deploy.stack import EXPOSURES, PROVIDERS + + def stack_command(name: str, handler: str, help_text: str) -> argparse.ArgumentParser: + parser = commands.add_parser(name, help=help_text) + parser.add_argument("--dir", help="stack directory (default: DOCSGPT_HOME, else ~/.docsgpt/server)") + parser.set_defaults(func=_deploy(handler), deploy=True) + return parser + + up = stack_command("up", "up", "install or update DocsGPT on Docker and start it") + up.add_argument("--expose", choices=EXPOSURES, help="who can reach it: local (default), network or domain") + up.add_argument("--domain", help="public domain served over HTTPS by Caddy (implies --expose domain)") + up.add_argument("--port", type=int, help="host port for the UI and API (default: 7091)") + up.add_argument("--provider", choices=list(PROVIDERS), help="model provider (default: the DocsGPT public API)") + up.add_argument("--api-key", help="the provider's API key (or set DOCSGPT_API_KEY)") + up.add_argument("--model", help="model name (required for openai-compatible)") + up.add_argument("--base-url", help="base URL of an OpenAI-compatible server") + docling = up.add_mutually_exclusive_group() + docling.add_argument("--docling", dest="docling", action="store_const", const=True, + help="run the image with the docling parser engine and OCR (several GB larger)") + docling.add_argument("--no-docling", dest="docling", action="store_const", const=False, + help="go back to the default image") + up.set_defaults(docling=None) + up.add_argument("--image-tag", help="image tag to run instead of this package's version, e.g. develop") + up.add_argument("-y", "--yes", action="store_true", help="ask nothing: use the flags, then the defaults") + up.add_argument("--reconfigure", action="store_true", help="ask the setup questions again") + up.add_argument("--adopt", action="store_true", help="take over a DocsGPT stack started from another folder") + up.add_argument("--no-open", action="store_true", help="do not open the browser after the first install") + up.add_argument("--timeout", type=int, default=300, help="seconds to wait for the API to answer (default: 300)") + + stack_command("down", "down", "stop the Docker stack (data and settings stay)") + stack_command("status", "status", "show the stack's version, address, containers and health") + + logs = stack_command("logs", "logs", "show the stack's logs") + logs.add_argument("-f", "--follow", action="store_true", help="keep printing new lines") + logs.add_argument("--tail", type=int, help="only the last N lines of each service") + logs.add_argument("services", nargs="*", help="services to show, e.g. backend worker") + + stack_command("token", "token", "print the access token (installs reachable beyond this computer)") + stack_command("open", "open_ui", "open DocsGPT in the browser") + + upgrade = stack_command("upgrade", "upgrade", "upgrade the package and restart the stack on the new version") + upgrade.add_argument("--version", help="version to install (default: the latest release)") + + uninstall = stack_command("uninstall", "uninstall", "remove the Docker stack") + uninstall.add_argument("-y", "--yes", action="store_true", help="do not ask for confirmation") + uninstall.add_argument("--purge", action="store_true", help="also delete the settings and all data") + + env = stack_command("env", "env", "show, get or set the stack's settings") + env_actions = env.add_subparsers(dest="env_action", metavar="") + get = env_actions.add_parser("get", help="print one setting") + get.add_argument("key") + set_ = env_actions.add_parser("set", help="set settings (KEY=VALUE ...); run `docsgpt up` to apply") + set_.add_argument("pairs", nargs="+", metavar="KEY=VALUE") + + def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="docsgpt", description="DocsGPT: private AI for agents, assistants and search.") parser.add_argument("--version", action="version", version=f"docsgpt {__version__}") commands = parser.add_subparsers(dest="command", metavar="") + _add_deploy_commands(commands) + api = commands.add_parser("api", help="serve the HTTP API") api.add_argument("--host", default=DEFAULT_HOST, help="interface to listen on (default: localhost; 0.0.0.0 for all)") api.add_argument("--port", type=int, default=DEFAULT_PORT) @@ -212,7 +283,18 @@ def main(argv: Optional[Sequence[str]] = None) -> int: if not args.command: parser.print_help() return 2 - return args.func(args) + if not getattr(args, "deploy", False): + return args.func(args) + from docsgpt.deploy.docker import DeployError + + try: + return args.func(args) + except DeployError as exc: + print(f"docsgpt: {exc}", file=sys.stderr) + return 1 + except KeyboardInterrupt: + print(file=sys.stderr) + return 130 if __name__ == "__main__": diff --git a/docsgpt/deploy/__init__.py b/docsgpt/deploy/__init__.py new file mode 100644 index 00000000..18c30821 --- /dev/null +++ b/docsgpt/deploy/__init__.py @@ -0,0 +1,5 @@ +"""Run DocsGPT on Docker from the installed package: ``docsgpt up`` and the commands that manage it. + +Nothing here imports the Flask app, Celery or the settings module, so these +commands start instantly and work before any configuration exists. +""" diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py new file mode 100644 index 00000000..140b7929 --- /dev/null +++ b/docsgpt/deploy/commands.py @@ -0,0 +1,381 @@ +"""``docsgpt up`` and the commands that manage the stack it starts.""" + +from __future__ import annotations + +import getpass +import json +import os +import re +import shutil +import subprocess +import sys +import webbrowser +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Optional + +from docsgpt.deploy import envfile, stack +from docsgpt.deploy.docker import DeployError, Docker, lan_ip, wait_healthy + +PROJECT = "docsgpt" +DATABASE_VOLUME = f"{PROJECT}_postgres_data" +_MOVING_TAGS = ("latest", "develop") +_KEY = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") + +EXPOSURE_CHOICES = [ + ("local", "Only this computer"), + ("network", "Other machines on the network (plain HTTP, access token)"), + ("domain", "A domain name with HTTPS (Caddy certificate, access token)"), +] + + +class Prompter: + """Questions on the terminal. The installer hands its terminal to ``docsgpt up``.""" + + def choose(self, question: str, options: list[tuple[str, str]], default: str) -> str: + keys = [key for key, _ in options] + print(question) + for number, (key, label) in enumerate(options, 1): + print(f" {number}) {label}{' (default)' if key == default else ''}") + while True: + answer = input(f"Choose 1-{len(options)} [{keys.index(default) + 1}]: ").strip() + if not answer: + return default + if answer.isdigit() and 1 <= int(answer) <= len(options): + return keys[int(answer) - 1] + if answer in keys: + return answer + print("Enter one of the numbers above.") + + def text(self, question: str, default: Optional[str] = None, secret: bool = False) -> str: + ask = getpass.getpass if secret else input + while True: + answer = ask(f"{question}{f' [{default}]' if default else ''}: ").strip() + if answer: + return answer + if default: + return default + + def confirm(self, question: str, default: bool = False) -> bool: + answer = input(f"{question} [{'Y/n' if default else 'y/N'}]: ").strip().lower() + return default if not answer else answer in ("y", "yes") + + +def detect_installer() -> str: + """How this ``docsgpt`` was installed: ``uv`` (uv tool), ``pipx`` or ``pip``.""" + prefix = Path(sys.prefix) + if (prefix / "uv-receipt.toml").is_file(): + return "uv" + if "pipx" in prefix.parts: + return "pipx" + return "pip" + + +def _run_command(args: list[str]) -> int: + try: + return subprocess.call(args) + except FileNotFoundError as exc: + raise DeployError(f"{args[0]} is not on PATH") from exc + + +def _exec_command(argv: list[str]) -> int: + """Replace this process with ``argv`` (the upgraded ``docsgpt``); Windows runs it and waits.""" + executable = shutil.which(argv[0]) or argv[0] + if sys.platform == "win32": + return subprocess.call([executable, *argv[1:]]) + os.execv(executable, [executable, *argv[1:]]) + return 0 # not reached + + +@dataclass +class Context: + """What the commands talk to; tests replace the parts that touch Docker, the terminal or the network.""" + + docker: Any + prompter: Any + interactive: bool + version: str + lan_ip: Callable[[], str] = lan_ip + wait: Callable[[str, float], bool] = wait_healthy + open_browser: Callable[[str], Any] = webbrowser.open + installer: Callable[[], str] = detect_installer + run: Callable[[list[str]], int] = _run_command + exec_up: Callable[[list[str]], int] = _exec_command + + @classmethod + def default(cls, args) -> "Context": + from docsgpt.version import __version__ + + interactive = sys.stdin.isatty() and not getattr(args, "yes", False) + return cls(docker=Docker(), prompter=Prompter(), interactive=interactive, version=__version__) + + +def _installed(directory: Path) -> Optional[dict[str, str]]: + """The stack's settings, or None (with a message) when there is no install in ``directory``.""" + if not (directory / stack.COMPOSE_FILE).is_file(): + print(f"No DocsGPT install in {directory}. Run `docsgpt up` first, or pass --dir.", file=sys.stderr) + return None + return envfile.read(directory / ".env") + + +def _current_provider(env: Mapping[str, str]) -> str: + name = env.get("LLM_PROVIDER", "docsgpt") + if name == "openai" and env.get("OPENAI_BASE_URL"): + return "openai-compatible" + return name if name in stack.PROVIDERS else "docsgpt" + + +def _check_other_stacks(args, context: Context, directory: Path) -> None: + others = {path for path in context.docker.project_dirs(PROJECT) if path.resolve() != directory.resolve()} + if not others or args.adopt: + return + where = ", ".join(sorted(str(path) for path in others)) + question = ( + f"Docker already runs a DocsGPT stack started from {where}; it uses the same data volumes. " + f"Manage it from {directory} instead?" + ) + if context.interactive and context.prompter.confirm(question, default=False): + return + raise DeployError( + f"Docker already runs a DocsGPT stack started from {where}. Stop it with `docker compose down` " + f"in that folder, or run again with --adopt to manage it from {directory}." + ) + + +def _choose_provider(args, context: Context, existing: Mapping[str, str], ask: bool): + name = args.provider + if name is None and ask: + name = context.prompter.choose("Which model provider?", list(stack.PROVIDERS.items()), _current_provider(existing)) + if name is None: + return None + api_key = args.api_key or os.environ.get("DOCSGPT_API_KEY") + model, base_url = args.model, args.base_url + if context.interactive: + if name == "openai-compatible": + base_url = base_url or context.prompter.text( + "Server base URL (Ollama on this machine: http://host.docker.internal:11434/v1)" + ) + model = model or context.prompter.text("Model name") + elif name != "docsgpt" and not api_key: + api_key = context.prompter.text(f"{stack.PROVIDERS[name]} API key", secret=True) + try: + return stack.provider_settings(name, api_key=api_key, model=model, base_url=base_url) + except ValueError as exc: + raise DeployError(str(exc)) from exc + + +def up(args, context: Optional[Context] = None) -> int: + """Install or update the stack in its directory and start it.""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + env_path = directory / ".env" + record_path = directory / stack.RECORD_FILE + existing = envfile.read(env_path) + configured = record_path.is_file() + + context.docker.preflight(context.interactive) + _check_other_stacks(args, context, directory) + + ask = context.interactive and (not configured or args.reconfigure) + expose, domain = args.expose, args.domain + if ask and expose is None and domain is None: + expose = context.prompter.choose("Who should reach DocsGPT?", EXPOSURE_CHOICES, stack.exposure(existing)) + if expose == "domain" and not domain: + if not context.interactive: + raise DeployError("--expose domain needs --domain") + domain = context.prompter.text("Domain name (its DNS must point at this machine)", existing.get("DOCSGPT_DOMAIN")) + provider = _choose_provider(args, context, existing, ask) + + image_tag = args.image_tag or context.version + try: + updates = stack.plan( + existing, + image_tag=image_tag, + fresh_database=not context.docker.volume_exists(DATABASE_VOLUME), + expose=expose, + domain=domain, + port=args.port, + provider=provider, + docling=args.docling, + ) + except ValueError as exc: + raise DeployError(str(exc)) from exc + + directory.mkdir(parents=True, exist_ok=True) + shutil.copyfile(stack.compose_source(), directory / stack.COMPOSE_FILE) + envfile.update(env_path, updates) + env = envfile.read(env_path) + + up_args = ["up", "-d", "--remove-orphans"] + if image_tag in _MOVING_TAGS: + up_args += ["--pull", "always"] + print(f"Starting DocsGPT {image_tag} from {directory} ...") + context.docker.compose(directory, *up_args) + + health = stack.health_url(env) + print("Waiting for DocsGPT to answer (the first start also sets up the database) ...") + if not context.wait(health, args.timeout): + print( + f"DocsGPT did not answer at {health} within {args.timeout} seconds. " + "See what happened with `docsgpt logs backend`.", + file=sys.stderr, + ) + return 1 + + now = datetime.now(timezone.utc).isoformat(timespec="seconds") + record = json.loads(record_path.read_text(encoding="utf-8")) if configured else {"installed_at": now} + record.update(version=context.version, image_tag=image_tag, updated_at=now) + record_path.write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") + + address = stack.url(env, context.lan_ip()) + mode = stack.exposure(env) + print(f"\nDocsGPT is running at {address}") + if env.get("AUTH_TYPE") == "simple_jwt" and env.get("JWT_SECRET_KEY"): + print("Access token (the page asks for it; `docsgpt token` prints it again):") + print(f" {stack.simple_jwt_token(env['JWT_SECRET_KEY'])}") + if mode == "network": + print("Traffic is plain HTTP. Outside a trusted network, use a domain with HTTPS: docsgpt up --domain ") + elif mode == "domain": + print("Caddy gets the certificate when it starts: DNS must point at this machine and ports 80 and 443 be open.") + print(f"Settings: {env_path} (change the model provider or access with `docsgpt up --reconfigure`)") + print("Manage it with: docsgpt status | logs | upgrade | down | uninstall") + if context.interactive and not configured and not args.no_open and mode == "local": + context.open_browser(address) + return 0 + + +def down(args, context: Optional[Context] = None) -> int: + """Stop the stack; data and settings stay.""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + if _installed(directory) is None: + return 1 + context.docker.compose(directory, "down") + return 0 + + +def status(args, context: Optional[Context] = None) -> int: + """Version, address, containers and whether the API answers (exit 1 when it does not).""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + env = _installed(directory) + if env is None: + return 1 + tag = env.get("DOCSGPT_IMAGE_TAG", "unknown") + env.get("DOCSGPT_IMAGE_VARIANT", "") + print(f"DocsGPT {tag} in {directory}") + print(f"Address: {stack.url(env, context.lan_ip())} ({stack.exposure(env)})") + context.docker.compose(directory, "ps", check=False) + healthy = context.wait(stack.health_url(env), 0) + print("API: answering" if healthy else "API: not answering (see `docsgpt logs backend`)") + return 0 if healthy else 1 + + +def logs(args, context: Optional[Context] = None) -> int: + """``docker compose logs`` for the stack.""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + if _installed(directory) is None: + return 1 + options = (["--follow"] if args.follow else []) + (["--tail", str(args.tail)] if args.tail else []) + return context.docker.compose(directory, "logs", *options, *args.services, check=False).returncode + + +def token(args, context: Optional[Context] = None) -> int: + """Print the access token of a ``simple_jwt`` install.""" + directory = stack.stack_dir(args.dir) + env = _installed(directory) + if env is None: + return 1 + if env.get("AUTH_TYPE") != "simple_jwt" or not env.get("JWT_SECRET_KEY"): + print(f"This install has no access token (AUTH_TYPE={env.get('AUTH_TYPE') or 'none'}).", file=sys.stderr) + return 1 + print(stack.simple_jwt_token(env["JWT_SECRET_KEY"])) + return 0 + + +def open_ui(args, context: Optional[Context] = None) -> int: + """Open DocsGPT in the browser.""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + env = _installed(directory) + if env is None: + return 1 + address = stack.url(env, context.lan_ip()) + print(address) + context.open_browser(address) + return 0 + + +def env(args, context: Optional[Context] = None) -> int: + """Show where the settings are, or get and set them.""" + env_path = stack.stack_dir(args.dir) / ".env" + if args.env_action is None: + print(env_path) + return 0 + if args.env_action == "get": + values = envfile.read(env_path) + if args.key not in values: + print(f"{args.key} is not set in {env_path}", file=sys.stderr) + return 1 + print(values[args.key]) + return 0 + updates = {} + for pair in args.pairs: + key, separator, value = pair.partition("=") + if not separator or not _KEY.match(key): + raise DeployError(f"expected KEY=VALUE, got {pair!r}") + updates[key] = value + try: + envfile.update(env_path, updates) + except ValueError as exc: + raise DeployError(str(exc)) from exc + print(f"Saved to {env_path}. Run `docsgpt up` to apply.") + return 0 + + +def _uninstall_hint(installer: str) -> str: + return {"uv": "uv tool uninstall docsgpt", "pipx": "pipx uninstall docsgpt"}.get(installer, "pip uninstall docsgpt") + + +def upgrade(args, context: Optional[Context] = None) -> int: + """Upgrade the package, then run the new version's ``docsgpt up``.""" + context = context or Context.default(args) + spec = f"docsgpt=={args.version}" if args.version else "docsgpt" + installer = context.installer() + if installer == "uv": + if context.run(["uv", "tool", "install", "--force", spec]) != 0: + raise DeployError(f"uv could not install {spec}") + return context.exec_up(["docsgpt", "up", "--dir", str(stack.stack_dir(args.dir))]) + command = f"pipx install --force {spec}" if installer == "pipx" else f"pip install -U {spec}" + print(f"Upgrade the package with `{command}`, then run `docsgpt up` to move the stack to it.", file=sys.stderr) + return 1 + + +def uninstall(args, context: Optional[Context] = None) -> int: + """Remove the containers and the stack files; ``--purge`` also deletes settings and data.""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + if _installed(directory) is None: + return 1 + if args.purge: + what = "containers, settings and data (documents, conversations, the database)" + else: + what = "containers (the settings in .env and the data volumes are kept)" + if not args.yes: + if not context.interactive: + raise DeployError("uninstall needs --yes when there is no terminal to confirm on") + if not context.prompter.confirm(f"Remove the DocsGPT {what} in {directory}?", default=False): + print("Nothing removed.") + return 1 + context.docker.compose(directory, "down", "--remove-orphans", *(["-v"] if args.purge else [])) + if args.purge: + shutil.rmtree(directory) + print(f"Removed DocsGPT and its data from {directory}.") + else: + for name in (stack.COMPOSE_FILE, stack.RECORD_FILE): + (directory / name).unlink(missing_ok=True) + print(f"Removed the containers. Settings stay in {directory / '.env'} and data in the Docker volumes.") + print(f"To remove the docsgpt command too: {_uninstall_hint(context.installer())}") + return 0 diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py new file mode 100644 index 00000000..0b3a1ea4 --- /dev/null +++ b/docsgpt/deploy/docker.py @@ -0,0 +1,143 @@ +"""The docker CLI calls behind ``docsgpt up``.""" + +from __future__ import annotations + +import http.client +import re +import shutil +import socket +import subprocess +import sys +import time +import urllib.request +from collections.abc import Callable, Sequence +from pathlib import Path +from typing import Optional + +MIN_COMPOSE = (2, 24, 0) +DAEMON_START_SECONDS = 120 + + +class DeployError(Exception): + """A problem the user can act on; the command prints it without a traceback.""" + + +def run(args: Sequence[str], *, cwd: Optional[Path] = None, capture: bool = False, check: bool = True): + """Run a command, streaming its output unless ``capture``; with ``check`` a failure raises DeployError.""" + try: + result = subprocess.run(list(args), cwd=cwd, text=True, capture_output=capture, check=False) + except FileNotFoundError as exc: + raise DeployError(f"{args[0]} is not installed or not on PATH") from exc + if check and result.returncode != 0: + detail = (result.stderr or "").strip() if capture else "" + message = f"`{' '.join(args)}` failed with exit code {result.returncode}" + raise DeployError(f"{message}: {detail}" if detail else message) + return result + + +def _parse_version(text: str) -> Optional[tuple[int, ...]]: + match = re.search(r"(\d+)\.(\d+)\.(\d+)", text or "") + return tuple(int(part) for part in match.groups()) if match else None + + +class Docker: + """Docker and Docker Compose, through their command-line tools.""" + + def __init__( + self, + runner: Callable[..., subprocess.CompletedProcess] = run, + which: Callable[[str], Optional[str]] = shutil.which, + sleep: Callable[[float], None] = time.sleep, + platform: str = sys.platform, + ) -> None: + self._run = runner + self._which = which + self._sleep = sleep + self._platform = platform + + def preflight(self, interactive: bool = False) -> None: + """Make sure Docker is installed and running and Compose is new enough, starting Docker Desktop on macOS.""" + if not self._which("docker"): + raise DeployError("Docker is not installed. Get it from https://docs.docker.com/get-docker/ and run this again.") + if not self.daemon_running(): + self._start_daemon() + result = self._run(["docker", "compose", "version", "--short"], capture=True, check=False) + version = _parse_version(result.stdout) if result.returncode == 0 else None + if version is None: + raise DeployError( + "Docker Compose v2 is not available (`docker compose version` failed). " + "Install the Compose plugin: https://docs.docker.com/compose/install/" + ) + if version < MIN_COMPOSE: + found = ".".join(str(part) for part in version) + raise DeployError(f"Docker Compose {found} is too old; DocsGPT needs 2.24 or newer.") + + def daemon_running(self) -> bool: + return self._run(["docker", "info"], capture=True, check=False).returncode == 0 + + def _start_daemon(self) -> None: + if self._platform == "darwin": + print("Docker is not running; starting Docker Desktop ...", file=sys.stderr) + self._run(["open", "-a", "Docker"], check=False) + for _ in range(DAEMON_START_SECONDS // 2): + self._sleep(2) + if self.daemon_running(): + return + raise DeployError("Docker Desktop did not start within two minutes. Start it and run this again.") + if self._platform.startswith("linux"): + raise DeployError("Docker is not running. Start it with `sudo systemctl start docker` and run this again.") + raise DeployError("Docker is not running. Start Docker Desktop and run this again.") + + def compose(self, directory: Path, *args: str, capture: bool = False, check: bool = True): + """``docker compose `` in ``directory``, which holds the Compose file and its ``.env``.""" + return self._run(["docker", "compose", *args], cwd=directory, capture=capture, check=check) + + def volume_exists(self, name: str) -> bool: + return self._run(["docker", "volume", "inspect", name], capture=True, check=False).returncode == 0 + + def project_dirs(self, project: str) -> set[Path]: + """The folders containers of Compose project ``project`` were started from.""" + result = self._run( + [ + "docker", "ps", "-a", + "--filter", f"label=com.docker.compose.project={project}", + "--format", '{{.Label "com.docker.compose.project.working_dir"}}', + ], + capture=True, + check=False, + ) + if result.returncode != 0: + return set() + return {Path(line.strip()) for line in result.stdout.splitlines() if line.strip()} + + +def wait_healthy( + url: str, + timeout: float, + *, + opener: Callable = urllib.request.urlopen, + sleep: Callable[[float], None] = time.sleep, + clock: Callable[[], float] = time.monotonic, +) -> bool: + """Poll ``url`` until it answers 2xx (True) or ``timeout`` seconds pass (False); tries at least once.""" + deadline = clock() + timeout + while True: + try: + with opener(url, timeout=5) as response: + if 200 <= response.status < 300: + return True + except (OSError, http.client.HTTPException): + pass + if clock() >= deadline: + return False + sleep(2) + + +def lan_ip() -> str: + """This machine's address on its network, or ``localhost``. No packet is sent.""" + with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as probe: + try: + probe.connect(("192.0.2.1", 80)) + return probe.getsockname()[0] + except OSError: + return "localhost" diff --git a/docsgpt/deploy/envfile.py b/docsgpt/deploy/envfile.py new file mode 100644 index 00000000..9d51b946 --- /dev/null +++ b/docsgpt/deploy/envfile.py @@ -0,0 +1,74 @@ +"""Read and update a ``.env`` file in place, keeping the lines the user wrote.""" + +from __future__ import annotations + +import os +import re +from collections.abc import Mapping +from pathlib import Path +from typing import Optional + +_ASSIGNMENT = re.compile(r"^\s*(?:export\s+)?([A-Za-z_][A-Za-z0-9_]*)\s*=(.*)$") +# Values written without quotes: nothing Compose would interpolate, strip or treat as a comment. +_PLAIN = re.compile(r"^[A-Za-z0-9_./:@+,=-]*$") + + +def _parse_value(raw: str) -> str: + """The value of one assignment, with Compose's quoting rules.""" + value = raw.strip() + if len(value) >= 2 and value[0] == value[-1] == "'": + return value[1:-1] + if len(value) >= 2 and value[0] == value[-1] == '"': + return re.sub(r'\\(["\\])', r"\1", value[1:-1]) + return value.split(" #", 1)[0].rstrip() + + +def _format_value(value: str) -> str: + """``value`` quoted so it reads back unchanged.""" + if "\n" in value or "\r" in value: + raise ValueError("a .env value cannot contain a newline") + if _PLAIN.match(value): + return value + if "'" not in value: + return f"'{value}'" + return '"' + value.replace("\\", "\\\\").replace('"', '\\"') + '"' + + +def read(path: Path) -> dict[str, str]: + """The assignments in ``path`` (empty when it does not exist); a repeated key keeps its last value.""" + path = Path(path) + if not path.is_file(): + return {} + values: dict[str, str] = {} + for line in path.read_text(encoding="utf-8").splitlines(): + match = _ASSIGNMENT.match(line) + if match: + values[match.group(1)] = _parse_value(match.group(2)) + return values + + +def update(path: Path, values: Mapping[str, Optional[str]]) -> None: + """Set each key in place (``None`` removes it), append new keys, and leave every other line alone. + + A new file is created readable by its owner only: it holds secrets. + """ + formatted = {key: None if value is None else _format_value(value) for key, value in values.items()} + path = Path(path) + lines = path.read_text(encoding="utf-8").splitlines() if path.is_file() else [] + written: set[str] = set() + out: list[str] = [] + for line in lines: + match = _ASSIGNMENT.match(line) + key = match.group(1) if match else None + if key not in formatted: + out.append(line) + continue + if key not in written and formatted[key] is not None: + out.append(f"{key}={formatted[key]}") + written.add(key) + out.extend(f"{key}={value}" for key, value in formatted.items() if key not in written and value is not None) + + path.parent.mkdir(parents=True, exist_ok=True) + descriptor = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600) + with os.fdopen(descriptor, "w", encoding="utf-8") as handle: + handle.write("\n".join(out) + "\n" if out else "") diff --git a/docsgpt/deploy/stack.py b/docsgpt/deploy/stack.py new file mode 100644 index 00000000..bb577c56 --- /dev/null +++ b/docsgpt/deploy/stack.py @@ -0,0 +1,184 @@ +"""The stack ``docsgpt up`` runs: where it lives and what its ``.env`` holds.""" + +from __future__ import annotations + +import os +import secrets +from collections.abc import Callable, Mapping +from pathlib import Path +from typing import Optional + +from docsgpt.core import paths + +COMPOSE_FILE = "docker-compose.yaml" +RECORD_FILE = "install.json" +DEFAULT_PORT = 7091 +EXPOSURES = ("local", "network", "domain") + +PROVIDERS = { + "docsgpt": "DocsGPT public API (free, no key)", + "openai": "OpenAI", + "anthropic": "Anthropic", + "google": "Google Gemini", + "openrouter": "OpenRouter", + "groq": "Groq", + "openai-compatible": "OpenAI-compatible server (Ollama, vLLM, LM Studio, ...)", +} + +_LOCAL_BINDS = ("", "127.0.0.1", "localhost", "::1") +_ALL_INTERFACES = ("0.0.0.0", "::") + + +def compose_source() -> Path: + """The Compose file for this version: shipped in the wheel, or ``deployment/`` in a checkout.""" + packaged = paths.package_dir() / "deploy" / COMPOSE_FILE + if packaged.is_file(): + return packaged + root = paths.checkout_root() + if root is not None: + in_checkout = root / "deployment" / "docker-compose-standalone.yaml" + if in_checkout.is_file(): + return in_checkout + raise FileNotFoundError("the Compose file is missing from this docsgpt installation; reinstall the package") + + +def stack_dir(explicit: Optional[str]) -> Path: + """Where the stack lives: ``--dir``, else ``DOCSGPT_HOME``, else the default home (never a checkout).""" + if explicit: + return Path(explicit).expanduser().resolve() + configured = os.environ.get(paths.HOME_ENV) + if configured: + return Path(configured).expanduser().resolve() + return paths.default_home() + + +def _profiles(env: Mapping[str, str]) -> set[str]: + return {name.strip() for name in env.get("COMPOSE_PROFILES", "").split(",") if name.strip()} + + +def exposure(env: Mapping[str, str]) -> str: + """Who can reach the stack, read back from its ``.env``: local, network or domain.""" + if "https" in _profiles(env) and env.get("DOCSGPT_DOMAIN"): + return "domain" + if env.get("DOCSGPT_BIND", "") not in _LOCAL_BINDS: + return "network" + return "local" + + +def provider_settings( + name: str, + api_key: Optional[str] = None, + model: Optional[str] = None, + base_url: Optional[str] = None, +) -> dict[str, Optional[str]]: + """The model settings for a provider choice; keys another provider used are cleared.""" + if name not in PROVIDERS: + raise ValueError(f"unknown provider {name!r}; choose one of: {', '.join(PROVIDERS)}") + if name == "docsgpt": + return {"LLM_PROVIDER": "docsgpt", "API_KEY": None, "LLM_NAME": None, "OPENAI_BASE_URL": None} + if name == "openai-compatible": + if not base_url: + raise ValueError("an OpenAI-compatible server needs a base URL, e.g. http://host.docker.internal:11434/v1") + if not model: + raise ValueError("an OpenAI-compatible server needs a model name") + return { + "LLM_PROVIDER": "openai", + "API_KEY": api_key or "not-needed", + "LLM_NAME": model, + "OPENAI_BASE_URL": base_url, + } + if not api_key: + raise ValueError(f"{PROVIDERS[name]} needs an API key") + # Without LLM_NAME the model catalog picks the provider's first model. + return {"LLM_PROVIDER": name, "API_KEY": api_key, "LLM_NAME": model or None, "OPENAI_BASE_URL": None} + + +def plan( + existing: Mapping[str, str], + *, + image_tag: str, + fresh_database: bool, + expose: Optional[str] = None, + domain: Optional[str] = None, + port: Optional[int] = None, + provider: Optional[Mapping[str, Optional[str]]] = None, + docling: Optional[bool] = None, + secret: Optional[Callable[[], str]] = None, +) -> dict[str, Optional[str]]: + """The ``.env`` changes for an ``up``: only keys that change, ``None`` for a key to remove. + + Settings the user did not ask to change are left alone, secrets are generated + once, and the database password is only set for a database that does not exist + yet (Postgres reads it when the volume is created). + """ + secret = secret or (lambda: secrets.token_hex(32)) + wanted: dict[str, Optional[str]] = {"DOCSGPT_IMAGE_TAG": image_tag} + + if domain and expose is None: + expose = "domain" + if expose is None and "DOCSGPT_BIND" not in existing and "COMPOSE_PROFILES" not in existing: + expose = "local" + if expose == "local": + wanted.update(DOCSGPT_BIND="127.0.0.1", COMPOSE_PROFILES=None, DOCSGPT_DOMAIN=None) + elif expose == "network": + wanted.update(DOCSGPT_BIND="0.0.0.0", COMPOSE_PROFILES=None, DOCSGPT_DOMAIN=None) + elif expose == "domain": + if not domain: + raise ValueError("exposing DocsGPT on a domain needs the domain name") + wanted.update(DOCSGPT_BIND="127.0.0.1", COMPOSE_PROFILES="https", DOCSGPT_DOMAIN=domain) + elif expose is not None: + raise ValueError(f"unknown exposure {expose!r}; choose one of: {', '.join(EXPOSURES)}") + if expose in ("network", "domain") and not existing.get("AUTH_TYPE"): + wanted["AUTH_TYPE"] = "simple_jwt" + + for key in ("INTERNAL_KEY", "JWT_SECRET_KEY"): + if not existing.get(key): + wanted[key] = secret() + if not existing.get("POSTGRES_PASSWORD") and fresh_database: + wanted["POSTGRES_PASSWORD"] = secret() + if "VITE_API_STREAMING" not in existing: + wanted["VITE_API_STREAMING"] = "true" + if port is not None: + wanted["DOCSGPT_PORT"] = str(port) + if docling is not None: + wanted["DOCSGPT_IMAGE_VARIANT"] = "-docling" if docling else None + if provider is None and "LLM_PROVIDER" not in existing: + provider = provider_settings("docsgpt") + if provider: + wanted.update(provider) + + return { + key: value + for key, value in wanted.items() + if (value is None and key in existing) or (value is not None and existing.get(key) != value) + } + + +def _port(env: Mapping[str, str]) -> str: + return env.get("DOCSGPT_PORT") or str(DEFAULT_PORT) + + +def url(env: Mapping[str, str], lan_ip: str) -> str: + """The address to open DocsGPT at.""" + mode = exposure(env) + if mode == "domain": + return f"https://{env['DOCSGPT_DOMAIN']}" + if mode == "network": + bind = env.get("DOCSGPT_BIND", "") + host = lan_ip if bind in _ALL_INTERFACES else bind + return f"http://{host}:{_port(env)}" + return f"http://localhost:{_port(env)}" + + +def health_url(env: Mapping[str, str]) -> str: + """The API health check, reached from this machine whatever the exposure.""" + bind = env.get("DOCSGPT_BIND", "") + host = "127.0.0.1" if bind in _LOCAL_BINDS or bind in _ALL_INTERFACES else bind + return f"http://{host}:{_port(env)}/api/health" + + +def simple_jwt_token(secret_key: str) -> str: + """The token the API accepts under ``AUTH_TYPE=simple_jwt`` (it signs the same payload at start).""" + from jose import jwt + + return jwt.encode({"sub": "local"}, secret_key, algorithm="HS256") diff --git a/pyproject.toml b/pyproject.toml index 58493952..30f07df6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -194,8 +194,14 @@ exclude = [ "docsgpt/vectors/", ] +# `docsgpt up` runs the standalone Compose file of its own version. The file +# stays in deployment/ (the release asset and the curl instructions use it +# there); the wheel carries a copy inside the package. +[tool.hatch.build.targets.wheel.force-include] +"deployment/docker-compose-standalone.yaml" = "docsgpt/deploy/docker-compose.yaml" + [tool.hatch.build.targets.sdist] -include = ["/docsgpt"] +include = ["/docsgpt", "/deployment/docker-compose-standalone.yaml"] artifacts = ["docsgpt/static/**"] exclude = [ "docsgpt/Dockerfile", diff --git a/tests/deploy/__init__.py b/tests/deploy/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/deploy/test_commands.py b/tests/deploy/test_commands.py new file mode 100644 index 00000000..3aa80d5c --- /dev/null +++ b/tests/deploy/test_commands.py @@ -0,0 +1,286 @@ +"""`docsgpt up` and the commands that manage the stack, against a fake Docker.""" + +import json +import subprocess +import sys +from pathlib import Path + +import pytest + +from docsgpt import cli +from docsgpt.deploy import commands, envfile, stack +from docsgpt.deploy.docker import DeployError + + +class FakeDocker: + def __init__(self, volumes=(), project_dirs=()): + self.volumes = set(volumes) + self.dirs = {Path(d) for d in project_dirs} + self.calls = [] + self.preflights = 0 + + def preflight(self, interactive=False): + self.preflights += 1 + + def compose(self, directory, *args, capture=False, check=True): + self.calls.append((Path(directory), list(args))) + if args and args[0] == "down" and "-v" in args: + self.volumes.clear() + return subprocess.CompletedProcess(["docker", "compose", *args], 0, stdout="", stderr="") + + def volume_exists(self, name): + return name in self.volumes + + def project_dirs(self, project): + return set(self.dirs) + + +class FakePrompter: + def __init__(self, answers=()): + self.answers = list(answers) + self.questions = [] + + def _next(self, question): + self.questions.append(question) + if not self.answers: + raise AssertionError(f"unexpected question: {question}") + return self.answers.pop(0) + + def choose(self, question, options, default): + return self._next(question) + + def text(self, question, default=None, secret=False): + return self._next(question) + + def confirm(self, question, default=False): + return self._next(question) + + +def _context(docker=None, prompter=None, interactive=False, healthy=True, **overrides): + context = commands.Context( + docker=docker or FakeDocker(), + prompter=prompter or FakePrompter(), + interactive=interactive, + version="0.21.0", + lan_ip=lambda: "192.168.1.10", + wait=lambda url, timeout: healthy, + open_browser=lambda url: None, + ) + for key, value in overrides.items(): + setattr(context, key, value) + return context + + +def _run(argv, context): + args = cli.build_parser().parse_args(argv) + return args.func(args, context) + + +class TestUpFirstInstall: + def test_yes_installs_locally_with_the_public_api(self, tmp_path, capsys): + docker = FakeDocker() + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context(docker)) == 0 + + assert (tmp_path / "docker-compose.yaml").read_text() == stack.compose_source().read_text() + env = envfile.read(tmp_path / ".env") + assert env["DOCSGPT_IMAGE_TAG"] == "0.21.0" + assert env["DOCSGPT_BIND"] == "127.0.0.1" + assert env["LLM_PROVIDER"] == "docsgpt" + assert env["INTERNAL_KEY"] and env["JWT_SECRET_KEY"] and env["POSTGRES_PASSWORD"] + assert docker.preflights == 1 + assert (tmp_path, ["up", "-d", "--remove-orphans"]) in docker.calls + record = json.loads((tmp_path / "install.json").read_text()) + assert record["version"] == "0.21.0" + assert "http://localhost:7091" in capsys.readouterr().out + + def test_interactive_asks_who_reaches_it_and_which_model(self, tmp_path, capsys): + prompter = FakePrompter(["network", "anthropic", "sk-ant"]) + assert _run(["up", "--dir", str(tmp_path)], _context(prompter=prompter, interactive=True)) == 0 + env = envfile.read(tmp_path / ".env") + assert env["DOCSGPT_BIND"] == "0.0.0.0" + assert env["AUTH_TYPE"] == "simple_jwt" + assert env["LLM_PROVIDER"] == "anthropic" + assert env["API_KEY"] == "sk-ant" + out = capsys.readouterr().out + assert "http://192.168.1.10:7091" in out + assert stack.simple_jwt_token(env["JWT_SECRET_KEY"]) in out + + def test_flags_answer_the_questions(self, tmp_path): + argv = ["up", "--dir", str(tmp_path), "--domain", "docs.example.com", "--provider", "openai", "--api-key", "sk"] + assert _run(argv, _context(interactive=True)) == 0 + env = envfile.read(tmp_path / ".env") + assert env["COMPOSE_PROFILES"] == "https" + assert env["DOCSGPT_DOMAIN"] == "docs.example.com" + assert env["LLM_PROVIDER"] == "openai" + + def test_a_missing_api_key_is_an_error_without_a_terminal(self, tmp_path): + with pytest.raises(DeployError, match="API key"): + _run(["up", "--yes", "--dir", str(tmp_path), "--provider", "openai"], _context()) + + def test_the_api_key_can_come_from_the_environment(self, tmp_path, monkeypatch): + monkeypatch.setenv("DOCSGPT_API_KEY", "sk-env") + assert _run(["up", "--yes", "--dir", str(tmp_path), "--provider", "openai"], _context()) == 0 + assert envfile.read(tmp_path / ".env")["API_KEY"] == "sk-env" + + def test_an_existing_database_keeps_its_password(self, tmp_path): + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context(docker)) == 0 + assert "POSTGRES_PASSWORD" not in envfile.read(tmp_path / ".env") + + def test_image_tag_and_docling(self, tmp_path): + argv = ["up", "--yes", "--dir", str(tmp_path), "--image-tag", "develop", "--docling"] + docker = FakeDocker() + assert _run(argv, _context(docker)) == 0 + env = envfile.read(tmp_path / ".env") + assert env["DOCSGPT_IMAGE_TAG"] == "develop" + assert env["DOCSGPT_IMAGE_VARIANT"] == "-docling" + # A moving tag is pulled every time, not only when missing. + assert (tmp_path, ["up", "-d", "--remove-orphans", "--pull", "always"]) in docker.calls + + def test_an_unhealthy_start_points_at_the_logs(self, tmp_path, capsys): + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context(healthy=False)) == 1 + assert "docsgpt logs" in capsys.readouterr().err + assert not (tmp_path / "install.json").exists() + + +class TestUpAgain: + def test_keeps_the_settings_and_moves_the_version(self, tmp_path): + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context()) == 0 + before = envfile.read(tmp_path / ".env") + envfile.update(tmp_path / ".env", {"CUSTOM": "mine"}) + + context = _context(interactive=True, docker=FakeDocker(volumes={"docsgpt_postgres_data"})) + context.version = "0.22.0" + assert _run(["up", "--dir", str(tmp_path)], context) == 0 + after = envfile.read(tmp_path / ".env") + assert after["DOCSGPT_IMAGE_TAG"] == "0.22.0" + assert after["CUSTOM"] == "mine" + for key in ("INTERNAL_KEY", "JWT_SECRET_KEY", "POSTGRES_PASSWORD"): + assert after[key] == before[key] + assert context.prompter.questions == [], "a configured install is not asked again" + + def test_another_projects_containers_need_consent(self, tmp_path): + docker = FakeDocker(project_dirs={"/srv/old-docsgpt"}) + with pytest.raises(DeployError, match="--adopt"): + _run(["up", "--yes", "--dir", str(tmp_path)], _context(docker)) + assert docker.calls == [] + assert _run(["up", "--yes", "--adopt", "--dir", str(tmp_path)], _context(docker)) == 0 + + def test_consent_can_be_given_at_the_prompt(self, tmp_path): + docker = FakeDocker(project_dirs={"/srv/old-docsgpt"}) + prompter = FakePrompter([True, "local", "docsgpt"]) + assert _run(["up", "--dir", str(tmp_path)], _context(docker, prompter, interactive=True)) == 0 + + +class TestManage: + @staticmethod + def _installed(tmp_path, *extra): + assert _run(["up", "--yes", "--dir", str(tmp_path), *extra], _context()) == 0 + + def test_down(self, tmp_path): + self._installed(tmp_path) + docker = FakeDocker() + assert _run(["down", "--dir", str(tmp_path)], _context(docker)) == 0 + assert docker.calls == [(tmp_path, ["down"])] + + def test_commands_on_a_missing_install(self, tmp_path, capsys): + assert _run(["status", "--dir", str(tmp_path)], _context()) == 1 + assert "docsgpt up" in capsys.readouterr().err + + def test_status(self, tmp_path, capsys): + self._installed(tmp_path) + docker = FakeDocker() + assert _run(["status", "--dir", str(tmp_path)], _context(docker)) == 0 + out = capsys.readouterr().out + assert "0.21.0" in out and "http://localhost:7091" in out + assert (tmp_path, ["ps"]) in docker.calls + + def test_logs_pass_through(self, tmp_path): + self._installed(tmp_path) + docker = FakeDocker() + assert _run(["logs", "--dir", str(tmp_path), "-f", "--tail", "50", "backend"], _context(docker)) == 0 + assert docker.calls == [(tmp_path, ["logs", "--follow", "--tail", "50", "backend"])] + + def test_token(self, tmp_path, capsys): + self._installed(tmp_path, "--expose", "network") + capsys.readouterr() + assert _run(["token", "--dir", str(tmp_path)], _context()) == 0 + secret = envfile.read(tmp_path / ".env")["JWT_SECRET_KEY"] + assert capsys.readouterr().out.strip() == stack.simple_jwt_token(secret) + + def test_no_token_without_simple_jwt(self, tmp_path, capsys): + self._installed(tmp_path) + assert _run(["token", "--dir", str(tmp_path)], _context()) == 1 + assert "AUTH_TYPE" in capsys.readouterr().err + + def test_env_get_and_set(self, tmp_path, capsys): + self._installed(tmp_path) + capsys.readouterr() + assert _run(["env", "--dir", str(tmp_path), "set", "LLM_NAME=gpt-5.5", "OCR_ENABLED=true"], _context()) == 0 + assert "docsgpt up" in capsys.readouterr().out + assert _run(["env", "--dir", str(tmp_path), "get", "LLM_NAME"], _context()) == 0 + assert capsys.readouterr().out.strip() == "gpt-5.5" + assert _run(["env", "--dir", str(tmp_path), "get", "NOPE"], _context()) == 1 + with pytest.raises(DeployError, match="KEY=VALUE"): + _run(["env", "--dir", str(tmp_path), "set", "oops"], _context()) + + def test_uninstall_keeps_data_and_settings(self, tmp_path, capsys): + self._installed(tmp_path) + docker = FakeDocker() + assert _run(["uninstall", "--yes", "--dir", str(tmp_path)], _context(docker, installer=lambda: "uv")) == 0 + assert docker.calls == [(tmp_path, ["down", "--remove-orphans"])] + assert (tmp_path / ".env").is_file(), "the database password lives there" + assert not (tmp_path / "docker-compose.yaml").exists() + assert not (tmp_path / "install.json").exists() + assert "uv tool uninstall docsgpt" in capsys.readouterr().out + + def test_uninstall_purge_removes_everything(self, tmp_path): + directory = tmp_path / "stack" + self._installed(directory) + docker = FakeDocker() + assert _run(["uninstall", "--yes", "--purge", "--dir", str(directory)], _context(docker)) == 0 + assert docker.calls == [(directory, ["down", "--remove-orphans", "-v"])] + assert not directory.exists() + + def test_uninstall_asks_first(self, tmp_path): + self._installed(tmp_path) + docker = FakeDocker() + prompter = FakePrompter([False]) + assert _run(["uninstall", "--dir", str(tmp_path)], _context(docker, prompter, interactive=True)) == 1 + assert docker.calls == [] + + +class TestUpgrade: + def test_a_uv_tool_install_upgrades_and_runs_up_again(self, tmp_path): + self_calls = [] + context = _context( + installer=lambda: "uv", + run=lambda args: self_calls.append(args) or 0, + exec_up=lambda argv: self_calls.append(["exec", *argv]) or 0, + ) + assert _run(["upgrade", "--dir", str(tmp_path), "--version", "0.22.0"], context) == 0 + assert self_calls[0] == ["uv", "tool", "install", "--force", "docsgpt==0.22.0"] + assert self_calls[1] == ["exec", "docsgpt", "up", "--dir", str(tmp_path)] + + def test_latest_when_no_version_is_given(self, tmp_path): + self_calls = [] + context = _context(installer=lambda: "uv", run=lambda args: self_calls.append(args) or 0, + exec_up=lambda argv: 0) + assert _run(["upgrade", "--dir", str(tmp_path)], context) == 0 + assert self_calls[0] == ["uv", "tool", "install", "--force", "docsgpt"] + + def test_a_pip_install_is_told_what_to_run(self, tmp_path, capsys): + context = _context(installer=lambda: "pip", run=lambda args: pytest.fail("must not run")) + assert _run(["upgrade", "--dir", str(tmp_path)], context) == 1 + err = capsys.readouterr().err + assert "pip install -U docsgpt" in err and "docsgpt up" in err + + +class TestImports: + def test_the_deploy_commands_do_not_boot_the_app(self): + code = ( + "import sys, docsgpt.cli, docsgpt.deploy.commands; " + "loaded = {m for m in sys.modules if m in ('docsgpt.app', 'docsgpt.core.settings', 'celery', 'flask')}; " + "assert not loaded, loaded" + ) + subprocess.run([sys.executable, "-c", code], cwd=Path(__file__).resolve().parents[2], check=True) diff --git a/tests/deploy/test_docker.py b/tests/deploy/test_docker.py new file mode 100644 index 00000000..3354438c --- /dev/null +++ b/tests/deploy/test_docker.py @@ -0,0 +1,133 @@ +"""The docker CLI calls behind `docsgpt up`, against a fake runner.""" + +import subprocess +from pathlib import Path + +import pytest + +from docsgpt.deploy import docker as docker_module +from docsgpt.deploy.docker import DeployError, Docker + + +class FakeRunner: + """Answers docker commands from a table of (argv prefix -> result) and records every call.""" + + def __init__(self, answers=None): + self.answers = list((answers or {}).items()) + self.calls = [] + + def __call__(self, args, *, cwd=None, capture=False, check=True): + self.calls.append((list(args), cwd)) + for prefix, answer in self.answers: + if list(args[: len(prefix)]) == list(prefix): + if callable(answer): + answer = answer() + code, out = answer + if check and code != 0: + raise DeployError(f"{' '.join(args)} failed") + return subprocess.CompletedProcess(args, code, stdout=out, stderr="") + return subprocess.CompletedProcess(args, 0, stdout="", stderr="") + + +def _docker(runner, platform="linux", which=lambda name: "/usr/bin/docker"): + return Docker(runner=runner, which=which, sleep=lambda seconds: None, platform=platform) + + +class TestPreflight: + def test_docker_missing(self): + with pytest.raises(DeployError, match="Docker is not installed"): + _docker(FakeRunner(), which=lambda name: None).preflight() + + def test_daemon_down_on_linux_says_how_to_start_it(self): + runner = FakeRunner({("docker", "info"): (1, "")}) + with pytest.raises(DeployError, match="systemctl start docker"): + _docker(runner).preflight() + + def test_daemon_down_on_macos_starts_docker_desktop(self): + state = {"started": False} + + def info(): + return (0, "") if state["started"] else (1, "") + + def start(): + state["started"] = True + return (0, "") + + runner = FakeRunner( + { + ("docker", "info"): info, + ("open", "-a", "Docker"): start, + ("docker", "compose", "version"): (0, "v2.39.1-desktop.1\n"), + } + ) + _docker(runner, platform="darwin").preflight() + assert (["open", "-a", "Docker"], None) in runner.calls + + @pytest.mark.parametrize("version", ["v2.23.0", "2.20.2"]) + def test_compose_too_old(self, version): + runner = FakeRunner({("docker", "compose", "version"): (0, version + "\n")}) + with pytest.raises(DeployError, match="2.24"): + _docker(runner).preflight() + + def test_compose_missing(self): + runner = FakeRunner({("docker", "compose", "version"): (1, "")}) + with pytest.raises(DeployError, match="Docker Compose"): + _docker(runner).preflight() + + @pytest.mark.parametrize("version", ["v2.24.0", "2.39.1-desktop.1", "5.5.1"]) + def test_compose_new_enough(self, version): + runner = FakeRunner({("docker", "compose", "version"): (0, version + "\n")}) + _docker(runner).preflight() + + +class TestQueries: + def test_compose_runs_in_the_stack_directory(self, tmp_path): + runner = FakeRunner() + _docker(runner).compose(tmp_path, "up", "-d") + assert runner.calls == [(["docker", "compose", "up", "-d"], tmp_path)] + + def test_volume_exists(self): + runner = FakeRunner({("docker", "volume", "inspect", "docsgpt_postgres_data"): (0, "[]")}) + assert _docker(runner).volume_exists("docsgpt_postgres_data") + assert not _docker(FakeRunner({("docker", "volume"): (1, "")})).volume_exists("docsgpt_postgres_data") + + def test_project_directories(self): + runner = FakeRunner({("docker", "ps"): (0, "/srv/old\n/srv/old\n\n/home/me/.docsgpt/server\n")}) + assert _docker(runner).project_dirs("docsgpt") == {Path("/srv/old"), Path("/home/me/.docsgpt/server")} + args = runner.calls[0][0] + assert "label=com.docker.compose.project=docsgpt" in args + + +class TestWaitHealthy: + def test_succeeds_once_the_api_answers(self, monkeypatch): + attempts = {"n": 0} + + def opener(url, timeout): + attempts["n"] += 1 + if attempts["n"] < 3: + raise OSError("connection refused") + return _Response(200) + + clock = iter(range(0, 1000, 2)) + assert docker_module.wait_healthy("http://127.0.0.1:7091/api/health", 60, opener=opener, + sleep=lambda s: None, clock=lambda: next(clock)) + assert attempts["n"] == 3 + + def test_gives_up_after_the_timeout(self): + def opener(url, timeout): + raise OSError("connection refused") + + clock = iter(range(0, 1000, 10)) + assert not docker_module.wait_healthy("http://127.0.0.1:7091/api/health", 30, opener=opener, + sleep=lambda s: None, clock=lambda: next(clock)) + + +class _Response: + def __init__(self, status): + self.status = status + + def __enter__(self): + return self + + def __exit__(self, *exc): + return False diff --git a/tests/deploy/test_envfile.py b/tests/deploy/test_envfile.py new file mode 100644 index 00000000..d2997d1b --- /dev/null +++ b/tests/deploy/test_envfile.py @@ -0,0 +1,72 @@ +"""Reading and updating the stack's .env without losing what the user wrote.""" + +import os +import sys + +import pytest + +from docsgpt.deploy import envfile + + +class TestRead: + def test_comments_blanks_quotes_and_export(self, tmp_path): + path = tmp_path / ".env" + path.write_text( + "# settings\n" + "\n" + "LLM_PROVIDER=openai\n" + "export API_KEY='sk-123'\n" + 'LLM_NAME="gpt 5"\n' + "EMPTY=\n" + "not a line\n" + ) + assert envfile.read(path) == { + "LLM_PROVIDER": "openai", + "API_KEY": "sk-123", + "LLM_NAME": "gpt 5", + "EMPTY": "", + } + + def test_a_missing_file_is_empty(self, tmp_path): + assert envfile.read(tmp_path / ".env") == {} + + def test_the_last_duplicate_wins(self, tmp_path): + path = tmp_path / ".env" + path.write_text("A=1\nA=2\n") + assert envfile.read(path) == {"A": "2"} + + +class TestUpdate: + def test_changes_values_in_place_and_keeps_everything_else(self, tmp_path): + path = tmp_path / ".env" + path.write_text("# my settings\nLLM_PROVIDER=openai\n\nCUSTOM=keep me\nDOCSGPT_IMAGE_TAG=0.19.0\n") + envfile.update(path, {"DOCSGPT_IMAGE_TAG": "0.21.0"}) + assert path.read_text() == "# my settings\nLLM_PROVIDER=openai\n\nCUSTOM=keep me\nDOCSGPT_IMAGE_TAG=0.21.0\n" + + def test_appends_new_keys_and_removes_none(self, tmp_path): + path = tmp_path / ".env" + path.write_text("A=1\nB=2") + envfile.update(path, {"B": None, "C": "3"}) + assert path.read_text() == "A=1\nC=3\n" + + def test_duplicates_collapse_to_the_first_line(self, tmp_path): + path = tmp_path / ".env" + path.write_text("A=1\nX=y\nA=2\n") + envfile.update(path, {"A": "3"}) + assert path.read_text() == "A=3\nX=y\n" + + @pytest.mark.parametrize("value", ["plain", "with space", "hash # inside", "it's", 'say "hi"', "back\\slash", ""]) + def test_values_round_trip(self, tmp_path, value): + path = tmp_path / ".env" + envfile.update(path, {"VALUE": value}) + assert envfile.read(path)["VALUE"] == value + + @pytest.mark.skipif(sys.platform == "win32", reason="POSIX permissions") + def test_a_new_file_is_private(self, tmp_path): + path = tmp_path / "stack" / ".env" + envfile.update(path, {"JWT_SECRET_KEY": "s"}) + assert oct(os.stat(path).st_mode & 0o777) == oct(0o600) + + def test_rejects_a_newline_in_a_value(self, tmp_path): + with pytest.raises(ValueError, match="newline"): + envfile.update(tmp_path / ".env", {"A": "1\n2"}) diff --git a/tests/deploy/test_stack.py b/tests/deploy/test_stack.py new file mode 100644 index 00000000..e32c2393 --- /dev/null +++ b/tests/deploy/test_stack.py @@ -0,0 +1,196 @@ +"""What `docsgpt up` writes to the stack's .env, and where the stack lives.""" + +from itertools import count +from pathlib import Path + +import pytest + +from docsgpt.core import paths +from docsgpt.deploy import stack + +REPO_ROOT = Path(__file__).resolve().parents[2] + + +def _secrets(): + numbers = count(1) + return lambda: f"secret{next(numbers)}" + + +def _first_install(**overrides): + options = {"image_tag": "0.21.0", "fresh_database": True, "secret": _secrets()} + options.update(overrides) + return stack.plan({}, **options) + + +class TestFirstInstall: + def test_defaults_are_local_with_the_public_api(self): + updates = _first_install() + assert updates["DOCSGPT_IMAGE_TAG"] == "0.21.0" + assert updates["DOCSGPT_BIND"] == "127.0.0.1" + assert updates["LLM_PROVIDER"] == "docsgpt" + assert updates["VITE_API_STREAMING"] == "true" + assert "AUTH_TYPE" not in updates + + def test_secrets_are_generated_once_each(self): + updates = _first_install() + generated = {updates["INTERNAL_KEY"], updates["JWT_SECRET_KEY"], updates["POSTGRES_PASSWORD"]} + assert len(generated) == 3 + + def test_an_existing_database_keeps_its_password(self): + """Postgres reads the password only when its volume is created.""" + updates = _first_install(fresh_database=False) + assert "POSTGRES_PASSWORD" not in updates + + +class TestRerun: + def test_secrets_and_settings_are_left_alone(self): + existing = { + "DOCSGPT_IMAGE_TAG": "0.20.0", + "INTERNAL_KEY": "k", + "JWT_SECRET_KEY": "j", + "POSTGRES_PASSWORD": "p", + "VITE_API_STREAMING": "true", + "LLM_PROVIDER": "anthropic", + "API_KEY": "sk", + "DOCSGPT_BIND": "0.0.0.0", + "AUTH_TYPE": "simple_jwt", + } + updates = stack.plan(existing, image_tag="0.21.0", fresh_database=False, secret=_secrets()) + assert updates == {"DOCSGPT_IMAGE_TAG": "0.21.0"} + + def test_a_missing_password_is_not_invented_for_an_existing_database(self): + updates = stack.plan({"INTERNAL_KEY": "k", "JWT_SECRET_KEY": "j"}, image_tag="x", fresh_database=False) + assert "POSTGRES_PASSWORD" not in updates + + +class TestExposure: + def test_network_publishes_everywhere_and_turns_on_auth(self): + updates = _first_install(expose="network") + assert updates["DOCSGPT_BIND"] == "0.0.0.0" + assert updates["AUTH_TYPE"] == "simple_jwt" + assert updates.get("COMPOSE_PROFILES") is None + + def test_domain_adds_caddy_and_auth_and_keeps_the_port_local(self): + updates = _first_install(expose="domain", domain="docs.example.com") + assert updates["COMPOSE_PROFILES"] == "https" + assert updates["DOCSGPT_DOMAIN"] == "docs.example.com" + assert updates["DOCSGPT_BIND"] == "127.0.0.1" + assert updates["AUTH_TYPE"] == "simple_jwt" + + def test_domain_needs_a_domain(self): + with pytest.raises(ValueError, match="domain"): + _first_install(expose="domain") + + def test_an_existing_auth_mode_is_not_downgraded(self): + updates = stack.plan({"AUTH_TYPE": "oidc"}, image_tag="x", expose="network", fresh_database=False) + assert "AUTH_TYPE" not in updates + + def test_back_to_local_removes_the_proxy(self): + existing = {"COMPOSE_PROFILES": "https", "DOCSGPT_DOMAIN": "docs.example.com", "DOCSGPT_BIND": "127.0.0.1"} + updates = stack.plan(existing, image_tag="x", expose="local", fresh_database=False) + assert updates["COMPOSE_PROFILES"] is None + assert updates["DOCSGPT_DOMAIN"] is None + + def test_port_and_docling(self): + updates = _first_install(port=8080, docling=True) + assert updates["DOCSGPT_PORT"] == "8080" + assert updates["DOCSGPT_IMAGE_VARIANT"] == "-docling" + assert stack.plan({"DOCSGPT_IMAGE_VARIANT": "-docling"}, image_tag="x", docling=False, fresh_database=False)[ + "DOCSGPT_IMAGE_VARIANT" + ] is None + + @pytest.mark.parametrize( + "env, mode", + [ + ({}, "local"), + ({"DOCSGPT_BIND": "0.0.0.0"}, "network"), + ({"COMPOSE_PROFILES": "https", "DOCSGPT_DOMAIN": "d.example.com"}, "domain"), + ], + ) + def test_the_mode_is_read_back_from_the_env(self, env, mode): + assert stack.exposure(env) == mode + + +class TestProviders: + def test_switching_provider_drops_the_old_keys(self): + existing = {"LLM_PROVIDER": "openai", "API_KEY": "sk", "LLM_NAME": "m", "OPENAI_BASE_URL": "http://x/v1"} + provider = stack.provider_settings("anthropic", api_key="ak") + updates = stack.plan(existing, image_tag="x", provider=provider, fresh_database=False) + assert updates["LLM_PROVIDER"] == "anthropic" + assert updates["API_KEY"] == "ak" + assert updates["LLM_NAME"] is None + assert updates["OPENAI_BASE_URL"] is None + + def test_the_public_api_needs_no_key(self): + assert stack.provider_settings("docsgpt") == { + "LLM_PROVIDER": "docsgpt", + "API_KEY": None, + "LLM_NAME": None, + "OPENAI_BASE_URL": None, + } + + def test_a_hosted_provider_needs_a_key(self): + with pytest.raises(ValueError, match="API key"): + stack.provider_settings("openai") + + def test_an_openai_compatible_server_needs_a_url_and_a_model(self): + with pytest.raises(ValueError, match="base URL"): + stack.provider_settings("openai-compatible", model="llama3") + settings = stack.provider_settings("openai-compatible", base_url="http://host.docker.internal:11434/v1", model="llama3") + assert settings == { + "LLM_PROVIDER": "openai", + "API_KEY": "not-needed", + "LLM_NAME": "llama3", + "OPENAI_BASE_URL": "http://host.docker.internal:11434/v1", + } + + def test_an_unknown_provider(self): + with pytest.raises(ValueError, match="unknown provider"): + stack.provider_settings("nope") + + +class TestUrls: + def test_local(self): + assert stack.url({"DOCSGPT_PORT": "8080"}, lan_ip="10.0.0.5") == "http://localhost:8080" + + def test_network_uses_the_machine_address(self): + assert stack.url({"DOCSGPT_BIND": "0.0.0.0"}, lan_ip="10.0.0.5") == "http://10.0.0.5:7091" + + def test_domain(self): + env = {"COMPOSE_PROFILES": "https", "DOCSGPT_DOMAIN": "docs.example.com"} + assert stack.url(env, lan_ip="10.0.0.5") == "https://docs.example.com" + + def test_health_is_always_checked_on_this_machine(self): + assert stack.health_url({"DOCSGPT_BIND": "0.0.0.0", "DOCSGPT_PORT": "9000"}) == "http://127.0.0.1:9000/api/health" + + +class TestToken: + def test_matches_what_the_api_prints(self): + """docsgpt/app.py signs {"sub": "local"} with JWT_SECRET_KEY for AUTH_TYPE=simple_jwt.""" + from jose import jwt + + token = stack.simple_jwt_token("s3cret") + assert token == jwt.encode({"sub": "local"}, "s3cret", algorithm="HS256") + assert jwt.decode(token, "s3cret", algorithms=["HS256"]) == {"sub": "local"} + + +class TestLocations: + def test_the_stack_dir_is_never_the_checkout(self, monkeypatch, tmp_path): + monkeypatch.delenv(paths.HOME_ENV, raising=False) + monkeypatch.setattr(paths, "default_home", lambda: tmp_path / "home") + assert stack.stack_dir(None) == tmp_path / "home" + + def test_docsgpt_home_and_then_an_explicit_dir_win(self, monkeypatch, tmp_path): + monkeypatch.setenv(paths.HOME_ENV, str(tmp_path / "env-home")) + assert stack.stack_dir(None) == (tmp_path / "env-home").resolve() + assert stack.stack_dir(str(tmp_path / "flag")) == (tmp_path / "flag").resolve() + + def test_a_checkout_uses_the_deployment_compose_file(self): + assert stack.compose_source() == REPO_ROOT / "deployment" / "docker-compose-standalone.yaml" + + def test_the_packaged_compose_file_wins(self, monkeypatch, tmp_path): + packaged = tmp_path / "docsgpt" / "deploy" / "docker-compose.yaml" + packaged.parent.mkdir(parents=True) + packaged.write_text("name: docsgpt\n") + monkeypatch.setattr(paths, "package_dir", lambda: tmp_path / "docsgpt") + assert stack.compose_source() == packaged From 68d9772f2eeb31d452b1b4fd52a833495904ebf1 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:00:23 +0100 Subject: [PATCH 007/130] docs: docsgpt up and the new data home; CI runs up against the built image Docker-Deploying gains a `docsgpt up` section, Pip-Install and Upgrading describe the ~/.docsgpt/server data home, and the changelog covers both. docker-image-verify.yml installs the wheel and runs `docsgpt up`, `status`, a second `up` that must keep the secrets, and `uninstall --purge` against the image it built. The standalone Compose file maps host.docker.internal to the host gateway, so a model server on a Linux host is reachable the way `docsgpt up` suggests. --- .github/workflows/docker-image-verify.yml | 30 ++++++++++++ deployment/docker-compose-standalone.yaml | 15 ++++-- docs/content/Deploying/Docker-Deploying.mdx | 52 +++++++++++++++++++++ docs/content/Deploying/Pip-Install.mdx | 12 ++++- docs/content/changelog.mdx | 16 +++++++ docs/content/upgrading.mdx | 4 ++ 6 files changed, 123 insertions(+), 6 deletions(-) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index fec67707..3d162675 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -23,6 +23,10 @@ on: - 'frontend/**' - 'scripts/build_frontend.sh' - 'deployment/docker-compose-standalone.yaml' + - 'docsgpt/deploy/**' + - 'docsgpt/cli.py' + - 'docsgpt/core/paths.py' + - 'pyproject.toml' - '.github/workflows/docker-image-verify.yml' permissions: @@ -94,6 +98,32 @@ jobs: curl -fsS "$base/settings" | grep -q 'src="/config.js"' echo "API and UI served on $base" + - name: Set up uv + if: matrix.variant == '' + uses: astral-sh/setup-uv@d0cc045d04ccac9d8b7881df0226f9e82c39688e # v6.8.0 + + - name: docsgpt up runs the same image from the installed package + if: matrix.variant == '' + run: | + set -euo pipefail + # Same Compose project name as the step above; stop that stack first. + docker compose -f deployment/docker-compose-standalone.yaml down -v + uv build --wheel --out-dir "$RUNNER_TEMP/dist" + uv venv "$RUNNER_TEMP/venv" + uv pip install --python "$RUNNER_TEMP/venv/bin/python" "$RUNNER_TEMP"/dist/*.whl + docsgpt="$RUNNER_TEMP/venv/bin/docsgpt" + stack="$RUNNER_TEMP/stack" + "$docsgpt" up --yes --dir "$stack" --image-tag verify + "$docsgpt" status --dir "$stack" + curl -fsS http://127.0.0.1:7091/ | grep -q 'src="/config.js"' + grep -q '^POSTGRES_PASSWORD=' "$stack/.env" + # Running up again keeps the generated secrets. + before=$(grep '^JWT_SECRET_KEY=' "$stack/.env") + "$docsgpt" up --yes --dir "$stack" --image-tag verify + [ "$(grep '^JWT_SECRET_KEY=' "$stack/.env")" = "$before" ] + "$docsgpt" uninstall --yes --purge --dir "$stack" + test ! -e "$stack" + - name: Stack logs if: failure() env: diff --git a/deployment/docker-compose-standalone.yaml b/deployment/docker-compose-standalone.yaml index ca69253b..52d38d06 100644 --- a/deployment/docker-compose-standalone.yaml +++ b/deployment/docker-compose-standalone.yaml @@ -9,10 +9,11 @@ # the API; without it every ingest fails with a 401 (setup.sh generates one). # # The backend image serves the web UI and the API on one port. Every release -# also attaches this file as an asset. Settings come from .env next to this -# file (any DocsGPT setting, VITE_* included; the compose-internal service URLs -# below take precedence). Data lives in named volumes, so `docker compose down` -# keeps it and `docker compose down -v` removes it. +# also attaches this file as an asset, and `docsgpt up` runs it from the Python +# package. Settings come from .env next to this file (any DocsGPT setting, +# VITE_* included; the compose-internal service URLs below take precedence). +# Data lives in named volumes, so `docker compose down` keeps it and +# `docker compose down -v` removes it. # # DOCSGPT_IMAGE_TAG release to run, e.g. 0.20.0 (default: latest release); # develop follows the main branch @@ -55,6 +56,10 @@ services: - CACHE_REDIS_URL=redis://redis:6379/2 - POSTGRES_URI=postgresql://docsgpt:${POSTGRES_PASSWORD:-docsgpt}@postgres:5432/docsgpt - EMBEDDINGS_NAME=${EMBEDDINGS_NAME:-ibm-granite/granite-embedding-311m-multilingual-r2} + # A model server on the host (Ollama, vLLM, ...) is reachable as + # host.docker.internal on Linux too, as it is on Docker Desktop. + extra_hosts: + - "host.docker.internal:host-gateway" ports: - "${DOCSGPT_BIND:-127.0.0.1}:${DOCSGPT_PORT:-7091}:7091" volumes: @@ -84,6 +89,8 @@ services: - POSTGRES_URI=postgresql://docsgpt:${POSTGRES_PASSWORD:-docsgpt}@postgres:5432/docsgpt - API_URL=http://backend:7091 - EMBEDDINGS_NAME=${EMBEDDINGS_NAME:-ibm-granite/granite-embedding-311m-multilingual-r2} + extra_hosts: + - "host.docker.internal:host-gateway" volumes: - indexes:/app/indexes - inputs:/app/inputs diff --git a/docs/content/Deploying/Docker-Deploying.mdx b/docs/content/Deploying/Docker-Deploying.mdx index b9e24c90..fdc91229 100644 --- a/docs/content/Deploying/Docker-Deploying.mdx +++ b/docs/content/Deploying/Docker-Deploying.mdx @@ -17,6 +17,58 @@ Docker is the recommended method for deploying DocsGPT, providing a consistent a **Important Note for Windows Users:** Docker Desktop on Windows generally requires the WSL 2 backend to function correctly, especially when using features like host networking which are utilized in DocsGPT's Docker Compose setup. Ensure WSL 2 is enabled and configured in Docker Desktop settings. +## Run it with `docsgpt up` + +The `docsgpt` Python package can set up and run the stack described below for +you. It needs Docker with Compose 2.24 or newer, and Python 3.12 or newer (uv +installs one when it is missing): + +```bash +uv tool install docsgpt # or: pipx install docsgpt +docsgpt up +``` + +`docsgpt up` keeps the stack in `~/.docsgpt/server` (`/opt/docsgpt` when run as +root on Linux; `--dir` or `DOCSGPT_HOME` choose another folder): the Compose +file of the installed version, a `.env` with your settings and the generated +secrets, and `install.json`. Data lives in named Docker volumes. The first run +asks two questions: + +- **Who should reach DocsGPT:** only this computer; other machines on the + network (plain HTTP, with `AUTH_TYPE=simple_jwt` and an access token); or a + domain name with HTTPS (Caddy gets the certificate, access token as well). +- **Which model provider:** the DocsGPT public API (no key needed), OpenAI, + Anthropic, Google Gemini, OpenRouter, Groq, or an OpenAI-compatible server + such as Ollama or vLLM. + +Flags answer the same questions, for scripts and servers: + +```bash +docsgpt up --yes --domain docs.example.com --provider openai --api-key "$OPENAI_API_KEY" +``` + +Running `docsgpt up` again is safe: it keeps `.env` and the secrets and runs +the images of the installed package version. `docsgpt up --reconfigure` asks +the questions again. + +| Command | What it does | +| --- | --- | +| `docsgpt status` | Version, address, containers, and whether the API answers | +| `docsgpt logs [-f] [service]` | Container logs | +| `docsgpt token` | The access token, for installs reachable beyond this computer | +| `docsgpt open` | Open DocsGPT in the browser | +| `docsgpt env set KEY=VALUE` | Change a setting; `docsgpt up` applies it | +| `docsgpt upgrade` | Upgrade the package (for `uv tool` installs) and restart on the new version | +| `docsgpt down` | Stop the stack; data and settings stay | +| `docsgpt uninstall [--purge]` | Remove the containers; `--purge` also deletes the settings and all data | + +More `docsgpt up` options: `--port`, `--docling` (the image with the docling +parser engine and OCR), `--image-tag develop` (follow the `main` branch) and +`--adopt` (manage a stack you started from the standalone Compose file in +another folder; both use the same data volumes). For Ollama on the same +machine, use the base URL `http://host.docker.internal:11434/v1`; on Linux, +also make Ollama listen beyond localhost (`OLLAMA_HOST=0.0.0.0`). + ## Quickest Setup: Pre-built Images, No Checkout Every release publishes ready-to-run images to Docker Hub (`arc53/docsgpt`, diff --git a/docs/content/Deploying/Pip-Install.mdx b/docs/content/Deploying/Pip-Install.mdx index ffe92c1d..1b946687 100644 --- a/docs/content/Deploying/Pip-Install.mdx +++ b/docs/content/Deploying/Pip-Install.mdx @@ -13,6 +13,10 @@ DocsGPT is on PyPI as [`docsgpt`](https://pypi.org/project/docsgpt/): the API se `docsgpt api` serves the web UI on the same port as the API. Set `SERVE_UI=false` to run the API alone, for example behind the frontend Docker image or a UI you host yourself. + + With Docker available, the same package can run the whole stack for you, including Postgres and Redis: `docsgpt up`. See [Run it with `docsgpt up`](/Deploying/Docker-Deploying#run-it-with-docsgpt-up). + + ## Requirements - Python 3.12 or newer @@ -49,7 +53,11 @@ pipx runpip docsgpt install --force-reinstall --no-deps --index-url https://down ## Configure -DocsGPT keeps its runtime files in a **data home**: the `.env` file it reads settings from, uploaded files under `inputs/`, vector indexes under `indexes/` and downloaded embedding models under `models/`. The data home is the directory you run the commands from, or the directory `DOCSGPT_HOME` points to. `DOCSGPT_ENV_FILE` points at a `.env` kept somewhere else. Both variables must be set in the process environment, not in `.env`: they decide where `.env` is read from. +DocsGPT keeps its runtime files in a **data home**: the `.env` file it reads settings from, uploaded files under `inputs/`, vector indexes under `indexes/` and downloaded embedding models under `models/`. The data home is `~/.docsgpt/server` (`/opt/docsgpt` when you run as root on Linux), whatever directory you run the commands from. `DOCSGPT_HOME` moves it, and `DOCSGPT_ENV_FILE` points at a `.env` kept somewhere else. Both variables must be set in the process environment, not in `.env`: they decide where `.env` is read from. In a source checkout the data home is the checkout. + + + Up to 0.20 the data home of an installed package was the directory you ran the command from. If you kept `.env` and your data there, move them to `~/.docsgpt/server` or set `DOCSGPT_HOME` to that directory. `docsgpt api` and `docsgpt worker` point out a `.env` in the working directory that they no longer read. + Create a `.env` in the data home. The minimum for a hosted LLM: @@ -80,7 +88,7 @@ docsgpt worker # in a second terminal: the Celery worker, with the scheduler The API applies pending migrations when it starts (`AUTO_MIGRATE`), so `docsgpt migrate` is the explicit step for deployments that want the schema in place before the first request or that run the API with a restricted database role. -Both commands print the data home they resolved on start-up. Run them from the same directory, or set `DOCSGPT_HOME` for both, so the worker finds the files the API stores and the API finds the indexes the worker builds. +Both commands print the data home they resolved on start-up. They share it as long as `DOCSGPT_HOME` is the same for both (or unset), so the worker finds the files the API stores and the API finds the indexes the worker builds. The worker is not optional: query embedding runs on it, so search fails without one. `docsgpt worker --help` lists the queue, concurrency and pool options; `--no-beat` starts a worker without the scheduler when another worker already runs it. On Windows the scheduler cannot be embedded, so run `docsgpt beat` in a third terminal. diff --git a/docs/content/changelog.mdx b/docs/content/changelog.mdx index 0aa9324a..a655b79a 100644 --- a/docs/content/changelog.mdx +++ b/docs/content/changelog.mdx @@ -13,6 +13,22 @@ request, and [Upgrading](/upgrading) covers the steps an existing deployment has ## Unreleased +### `docsgpt up` runs DocsGPT on Docker + +The Python package now sets up and runs the Docker stack: `uv tool install docsgpt`, then +`docsgpt up`. The first run asks who should reach DocsGPT (this computer, the network with an +access token, or a domain with HTTPS) and which model provider to use, writes the settings and +secrets to `~/.docsgpt/server/.env`, and starts the images of the installed version. `docsgpt status`, +`logs`, `token`, `upgrade`, `down` and `uninstall` manage it afterwards. See +[Run it with `docsgpt up`](/Deploying/Docker-Deploying#run-it-with-docsgpt-up). + +### An installed package keeps its data in `~/.docsgpt/server` + +Outside a source checkout, the data home (`.env`, uploads, indexes, models) was the directory +the command ran from, so starting `docsgpt api` from another folder silently used other +settings. It is now `~/.docsgpt/server`, or `/opt/docsgpt` for root on Linux; `DOCSGPT_HOME` +still overrides it. See [Upgrading](/upgrading#pip-installs-data-home-moved). + ### The standalone Docker stack runs on one port The `arc53/docsgpt` image now serves the web UI next to the API, the way `docsgpt api` does from diff --git a/docs/content/upgrading.mdx b/docs/content/upgrading.mdx index 7df9a0e3..2e5d3032 100644 --- a/docs/content/upgrading.mdx +++ b/docs/content/upgrading.mdx @@ -11,6 +11,10 @@ import { Callout } from 'nextra/components' **Upgrading from 0.16.x?** User data moved from MongoDB to Postgres in 0.17.0. Follow the [Postgres Migration guide](/Deploying/Postgres-Migration) before running `docker compose pull` or `git pull` — existing deployments will not start cleanly without it. +## pip installs: data home moved + +An installed package (`pip install docsgpt`, pipx, `uv tool`) used to keep its data home, meaning `.env`, `inputs/`, `indexes/` and `models/`, in the directory you ran `docsgpt api` and `docsgpt worker` from. It is now `~/.docsgpt/server` (`/opt/docsgpt` for root on Linux). Either move those files there, or set `DOCSGPT_HOME` to the old directory in the environment of both commands. Both commands print the data home they use, and point out a `.env` in the working directory that they no longer read. Source checkouts and the Docker images are not affected. + ## Standalone Compose file: one port The standalone Compose file (`docker-compose-standalone.yaml`) no longer runs a frontend container. The backend image serves the web UI on port 7091, and the port is published on `127.0.0.1` unless you set `DOCSGPT_BIND`. After downloading the new file, start it with `--remove-orphans` to remove the old frontend container, then open port 7091 instead of 5173. If you opened DocsGPT from other machines, see [Upgrading from an earlier standalone file](/Deploying/Docker-Deploying#upgrading-from-an-earlier-standalone-file). The checkout Compose files and the Kubernetes manifests are unchanged. From 0bd0afc1f8a0bf192b2417de2e1eb2bcb4965077 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:06:27 +0100 Subject: [PATCH 008/130] fix: docsgpt up removes Caddy when leaving a domain; clearer Docker permission error Caddy sits behind the https profile, so once COMPOSE_PROFILES no longer enables it, `up --remove-orphans` left it running on ports 80 and 443 and `down -v` left its volumes. `up` now removes Caddy when an install moves off its domain, and down and uninstall name the profile explicitly. A user outside the docker group was told Docker is not running; the error now says how to get access to the socket. --- docsgpt/deploy/commands.py | 10 ++++++++-- docsgpt/deploy/docker.py | 8 +++++++- tests/deploy/test_commands.py | 33 ++++++++++++++++++++++++++------- tests/deploy/test_docker.py | 15 +++++++++++---- 4 files changed, 52 insertions(+), 14 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 140b7929..4367221b 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -23,6 +23,9 @@ PROJECT = "docsgpt" DATABASE_VOLUME = f"{PROJECT}_postgres_data" _MOVING_TAGS = ("latest", "develop") _KEY = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") +# Commands that stop or remove services name every profile, so Caddy goes too +# even when COMPOSE_PROFILES no longer enables it. +_EVERY_PROFILE = ("--profile", "https") EXPOSURE_CHOICES = [ ("local", "Only this computer"), @@ -208,6 +211,9 @@ def up(args, context: Optional[Context] = None) -> int: envfile.update(env_path, updates) env = envfile.read(env_path) + if stack.exposure(existing) == "domain" and stack.exposure(env) != "domain": + # With the https profile off, `up --remove-orphans` would leave Caddy running on ports 80 and 443. + context.docker.compose(directory, *_EVERY_PROFILE, "rm", "--stop", "--force", "caddy", check=False) up_args = ["up", "-d", "--remove-orphans"] if image_tag in _MOVING_TAGS: up_args += ["--pull", "always"] @@ -252,7 +258,7 @@ def down(args, context: Optional[Context] = None) -> int: directory = stack.stack_dir(args.dir) if _installed(directory) is None: return 1 - context.docker.compose(directory, "down") + context.docker.compose(directory, *_EVERY_PROFILE, "down") return 0 @@ -369,7 +375,7 @@ def uninstall(args, context: Optional[Context] = None) -> int: if not context.prompter.confirm(f"Remove the DocsGPT {what} in {directory}?", default=False): print("Nothing removed.") return 1 - context.docker.compose(directory, "down", "--remove-orphans", *(["-v"] if args.purge else [])) + context.docker.compose(directory, *_EVERY_PROFILE, "down", "--remove-orphans", *(["-v"] if args.purge else [])) if args.purge: shutil.rmtree(directory) print(f"Removed DocsGPT and its data from {directory}.") diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py index 0b3a1ea4..90ddeea0 100644 --- a/docsgpt/deploy/docker.py +++ b/docsgpt/deploy/docker.py @@ -59,7 +59,13 @@ class Docker: """Make sure Docker is installed and running and Compose is new enough, starting Docker Desktop on macOS.""" if not self._which("docker"): raise DeployError("Docker is not installed. Get it from https://docs.docker.com/get-docker/ and run this again.") - if not self.daemon_running(): + info = self._run(["docker", "info"], capture=True, check=False) + if info.returncode != 0: + if "permission denied" in (info.stderr or "").lower(): + raise DeployError( + "Your user cannot use Docker (permission denied on its socket). Add it to the docker group " + "with `sudo usermod -aG docker $USER`, log out and back in, and run this again." + ) self._start_daemon() result = self._run(["docker", "compose", "version", "--short"], capture=True, check=False) version = _parse_version(result.stdout) if result.returncode == 0 else None diff --git a/tests/deploy/test_commands.py b/tests/deploy/test_commands.py index 3aa80d5c..4cc963b6 100644 --- a/tests/deploy/test_commands.py +++ b/tests/deploy/test_commands.py @@ -11,6 +11,8 @@ from docsgpt import cli from docsgpt.deploy import commands, envfile, stack from docsgpt.deploy.docker import DeployError +EVERY_PROFILE = ["--profile", "https"] + class FakeDocker: def __init__(self, volumes=(), project_dirs=()): @@ -24,7 +26,7 @@ class FakeDocker: def compose(self, directory, *args, capture=False, check=True): self.calls.append((Path(directory), list(args))) - if args and args[0] == "down" and "-v" in args: + if "down" in args and "-v" in args: self.volumes.clear() return subprocess.CompletedProcess(["docker", "compose", *args], 0, stdout="", stderr="") @@ -88,7 +90,7 @@ class TestUpFirstInstall: assert env["LLM_PROVIDER"] == "docsgpt" assert env["INTERNAL_KEY"] and env["JWT_SECRET_KEY"] and env["POSTGRES_PASSWORD"] assert docker.preflights == 1 - assert (tmp_path, ["up", "-d", "--remove-orphans"]) in docker.calls + assert docker.calls == [(tmp_path, ["up", "-d", "--remove-orphans"])] record = json.loads((tmp_path / "install.json").read_text()) assert record["version"] == "0.21.0" assert "http://localhost:7091" in capsys.readouterr().out @@ -113,7 +115,8 @@ class TestUpFirstInstall: assert env["DOCSGPT_DOMAIN"] == "docs.example.com" assert env["LLM_PROVIDER"] == "openai" - def test_a_missing_api_key_is_an_error_without_a_terminal(self, tmp_path): + def test_a_missing_api_key_is_an_error_without_a_terminal(self, tmp_path, monkeypatch): + monkeypatch.delenv("DOCSGPT_API_KEY", raising=False) with pytest.raises(DeployError, match="API key"): _run(["up", "--yes", "--dir", str(tmp_path), "--provider", "openai"], _context()) @@ -159,6 +162,22 @@ class TestUpAgain: assert after[key] == before[key] assert context.prompter.questions == [], "a configured install is not asked again" + def test_leaving_the_domain_removes_caddy_before_starting(self, tmp_path): + """With the https profile off, `up --remove-orphans` alone would leave Caddy on ports 80 and 443.""" + assert _run(["up", "--yes", "--dir", str(tmp_path), "--domain", "docs.example.com"], _context()) == 0 + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + assert _run(["up", "--yes", "--dir", str(tmp_path), "--expose", "local"], _context(docker)) == 0 + assert docker.calls == [ + (tmp_path, [*EVERY_PROFILE, "rm", "--stop", "--force", "caddy"]), + (tmp_path, ["up", "-d", "--remove-orphans"]), + ] + + def test_staying_on_the_domain_leaves_caddy_alone(self, tmp_path): + assert _run(["up", "--yes", "--dir", str(tmp_path), "--domain", "docs.example.com"], _context()) == 0 + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context(docker)) == 0 + assert docker.calls == [(tmp_path, ["up", "-d", "--remove-orphans"])] + def test_another_projects_containers_need_consent(self, tmp_path): docker = FakeDocker(project_dirs={"/srv/old-docsgpt"}) with pytest.raises(DeployError, match="--adopt"): @@ -177,11 +196,11 @@ class TestManage: def _installed(tmp_path, *extra): assert _run(["up", "--yes", "--dir", str(tmp_path), *extra], _context()) == 0 - def test_down(self, tmp_path): + def test_down_includes_caddy(self, tmp_path): self._installed(tmp_path) docker = FakeDocker() assert _run(["down", "--dir", str(tmp_path)], _context(docker)) == 0 - assert docker.calls == [(tmp_path, ["down"])] + assert docker.calls == [(tmp_path, [*EVERY_PROFILE, "down"])] def test_commands_on_a_missing_install(self, tmp_path, capsys): assert _run(["status", "--dir", str(tmp_path)], _context()) == 1 @@ -228,7 +247,7 @@ class TestManage: self._installed(tmp_path) docker = FakeDocker() assert _run(["uninstall", "--yes", "--dir", str(tmp_path)], _context(docker, installer=lambda: "uv")) == 0 - assert docker.calls == [(tmp_path, ["down", "--remove-orphans"])] + assert docker.calls == [(tmp_path, [*EVERY_PROFILE, "down", "--remove-orphans"])] assert (tmp_path / ".env").is_file(), "the database password lives there" assert not (tmp_path / "docker-compose.yaml").exists() assert not (tmp_path / "install.json").exists() @@ -239,7 +258,7 @@ class TestManage: self._installed(directory) docker = FakeDocker() assert _run(["uninstall", "--yes", "--purge", "--dir", str(directory)], _context(docker)) == 0 - assert docker.calls == [(directory, ["down", "--remove-orphans", "-v"])] + assert docker.calls == [(directory, [*EVERY_PROFILE, "down", "--remove-orphans", "-v"])] assert not directory.exists() def test_uninstall_asks_first(self, tmp_path): diff --git a/tests/deploy/test_docker.py b/tests/deploy/test_docker.py index 3354438c..d3de76b6 100644 --- a/tests/deploy/test_docker.py +++ b/tests/deploy/test_docker.py @@ -10,7 +10,7 @@ from docsgpt.deploy.docker import DeployError, Docker class FakeRunner: - """Answers docker commands from a table of (argv prefix -> result) and records every call.""" + """Answers docker commands from a table of argv prefix -> (code, stdout[, stderr]) and records every call.""" def __init__(self, answers=None): self.answers = list((answers or {}).items()) @@ -22,10 +22,10 @@ class FakeRunner: if list(args[: len(prefix)]) == list(prefix): if callable(answer): answer = answer() - code, out = answer + code, out, err = (*answer, "")[:3] if check and code != 0: raise DeployError(f"{' '.join(args)} failed") - return subprocess.CompletedProcess(args, code, stdout=out, stderr="") + return subprocess.CompletedProcess(args, code, stdout=out, stderr=err) return subprocess.CompletedProcess(args, 0, stdout="", stderr="") @@ -43,6 +43,13 @@ class TestPreflight: with pytest.raises(DeployError, match="systemctl start docker"): _docker(runner).preflight() + def test_no_permission_on_the_socket_says_how_to_get_it(self): + """A user outside the docker group is not told that Docker is down.""" + denied = "permission denied while trying to connect to the Docker daemon socket at unix:///var/run/docker.sock" + runner = FakeRunner({("docker", "info"): (1, "", denied)}) + with pytest.raises(DeployError, match="docker group"): + _docker(runner).preflight() + def test_daemon_down_on_macos_starts_docker_desktop(self): state = {"started": False} @@ -99,7 +106,7 @@ class TestQueries: class TestWaitHealthy: - def test_succeeds_once_the_api_answers(self, monkeypatch): + def test_succeeds_once_the_api_answers(self): attempts = {"n": 0} def opener(url, timeout): From c7af5f873a75bf08b542ad12d7fd29fc8df47791 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:11:32 +0100 Subject: [PATCH 009/130] fix: docsgpt up --adopt recreates every container Compose keeps containers whose configuration did not change, so after a takeover Redis and Postgres still carried the other folder's working directory label, and the next `docsgpt up` asked for --adopt again. --- docsgpt/deploy/commands.py | 17 +++++++++++------ tests/deploy/test_commands.py | 11 +++++++++++ 2 files changed, 22 insertions(+), 6 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 4367221b..19d4aae2 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -130,17 +130,20 @@ def _current_provider(env: Mapping[str, str]) -> str: return name if name in stack.PROVIDERS else "docsgpt" -def _check_other_stacks(args, context: Context, directory: Path) -> None: +def _check_other_stacks(args, context: Context, directory: Path) -> bool: + """True when containers of the project were started from another folder and ``up`` takes them over.""" others = {path for path in context.docker.project_dirs(PROJECT) if path.resolve() != directory.resolve()} - if not others or args.adopt: - return + if not others: + return False + if args.adopt: + return True where = ", ".join(sorted(str(path) for path in others)) question = ( f"Docker already runs a DocsGPT stack started from {where}; it uses the same data volumes. " f"Manage it from {directory} instead?" ) if context.interactive and context.prompter.confirm(question, default=False): - return + return True raise DeployError( f"Docker already runs a DocsGPT stack started from {where}. Stop it with `docker compose down` " f"in that folder, or run again with --adopt to manage it from {directory}." @@ -179,7 +182,9 @@ def up(args, context: Optional[Context] = None) -> int: configured = record_path.is_file() context.docker.preflight(context.interactive) - _check_other_stacks(args, context, directory) + # Compose keeps containers whose configuration did not change, and with them the other + # folder's working-directory label; recreating them all makes the takeover complete. + recreate = ["--force-recreate"] if _check_other_stacks(args, context, directory) else [] ask = context.interactive and (not configured or args.reconfigure) expose, domain = args.expose, args.domain @@ -214,7 +219,7 @@ def up(args, context: Optional[Context] = None) -> int: if stack.exposure(existing) == "domain" and stack.exposure(env) != "domain": # With the https profile off, `up --remove-orphans` would leave Caddy running on ports 80 and 443. context.docker.compose(directory, *_EVERY_PROFILE, "rm", "--stop", "--force", "caddy", check=False) - up_args = ["up", "-d", "--remove-orphans"] + up_args = ["up", "-d", "--remove-orphans", *recreate] if image_tag in _MOVING_TAGS: up_args += ["--pull", "always"] print(f"Starting DocsGPT {image_tag} from {directory} ...") diff --git a/tests/deploy/test_commands.py b/tests/deploy/test_commands.py index 4cc963b6..44398703 100644 --- a/tests/deploy/test_commands.py +++ b/tests/deploy/test_commands.py @@ -185,6 +185,17 @@ class TestUpAgain: assert docker.calls == [] assert _run(["up", "--yes", "--adopt", "--dir", str(tmp_path)], _context(docker)) == 0 + def test_adopting_recreates_every_container(self, tmp_path): + """Compose keeps unchanged containers, and with them the old folder's label; the next up would ask again.""" + docker = FakeDocker(project_dirs={"/srv/old-docsgpt"}) + assert _run(["up", "--yes", "--adopt", "--dir", str(tmp_path)], _context(docker)) == 0 + assert docker.calls == [(tmp_path, ["up", "-d", "--remove-orphans", "--force-recreate"])] + + def test_containers_from_this_folder_are_not_recreated(self, tmp_path): + docker = FakeDocker(project_dirs={str(tmp_path)}) + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context(docker)) == 0 + assert docker.calls == [(tmp_path, ["up", "-d", "--remove-orphans"])] + def test_consent_can_be_given_at_the_prompt(self, tmp_path): docker = FakeDocker(project_dirs={"/srv/old-docsgpt"}) prompter = FakePrompter([True, "local", "docsgpt"]) From 5e699b0168c01d007e52e47d35f14cd41fe3a4d4 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:20:16 +0100 Subject: [PATCH 010/130] ci: no uv cache in the image verify job (zizmor cache-poisoning) --- .github/workflows/docker-image-verify.yml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index 3d162675..f59bbf9a 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -101,6 +101,9 @@ jobs: - name: Set up uv if: matrix.variant == '' uses: astral-sh/setup-uv@d0cc045d04ccac9d8b7881df0226f9e82c39688e # v6.8.0 + with: + # No cache: a cache restored into a job that runs the built image is a poisoning risk. + enable-cache: false - name: docsgpt up runs the same image from the installed package if: matrix.variant == '' From f6bf4fefd149956f5af43e9f43906dc8a80948f1 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:32:37 +0100 Subject: [PATCH 011/130] fix: docsgpt up review follow-ups - wait_healthy starts no request once the deadline is reached, and neither its pauses nor the requests after the first run past it; the first attempt still always runs (status uses a zero timeout). - envfile writes $ as $$ inside double quotes, which Compose interpolates, and reads $$ back as $, so a value such as pa$w'rd reaches the container unchanged. - Upgrading no longer describes the working directory as the data home. --- docs/content/upgrading.mdx | 5 +++-- docsgpt/deploy/docker.py | 18 ++++++++++++++---- docsgpt/deploy/envfile.py | 8 +++++--- tests/deploy/test_docker.py | 32 ++++++++++++++++++++++++++++++++ tests/deploy/test_envfile.py | 11 ++++++++++- 5 files changed, 64 insertions(+), 10 deletions(-) diff --git a/docs/content/upgrading.mdx b/docs/content/upgrading.mdx index 2e5d3032..c5397390 100644 --- a/docs/content/upgrading.mdx +++ b/docs/content/upgrading.mdx @@ -121,8 +121,9 @@ alias, so nothing breaks on upgrade, but update these before the alias goes: and vectors under `application/` in the checkout, where they already are. - The backend is also a package now (`pip install docsgpt`, see [Install with pip](/Deploying/Pip-Install)). Runtime data lives in a data - home: `DOCSGPT_HOME`, else the checkout, else the directory the process - starts in. One consequence for a source checkout: the embedded Milvus + home: `DOCSGPT_HOME`, else the checkout, else `~/.docsgpt/server` + (`/opt/docsgpt` for root on Linux; see + [pip installs: data home moved](#pip-installs-data-home-moved)). One consequence for a source checkout: the embedded Milvus (`MILVUS_URI`) and LanceDB (`LANCEDB_PATH`) default paths now resolve under the checkout instead of the start directory. If you use either store at its default path and start DocsGPT from another directory, the old data is at diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py index 90ddeea0..e5ba9cb9 100644 --- a/docsgpt/deploy/docker.py +++ b/docsgpt/deploy/docker.py @@ -125,18 +125,28 @@ def wait_healthy( sleep: Callable[[float], None] = time.sleep, clock: Callable[[], float] = time.monotonic, ) -> bool: - """Poll ``url`` until it answers 2xx (True) or ``timeout`` seconds pass (False); tries at least once.""" + """Poll ``url`` until it answers 2xx (True) or ``timeout`` seconds pass (False); tries at least once. + + No request starts once the deadline is reached, and neither the pauses nor the + requests after the first run past it. + """ deadline = clock() + timeout + request_timeout = min(5.0, timeout) if timeout > 0 else 5.0 while True: try: - with opener(url, timeout=5) as response: + with opener(url, timeout=request_timeout) as response: if 200 <= response.status < 300: return True except (OSError, http.client.HTTPException): pass - if clock() >= deadline: + remaining = deadline - clock() + if remaining <= 0: return False - sleep(2) + sleep(min(2.0, remaining)) + remaining = deadline - clock() + if remaining <= 0: + return False + request_timeout = min(5.0, remaining) def lan_ip() -> str: diff --git a/docsgpt/deploy/envfile.py b/docsgpt/deploy/envfile.py index 9d51b946..fce83ae5 100644 --- a/docsgpt/deploy/envfile.py +++ b/docsgpt/deploy/envfile.py @@ -19,19 +19,21 @@ def _parse_value(raw: str) -> str: if len(value) >= 2 and value[0] == value[-1] == "'": return value[1:-1] if len(value) >= 2 and value[0] == value[-1] == '"': - return re.sub(r'\\(["\\])', r"\1", value[1:-1]) + return re.sub(r'\\(["\\])', r"\1", value[1:-1]).replace("$$", "$") return value.split(" #", 1)[0].rstrip() def _format_value(value: str) -> str: - """``value`` quoted so it reads back unchanged.""" + """``value`` quoted so it reads back unchanged, and so Compose passes it on unchanged.""" if "\n" in value or "\r" in value: raise ValueError("a .env value cannot contain a newline") if _PLAIN.match(value): return value if "'" not in value: + # Compose does not interpolate single-quoted values. return f"'{value}'" - return '"' + value.replace("\\", "\\\\").replace('"', '\\"') + '"' + # Double quotes are interpolated: $$ is Compose's literal dollar. + return '"' + value.replace("\\", "\\\\").replace('"', '\\"').replace("$", "$$") + '"' def read(path: Path) -> dict[str, str]: diff --git a/tests/deploy/test_docker.py b/tests/deploy/test_docker.py index d3de76b6..8e42556e 100644 --- a/tests/deploy/test_docker.py +++ b/tests/deploy/test_docker.py @@ -120,6 +120,38 @@ class TestWaitHealthy: sleep=lambda s: None, clock=lambda: next(clock)) assert attempts["n"] == 3 + def test_no_request_starts_at_or_after_the_deadline(self): + """A short timeout must not be overrun by a sleep and one more five-second request.""" + now = {"t": 0.0} + starts = [] + + def opener(url, timeout): + starts.append((now["t"], timeout)) + now["t"] += 1 + raise OSError("connection refused") + + def sleep(seconds): + now["t"] += seconds + + assert not docker_module.wait_healthy("http://127.0.0.1:7091/api/health", 3, opener=opener, + sleep=sleep, clock=lambda: now["t"]) + assert starts, "the first attempt always runs" + for started, timeout in starts[1:]: + assert started < 3 + assert started + timeout <= 3 + assert now["t"] <= 3 + + def test_a_zero_timeout_still_tries_once(self): + calls = [] + + def opener(url, timeout): + calls.append(timeout) + return _Response(200) + + assert docker_module.wait_healthy("http://127.0.0.1:7091/api/health", 0, opener=opener, + sleep=lambda s: None, clock=lambda: 0.0) + assert len(calls) == 1 and calls[0] > 0 + def test_gives_up_after_the_timeout(self): def opener(url, timeout): raise OSError("connection refused") diff --git a/tests/deploy/test_envfile.py b/tests/deploy/test_envfile.py index d2997d1b..0ccb3aae 100644 --- a/tests/deploy/test_envfile.py +++ b/tests/deploy/test_envfile.py @@ -55,12 +55,21 @@ class TestUpdate: envfile.update(path, {"A": "3"}) assert path.read_text() == "A=3\nX=y\n" - @pytest.mark.parametrize("value", ["plain", "with space", "hash # inside", "it's", 'say "hi"', "back\\slash", ""]) + @pytest.mark.parametrize( + "value", + ["plain", "with space", "hash # inside", "it's", 'say "hi"', "back\\slash", "", "pa$w'rd", "a$$b'c", "$HOME"], + ) def test_values_round_trip(self, tmp_path, value): path = tmp_path / ".env" envfile.update(path, {"VALUE": value}) assert envfile.read(path)["VALUE"] == value + def test_a_dollar_in_double_quotes_is_escaped_for_compose(self, tmp_path): + """Compose interpolates double-quoted .env values; $$ is its literal dollar.""" + path = tmp_path / ".env" + envfile.update(path, {"API_KEY": "pa$w'rd"}) + assert path.read_text() == "API_KEY=\"pa$$w'rd\"\n" + @pytest.mark.skipif(sys.platform == "win32", reason="POSIX permissions") def test_a_new_file_is_private(self, tmp_path): path = tmp_path / "stack" / ".env" From 303f239fcd2384e89b06df10e9cb47526ae268ca Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 00:11:38 +0100 Subject: [PATCH 012/130] fix: tighten an existing .env to 0600 before rewriting it; CI checks secrets have values envfile.update only applied 0600 when it created the file, so an existing .env with a wider mode kept it while secrets were written into it. The mode is now set on the open descriptor before the file is truncated and written. The CI step now requires POSTGRES_PASSWORD and JWT_SECRET_KEY to have values: an empty one falls back to a default without any check noticing. --- .github/workflows/docker-image-verify.yml | 5 +++-- docsgpt/deploy/envfile.py | 8 ++++++-- tests/deploy/test_envfile.py | 10 ++++++++++ 3 files changed, 19 insertions(+), 4 deletions(-) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index f59bbf9a..dad39d91 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -119,9 +119,10 @@ jobs: "$docsgpt" up --yes --dir "$stack" --image-tag verify "$docsgpt" status --dir "$stack" curl -fsS http://127.0.0.1:7091/ | grep -q 'src="/config.js"' - grep -q '^POSTGRES_PASSWORD=' "$stack/.env" + # The secrets must have values: an empty one would fall back to a default silently. + grep -Eq '^POSTGRES_PASSWORD=.+$' "$stack/.env" # Running up again keeps the generated secrets. - before=$(grep '^JWT_SECRET_KEY=' "$stack/.env") + before=$(grep -E '^JWT_SECRET_KEY=.+$' "$stack/.env") "$docsgpt" up --yes --dir "$stack" --image-tag verify [ "$(grep '^JWT_SECRET_KEY=' "$stack/.env")" = "$before" ] "$docsgpt" uninstall --yes --purge --dir "$stack" diff --git a/docsgpt/deploy/envfile.py b/docsgpt/deploy/envfile.py index fce83ae5..ff91e09b 100644 --- a/docsgpt/deploy/envfile.py +++ b/docsgpt/deploy/envfile.py @@ -52,7 +52,7 @@ def read(path: Path) -> dict[str, str]: def update(path: Path, values: Mapping[str, Optional[str]]) -> None: """Set each key in place (``None`` removes it), append new keys, and leave every other line alone. - A new file is created readable by its owner only: it holds secrets. + The file is left readable by its owner only, including one that existed with a wider mode: it holds secrets. """ formatted = {key: None if value is None else _format_value(value) for key, value in values.items()} path = Path(path) @@ -71,6 +71,10 @@ def update(path: Path, values: Mapping[str, Optional[str]]) -> None: out.extend(f"{key}={value}" for key, value in formatted.items() if key not in written and value is not None) path.parent.mkdir(parents=True, exist_ok=True) - descriptor = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600) + descriptor = os.open(path, os.O_WRONLY | os.O_CREAT, 0o600) + if hasattr(os, "fchmod"): + # The creation mode only applies to a new file; tighten an existing one before writing. + os.fchmod(descriptor, 0o600) + os.ftruncate(descriptor, 0) with os.fdopen(descriptor, "w", encoding="utf-8") as handle: handle.write("\n".join(out) + "\n" if out else "") diff --git a/tests/deploy/test_envfile.py b/tests/deploy/test_envfile.py index 0ccb3aae..e292be1c 100644 --- a/tests/deploy/test_envfile.py +++ b/tests/deploy/test_envfile.py @@ -76,6 +76,16 @@ class TestUpdate: envfile.update(path, {"JWT_SECRET_KEY": "s"}) assert oct(os.stat(path).st_mode & 0o777) == oct(0o600) + @pytest.mark.skipif(sys.platform == "win32", reason="POSIX permissions") + def test_an_existing_readable_file_is_made_private_before_writing(self, tmp_path): + """Secrets are rewritten into the file, so a permissive mode left from before must not stay.""" + path = tmp_path / ".env" + path.write_text("LLM_PROVIDER=openai\n") + os.chmod(path, 0o644) + envfile.update(path, {"JWT_SECRET_KEY": "s"}) + assert oct(os.stat(path).st_mode & 0o777) == oct(0o600) + assert envfile.read(path) == {"LLM_PROVIDER": "openai", "JWT_SECRET_KEY": "s"} + def test_rejects_a_newline_in_a_value(self, tmp_path): with pytest.raises(ValueError, match="newline"): envfile.update(tmp_path / ".env", {"A": "1\n2"}) From 9b5f1196fe836637cf02cd31fe363cc6ffa9207f Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 00:20:43 +0100 Subject: [PATCH 013/130] ci: the repeated startup must keep both generated secrets --- .github/workflows/docker-image-verify.yml | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index dad39d91..8e79b1ea 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -119,12 +119,12 @@ jobs: "$docsgpt" up --yes --dir "$stack" --image-tag verify "$docsgpt" status --dir "$stack" curl -fsS http://127.0.0.1:7091/ | grep -q 'src="/config.js"' - # The secrets must have values: an empty one would fall back to a default silently. - grep -Eq '^POSTGRES_PASSWORD=.+$' "$stack/.env" - # Running up again keeps the generated secrets. - before=$(grep -E '^JWT_SECRET_KEY=.+$' "$stack/.env") + # Both secrets must have values: an empty one falls back to a default silently. + secrets=$(grep -E '^(POSTGRES_PASSWORD|JWT_SECRET_KEY)=.+$' "$stack/.env" | sort) + [ "$(printf '%s\n' "$secrets" | wc -l)" -eq 2 ] + # Running up again keeps them; a changed database password locks the stack out of its volume. "$docsgpt" up --yes --dir "$stack" --image-tag verify - [ "$(grep '^JWT_SECRET_KEY=' "$stack/.env")" = "$before" ] + [ "$(grep -E '^(POSTGRES_PASSWORD|JWT_SECRET_KEY)=.+$' "$stack/.env" | sort)" = "$secrets" ] "$docsgpt" uninstall --yes --purge --dir "$stack" test ! -e "$stack" From db1e382502e9683f7d109eacbf12afe2b2841701 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 00:35:29 +0100 Subject: [PATCH 014/130] ci: each generated secret must appear exactly once --- .github/workflows/docker-image-verify.yml | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index 8e79b1ea..f592febf 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -119,11 +119,18 @@ jobs: "$docsgpt" up --yes --dir "$stack" --image-tag verify "$docsgpt" status --dir "$stack" curl -fsS http://127.0.0.1:7091/ | grep -q 'src="/config.js"' - # Both secrets must have values: an empty one falls back to a default silently. + # Each secret must appear exactly once with a value: a missing or empty one + # falls back to a default silently. + check_secrets() { + for key in POSTGRES_PASSWORD JWT_SECRET_KEY; do + [ "$(grep -Ec "^$key=.+$" "$stack/.env")" -eq 1 ] + done + } + check_secrets secrets=$(grep -E '^(POSTGRES_PASSWORD|JWT_SECRET_KEY)=.+$' "$stack/.env" | sort) - [ "$(printf '%s\n' "$secrets" | wc -l)" -eq 2 ] # Running up again keeps them; a changed database password locks the stack out of its volume. "$docsgpt" up --yes --dir "$stack" --image-tag verify + check_secrets [ "$(grep -E '^(POSTGRES_PASSWORD|JWT_SECRET_KEY)=.+$' "$stack/.env" | sort)" = "$secrets" ] "$docsgpt" uninstall --yes --purge --dir "$stack" test ! -e "$stack" From 63de66722ecf20255187aa11dd8a873cb113709f Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 00:53:03 +0100 Subject: [PATCH 015/130] fix: warn about plain HTTP before the stack starts, not only afterwards Network mode publishes the port on every interface and its access token travels as readable text, so `docsgpt up` says so before starting rather than in the summary at the end. The health poll's except clause says why it swallows the error. --- docsgpt/deploy/commands.py | 10 +++++++--- docsgpt/deploy/docker.py | 1 + tests/deploy/test_commands.py | 12 ++++++++++++ 3 files changed, 20 insertions(+), 3 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 19d4aae2..eaa38f44 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -219,6 +219,12 @@ def up(args, context: Optional[Context] = None) -> int: if stack.exposure(existing) == "domain" and stack.exposure(env) != "domain": # With the https profile off, `up --remove-orphans` would leave Caddy running on ports 80 and 443. context.docker.compose(directory, *_EVERY_PROFILE, "rm", "--stop", "--force", "caddy", check=False) + if stack.exposure(env) == "network": + print( + "DocsGPT will listen on every interface over plain HTTP: its access token travels as " + "readable text. Use `docsgpt up --domain ` for HTTPS outside a trusted network.", + file=sys.stderr, + ) up_args = ["up", "-d", "--remove-orphans", *recreate] if image_tag in _MOVING_TAGS: up_args += ["--pull", "always"] @@ -246,9 +252,7 @@ def up(args, context: Optional[Context] = None) -> int: if env.get("AUTH_TYPE") == "simple_jwt" and env.get("JWT_SECRET_KEY"): print("Access token (the page asks for it; `docsgpt token` prints it again):") print(f" {stack.simple_jwt_token(env['JWT_SECRET_KEY'])}") - if mode == "network": - print("Traffic is plain HTTP. Outside a trusted network, use a domain with HTTPS: docsgpt up --domain ") - elif mode == "domain": + if mode == "domain": print("Caddy gets the certificate when it starts: DNS must point at this machine and ports 80 and 443 be open.") print(f"Settings: {env_path} (change the model provider or access with `docsgpt up --reconfigure`)") print("Manage it with: docsgpt status | logs | upgrade | down | uninstall") diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py index e5ba9cb9..795aaeb9 100644 --- a/docsgpt/deploy/docker.py +++ b/docsgpt/deploy/docker.py @@ -138,6 +138,7 @@ def wait_healthy( if 200 <= response.status < 300: return True except (OSError, http.client.HTTPException): + # Not answering yet: refused, reset, timed out or a broken response. Keep polling. pass remaining = deadline - clock() if remaining <= 0: diff --git a/tests/deploy/test_commands.py b/tests/deploy/test_commands.py index 44398703..db9a9698 100644 --- a/tests/deploy/test_commands.py +++ b/tests/deploy/test_commands.py @@ -140,6 +140,18 @@ class TestUpFirstInstall: # A moving tag is pulled every time, not only when missing. assert (tmp_path, ["up", "-d", "--remove-orphans", "--pull", "always"]) in docker.calls + def test_network_mode_warns_about_plain_http_before_starting(self, tmp_path, capsys): + """The token travels as readable text, so say so before the stack is up, not only after.""" + docker = FakeDocker() + assert _run(["up", "--yes", "--dir", str(tmp_path), "--expose", "network"], _context(docker)) == 0 + err = capsys.readouterr().err + assert "plain HTTP" in err + assert "--domain" in err + + def test_a_local_install_does_not_warn(self, tmp_path, capsys): + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context()) == 0 + assert "plain HTTP" not in capsys.readouterr().err + def test_an_unhealthy_start_points_at_the_logs(self, tmp_path, capsys): assert _run(["up", "--yes", "--dir", str(tmp_path)], _context(healthy=False)) == 1 assert "docsgpt logs" in capsys.readouterr().err From 6b6bd1b0fb9d20a4ba12d07a679ad0801101ac00 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 01:10:09 +0100 Subject: [PATCH 016/130] test: assert the plain-HTTP warning is printed before compose starts --- tests/deploy/test_commands.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/tests/deploy/test_commands.py b/tests/deploy/test_commands.py index db9a9698..30faa350 100644 --- a/tests/deploy/test_commands.py +++ b/tests/deploy/test_commands.py @@ -143,10 +143,19 @@ class TestUpFirstInstall: def test_network_mode_warns_about_plain_http_before_starting(self, tmp_path, capsys): """The token travels as readable text, so say so before the stack is up, not only after.""" docker = FakeDocker() + start = docker.compose + seen = {} + + def recording(directory, *args, **kwargs): + # capsys hands over what was written so far, so the warning has to be in it already. + seen.setdefault("stderr", capsys.readouterr().err) + return start(directory, *args, **kwargs) + + docker.compose = recording assert _run(["up", "--yes", "--dir", str(tmp_path), "--expose", "network"], _context(docker)) == 0 - err = capsys.readouterr().err - assert "plain HTTP" in err - assert "--domain" in err + assert docker.calls, "compose was called" + assert "plain HTTP" in seen["stderr"] + assert "--domain" in seen["stderr"] def test_a_local_install_does_not_warn(self, tmp_path, capsys): assert _run(["up", "--yes", "--dir", str(tmp_path)], _context()) == 0 From 4f0bf2cca8a23f0bcdeeb53b60a62da8fabaaff8 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:17:58 +0100 Subject: [PATCH 017/130] feat: one-command installers for macOS, Linux and Windows deployment/install.sh (curl | bash) and install.ps1 (irm | iex) check for Docker, install uv when it is missing or older than 0.8 (pinned 0.12.15 via Astral's installer), install or upgrade the docsgpt package with `uv tool install`, and hand the terminal to `docsgpt up` with any arguments. On Linux without Docker the shell installer offers get.docker.com. Both run entirely inside a function, so a download cut short runs nothing. Releases attach both scripts next to the Compose file, which is where docs.ac/install and docs.ac/install.ps1 will point. installer-lint.yml runs shellcheck and the PowerShell parser; docker-image-verify.yml now installs through install.sh. README, Quickstart, Docker-Deploying and the changelog lead with the one-liner. --- .github/workflows/ci.yml | 11 +- .github/workflows/docker-image-verify.yml | 47 +++--- .github/workflows/installer-lint.yml | 47 ++++++ README.md | 28 +++- deployment/install.ps1 | 94 +++++++++++ deployment/install.sh | 169 ++++++++++++++++++++ docs/content/Deploying/Docker-Deploying.mdx | 14 +- docs/content/changelog.mdx | 7 + docs/content/quickstart.mdx | 139 ++++++++-------- 9 files changed, 455 insertions(+), 101 deletions(-) create mode 100644 .github/workflows/installer-lint.yml create mode 100644 deployment/install.ps1 create mode 100755 deployment/install.sh diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 5cd9c85e..fd042048 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -206,9 +206,16 @@ jobs: ref: ${{ inputs.version && format('refs/tags/{0}', inputs.version) || github.ref }} persist-credentials: false - - name: Attach the standalone compose file to the release + # The installers are served from these assets: docs.ac/install redirects to + # releases/latest/download/install.sh (and install.ps1), so the script a + # user runs always comes from the newest release. + - name: Attach the standalone compose file and the installers to the release env: GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} TAG: ${{ env.RELEASE_TAG }} run: | - gh release upload "$TAG" deployment/docker-compose-standalone.yaml --clobber + gh release upload "$TAG" \ + deployment/docker-compose-standalone.yaml \ + deployment/install.sh \ + deployment/install.ps1 \ + --clobber diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index f592febf..a2825d52 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -23,6 +23,7 @@ on: - 'frontend/**' - 'scripts/build_frontend.sh' - 'deployment/docker-compose-standalone.yaml' + - 'deployment/install.sh' - 'docsgpt/deploy/**' - 'docsgpt/cli.py' - 'docsgpt/core/paths.py' @@ -98,41 +99,29 @@ jobs: curl -fsS "$base/settings" | grep -q 'src="/config.js"' echo "API and UI served on $base" - - name: Set up uv - if: matrix.variant == '' - uses: astral-sh/setup-uv@d0cc045d04ccac9d8b7881df0226f9e82c39688e # v6.8.0 - with: - # No cache: a cache restored into a job that runs the built image is a poisoning risk. - enable-cache: false - - - name: docsgpt up runs the same image from the installed package + - name: The installer runs docsgpt up on the same image if: matrix.variant == '' + env: + DOCSGPT_NO_MODIFY_PATH: "1" run: | set -euo pipefail # Same Compose project name as the step above; stop that stack first. docker compose -f deployment/docker-compose-standalone.yaml down -v - uv build --wheel --out-dir "$RUNNER_TEMP/dist" - uv venv "$RUNNER_TEMP/venv" - uv pip install --python "$RUNNER_TEMP/venv/bin/python" "$RUNNER_TEMP"/dist/*.whl - docsgpt="$RUNNER_TEMP/venv/bin/docsgpt" - stack="$RUNNER_TEMP/stack" - "$docsgpt" up --yes --dir "$stack" --image-tag verify - "$docsgpt" status --dir "$stack" + pipx run build --wheel --outdir "$RUNNER_TEMP/dist" + export DOCSGPT_PACKAGE="$(ls "$RUNNER_TEMP"/dist/docsgpt-*.whl)" + # No uv is set up beforehand, so the installer's pinned uv download runs too. + # Without a terminal the installer passes --yes to docsgpt up. + bash deployment/install.sh --image-tag verify [!Note] -> Make sure you have [Docker](https://docs.docker.com/engine/install/) installed +> DocsGPT runs on [Docker](https://docs.docker.com/engine/install/). The installer checks for it first. -A more detailed [Quickstart](https://docs.docsgpt.cloud/quickstart) is available in our documentation +**macOS and Linux:** + +```bash +curl -fsSL https://docs.ac/install | bash +``` + +**Windows (PowerShell):** + +```powershell +irm https://docs.ac/install.ps1 | iex +``` + +The installer gets [uv](https://docs.astral.sh/uv/), installs the `docsgpt` Python package with it, and runs `docsgpt up`. That asks who should reach DocsGPT (only this computer, your network, or a domain with HTTPS) and which model provider to use, then starts it, at http://localhost:7091 for a local install. Afterwards, `docsgpt status`, `docsgpt logs`, `docsgpt upgrade`, `docsgpt down` and `docsgpt uninstall` manage it. + +To read the script before running it: + +```bash +curl -fsSL https://docs.ac/install -o install.sh +less install.sh +bash install.sh +``` + +A more detailed [Quickstart](https://docs.docsgpt.cloud/quickstart) is available in our documentation. + +### From a clone, with the setup script 1. **Clone the repository:** diff --git a/deployment/install.ps1 b/deployment/install.ps1 new file mode 100644 index 00000000..0c722b3b --- /dev/null +++ b/deployment/install.ps1 @@ -0,0 +1,94 @@ +# DocsGPT installer for Windows. +# +# irm https://docs.ac/install.ps1 | iex +# +# Installs uv when it is missing or too old, installs the docsgpt Python +# package with it, then runs `docsgpt up`, which sets up DocsGPT on Docker +# Desktop and starts it. Running it again upgrades the package and keeps your +# settings. To pass options to `docsgpt up`: +# +# & ([scriptblock]::Create((irm https://docs.ac/install.ps1))) --domain docs.example.com --yes +# +# Environment: +# DOCSGPT_VERSION package version to install (default: the latest release) +# DOCSGPT_PACKAGE install this instead of docsgpt from PyPI (a wheel path or URL) +# DOCSGPT_NO_MODIFY_PATH set to 1 to leave PATH alone +# +# Everything runs inside a function, so a download cut short runs nothing, and +# nothing calls `exit`, which would close the window `iex` runs in. + +function Install-DocsGPT { + param([string[]]$UpArguments) + + $ErrorActionPreference = 'Stop' + $UvVersion = '0.12.15' + $UvMinVersion = [version]'0.8.0' + + function Say([string]$Message) { Write-Host "==> $Message" } + + if (-not (Get-Command docker -ErrorAction SilentlyContinue)) { + throw 'DocsGPT runs on Docker. Install Docker Desktop (https://docs.docker.com/desktop/setup/install/windows-install/), start it, and run this again.' + } + + # uv installs and upgrades the package, and brings Python 3.12 when the system has none. + $uv = $null + $candidates = @( + (Get-Command uv -ErrorAction SilentlyContinue | Select-Object -ExpandProperty Source -First 1), + (Join-Path $HOME '.local\bin\uv.exe'), + (Join-Path $HOME '.cargo\bin\uv.exe') + ) | Where-Object { $_ -and (Test-Path $_) } + foreach ($candidate in $candidates) { + $found = "$(& $candidate --version 2>$null)" -replace '^uv\s+([0-9.]+).*$', '$1' + if ($found -match '^\d+\.\d+(\.\d+)?$' -and [version]$found -ge $UvMinVersion) { + $uv = $candidate + break + } + } + if (-not $uv) { + $uvDir = Join-Path $HOME '.local\bin' + Say "Installing uv $UvVersion into $uvDir" + $env:UV_INSTALL_DIR = $uvDir + $env:UV_NO_MODIFY_PATH = '1' + $env:UV_PRINT_QUIET = '1' + # A child PowerShell, so nothing the uv installer does can end this session. + $shell = (Get-Process -Id $PID).Path + & $shell -NoProfile -ExecutionPolicy Bypass -Command "irm https://astral.sh/uv/$UvVersion/install.ps1 | iex" + $uv = Join-Path $uvDir 'uv.exe' + if (-not (Test-Path $uv)) { throw "uv did not install into $uvDir" } + } + + if ($env:DOCSGPT_VERSION -and $env:DOCSGPT_PACKAGE) { + throw 'Set DOCSGPT_VERSION or DOCSGPT_PACKAGE, not both.' + } + if ($env:DOCSGPT_PACKAGE) { + Say "Installing docsgpt from $env:DOCSGPT_PACKAGE" + & $uv tool install --reinstall --python 3.12 $env:DOCSGPT_PACKAGE + } elseif ($env:DOCSGPT_VERSION) { + Say "Installing docsgpt $env:DOCSGPT_VERSION" + & $uv tool install --force --python 3.12 "docsgpt==$env:DOCSGPT_VERSION" + } else { + Say 'Installing the latest docsgpt' + & $uv tool install --upgrade --python 3.12 docsgpt + } + if ($LASTEXITCODE -ne 0) { throw 'Installing the docsgpt package failed.' } + + $binDir = "$(& $uv tool dir --bin)".Trim() + $docsgpt = Join-Path $binDir 'docsgpt.exe' + if (-not (Test-Path $docsgpt)) { throw "The docsgpt command is missing from $binDir." } + if (($env:Path -split ';') -notcontains $binDir) { + if ($env:DOCSGPT_NO_MODIFY_PATH -eq '1') { + Say "Add $binDir to PATH to run docsgpt from a new terminal" + } else { + & $uv tool update-shell *> $null + Say "Added $binDir to PATH for new terminals" + } + $env:Path = "$binDir;$env:Path" + } + + & $docsgpt up @UpArguments + if ($LASTEXITCODE -ne 0) { + Write-Error "docsgpt up exited with code $LASTEXITCODE. Run it again after fixing the problem above: docsgpt up" -ErrorAction Continue + } +} + +Install-DocsGPT -UpArguments $args diff --git a/deployment/install.sh b/deployment/install.sh new file mode 100755 index 00000000..95ed1401 --- /dev/null +++ b/deployment/install.sh @@ -0,0 +1,169 @@ +#!/usr/bin/env bash +# DocsGPT installer for macOS and Linux. +# +# curl -fsSL https://docs.ac/install | bash +# +# Installs uv when it is missing or too old, installs the `docsgpt` Python +# package with it, then runs `docsgpt up`, which sets up DocsGPT on Docker and +# starts it. Running it again upgrades the package and keeps your settings. +# Arguments go to `docsgpt up` (see `docsgpt up --help`): +# +# curl -fsSL https://docs.ac/install | bash -s -- --domain docs.example.com --yes +# +# Environment: +# DOCSGPT_VERSION package version to install (default: the latest release) +# DOCSGPT_PACKAGE install this instead of docsgpt from PyPI (a wheel path or URL) +# DOCSGPT_NO_MODIFY_PATH set to 1 to leave shell profiles alone +# DOCSGPT_INSTALL_DOCKER set to 1 to install Docker on Linux without asking +# +# Everything runs inside main(), so a download cut short runs nothing. + +UV_VERSION="0.12.15" +UV_MIN_VERSION="0.8.0" + +main() { + set -euo pipefail + + local bold="" red="" reset="" + if [ -t 2 ]; then + bold=$'\033[1m' red=$'\033[31m' reset=$'\033[0m' + fi + say() { printf '%s==>%s %s\n' "$bold" "$reset" "$*" >&2; } + die() { printf '%serror:%s %s\n' "$red" "$reset" "$*" >&2; exit 1; } + has() { command -v "$1" >/dev/null 2>&1; } + have_tty() { (exec /dev/null; } + ask_yes() { + local answer + printf '%s [y/N] ' "$1" >/dev/tty + read -r answer = B for dotted version numbers. + version_ge() { + local -a left right + IFS=. read -r -a left <<<"$1" + IFS=. read -r -a right <<<"$2" + local i x y + for i in 0 1 2; do + x="${left[i]:-0}" y="${right[i]:-0}" + x="${x%%[!0-9]*}" y="${y%%[!0-9]*}" + if (( 10#${x:-0} > 10#${y:-0} )); then return 0; fi + if (( 10#${x:-0} < 10#${y:-0} )); then return 1; fi + done + return 0 + } + + local os + os="$(uname -s)" + case "$os" in + Linux | Darwin) ;; + *) die "this installer is for macOS and Linux. On Windows, in PowerShell: irm https://docs.ac/install.ps1 | iex" ;; + esac + + # Docker first: without it nothing below is useful. + local docker_group_pending=0 + if ! has docker; then + if [ "$os" = Darwin ]; then + die "DocsGPT runs on Docker. Install Docker Desktop (https://docs.docker.com/desktop/setup/install/mac-install/) or OrbStack (https://orbstack.dev), start it, and run this again." + fi + if [ "${DOCSGPT_INSTALL_DOCKER:-}" = 1 ] || { have_tty && ask_yes "Docker is not installed. Install it now with Docker's script from get.docker.com?"; }; then + local sudo="" + if [ "$(id -u)" -ne 0 ]; then + has sudo || die "installing Docker needs root. Install it (https://docs.docker.com/engine/install/) and run this again." + sudo="sudo" + fi + say "Installing Docker" + download https://get.docker.com | $sudo sh + $sudo systemctl enable --now docker >/dev/null 2>&1 || true + if [ -n "$sudo" ]; then + $sudo usermod -aG docker "$(id -un)" + docker_group_pending=1 + fi + else + die "DocsGPT runs on Docker. Install it (https://docs.docker.com/engine/install/) and run this again." + fi + fi + + # uv installs and upgrades the package, and brings Python 3.12 when the system has none. + local uv="" candidate found + for candidate in "$(command -v uv 2>/dev/null || true)" "$HOME/.local/bin/uv" "$HOME/.cargo/bin/uv"; do + [ -n "$candidate" ] && [ -x "$candidate" ] || continue + found="$("$candidate" --version 2>/dev/null | awk '{print $2}')" || continue + if [ -n "$found" ] && version_ge "$found" "$UV_MIN_VERSION"; then + uv="$candidate" + break + fi + done + if [ -z "$uv" ]; then + local uv_dir="${XDG_BIN_HOME:-$HOME/.local/bin}" + say "Installing uv $UV_VERSION into $uv_dir" + download "https://astral.sh/uv/$UV_VERSION/install.sh" | env UV_INSTALL_DIR="$uv_dir" UV_NO_MODIFY_PATH=1 UV_PRINT_QUIET=1 sh + uv="$uv_dir/uv" + [ -x "$uv" ] || die "uv did not install into $uv_dir" + fi + + if [ -n "${DOCSGPT_VERSION:-}" ] && [ -n "${DOCSGPT_PACKAGE:-}" ]; then + die "set DOCSGPT_VERSION or DOCSGPT_PACKAGE, not both" + fi + if [ -n "${DOCSGPT_PACKAGE:-}" ]; then + say "Installing docsgpt from $DOCSGPT_PACKAGE" + "$uv" tool install --reinstall --python 3.12 "$DOCSGPT_PACKAGE" + elif [ -n "${DOCSGPT_VERSION:-}" ]; then + say "Installing docsgpt $DOCSGPT_VERSION" + "$uv" tool install --force --python 3.12 "docsgpt==$DOCSGPT_VERSION" + else + say "Installing the latest docsgpt" + "$uv" tool install --upgrade --python 3.12 docsgpt + fi + + local bin_dir docsgpt + bin_dir="$("$uv" tool dir --bin)" + docsgpt="$bin_dir/docsgpt" + [ -x "$docsgpt" ] || die "the docsgpt command is missing from $bin_dir" + case ":$PATH:" in + *":$bin_dir:"*) ;; + *) + if [ "${DOCSGPT_NO_MODIFY_PATH:-}" = 1 ]; then + say "Add $bin_dir to PATH to run docsgpt from a new terminal" + else + "$uv" tool update-shell >/dev/null 2>&1 || true + say "Added $bin_dir to PATH for new terminals" + fi + ;; + esac + + if [ "$docker_group_pending" = 1 ]; then + if has sg; then + # The docker group applies to new logins; sg gives it to this command now. + local command + command="$(printf '%q ' "$docsgpt" up "$@")" + if have_tty; then + exec sg docker -c "$command + To read the script before running it, download it first: `curl -fsSL https://docs.ac/install -o install.sh`, then `bash install.sh`. On Windows: `irm https://docs.ac/install.ps1 -OutFile install.ps1`, then `.\install.ps1`. + + +**Options.** Arguments after `bash -s --` go to `docsgpt up`, so a server can be set up without questions: + +```bash +curl -fsSL https://docs.ac/install | bash -s -- --yes --domain docs.example.com --provider openai --api-key "$OPENAI_API_KEY" +``` + +`DOCSGPT_VERSION` installs a specific release, and `DOCSGPT_NO_MODIFY_PATH=1` leaves your shell profile alone. `docsgpt up --help` lists every option. + +**Afterwards:** + +| Command | What it does | +| --- | --- | +| `docsgpt status` | Version, address, and whether DocsGPT answers | +| `docsgpt logs -f` | Follow the logs | +| `docsgpt token` | The access token, for installs reachable beyond this computer | +| `docsgpt up --reconfigure` | Ask the setup questions again | +| `docsgpt upgrade` | Upgrade to the latest release, keeping your settings and data | +| `docsgpt down` | Stop DocsGPT | +| `docsgpt uninstall` | Remove it; `--purge` also deletes settings and data | + +Running the install command again also upgrades. See [Run it with `docsgpt up`](/Deploying/Docker-Deploying#run-it-with-docsgpt-up) for the details. + +## From a clone, with the setup script + +To work from the source tree, for example to build the images yourself, use `setup.sh` (macOS and Linux) or `setup.ps1` (Windows). + +1. **Clone the repository:** ```bash git clone https://github.com/arc53/DocsGPT.git cd DocsGPT ``` -2. **Run the `setup.sh` script:** - - Navigate to the DocsGPT directory in your terminal and execute the `setup.sh` script: +2. **Run the setup script:** ```bash ./setup.sh ``` -3. **Follow the interactive setup:** + On Windows: - The `setup.sh` script will guide you through an interactive menu with the following options: + ```powershell + PowerShell -ExecutionPolicy Bypass -File .\setup.ps1 + ``` + +3. **Follow the interactive setup:** ``` Welcome to DocsGPT Setup! @@ -47,73 +98,29 @@ The easiest way to launch DocsGPT is using the provided `setup.sh` script. This Choose option (1-5): ``` - Let's break down each option: + * **1) Use DocsGPT Public API Endpoint (simple and free):** This is the simplest option to get started. It utilizes the DocsGPT public API, requiring no API keys or local model downloads. - * **1) Use DocsGPT Public API Endpoint (simple and free):** This is the simplest option to get started. It utilizes the DocsGPT public API, requiring no API keys or local model downloads. Choose this for a quick and easy setup. + * **2) Serve Local (with Ollama):** Runs a Large Language Model locally using [Ollama](https://ollama.com/). You'll be prompted to choose between CPU or GPU for Ollama and select a model to download. - * **2) Serve Local (with Ollama):** This option allows you to run a Large Language Model locally using [Ollama](https://ollama.com/). You'll be prompted to choose between CPU or GPU for Ollama and select a model to download. This is a good option for local processing and experimentation. + * **3) Connect Local Inference Engine:** If you already run a local inference engine like Llama.cpp, Text Generation Inference (TGI), vLLM, or others, choose this option and provide the connection details. - * **3) Connect Local Inference Engine:** If you are already running a local inference engine like Llama.cpp, Text Generation Inference (TGI), vLLM, or others, choose this option. You'll be asked to select your engine and provide the necessary connection details. This is for users with existing local LLM infrastructure. + * **4) Connect Cloud API Provider:** Connect DocsGPT to a Cloud API provider such as OpenAI, Google (Vertex AI/Gemini), Anthropic (Claude), Groq, HuggingFace Inference API, or Azure OpenAI. You will need an API key from your chosen provider. - * **4) Connect Cloud API Provider:** This option lets you connect DocsGPT to a commercial Cloud API provider such as OpenAI, Google (Vertex AI/Gemini), Anthropic (Claude), Groq, HuggingFace Inference API, or Azure OpenAI. You will need an API key from your chosen provider. Select this if you prefer to use a powerful cloud-based LLM. + * **5) Modify DocsGPT's source code and rebuild the Docker images locally.** Instead of pulling prebuilt images from Docker Hub, you build the backend and frontend from source, to customize how DocsGPT works internally or to run it in an environment without internet access. - * **5) Modify DocsGPT's source code and rebuild the Docker images locally.** Instead of pulling prebuilt images from Docker Hub or using the hosted/public API, you build the entire backend and frontend from source, customizing how DocsGPT works internally, or run it in an environment without internet access. + After selecting an option and providing any required information (like API keys or model names), the script configures your `.env` file and starts DocsGPT using Docker Compose. - After selecting an option and providing any required information (like API keys or model names), the script will configure your `.env` file and start DocsGPT using Docker Compose. +4. **Access DocsGPT in your browser:** open [http://localhost:5173/](http://localhost:5173/). -4. **Access DocsGPT in your browser:** - - Once the setup is complete and Docker containers are running, navigate to [http://localhost:5173/](http://localhost:5173/) in your web browser to access the DocsGPT web application. - -5. **Stopping DocsGPT:** - - To stop DocsGPT, simply open a new terminal in the `DocsGPT` directory and run: +5. **Stopping DocsGPT:** in the `DocsGPT` directory, run the `docker compose down` command the script printed at the end, for example: ```bash docker compose -f deployment/docker-compose-hub.yaml down ``` - (or the specific `docker compose` command shown at the end of the `setup.sh` execution, which may include optional compose files depending on your choices). -## Launching DocsGPT (Windows) +**Important for Windows:** Ensure Docker Desktop is installed and running before you start. The script tries to start Docker if it is not running, but you may need to start it manually. -For Windows users, we provide a PowerShell script that offers the same functionality as the macOS/Linux setup script. - -**Steps:** - -1. **Download the DocsGPT Repository:** - - First, you need to download the DocsGPT repository to your local machine. You can do this using Git: - - ```powershell - git clone https://github.com/arc53/DocsGPT.git - cd DocsGPT - ``` - -2. **Run the `setup.ps1` script:** - - Execute the PowerShell setup script: - - ```powershell - PowerShell -ExecutionPolicy Bypass -File .\setup.ps1 - ``` - -3. **Follow the interactive setup:** - - Just like the Linux/macOS script, the PowerShell script will guide you through setting DocsGPT. - The script will handle environment configuration and start DocsGPT based on your selections. - -4. **Access DocsGPT in your browser:** - - Once the setup is complete and Docker containers are running, navigate to [http://localhost:5173/](http://localhost:5173/) in your web browser to access the DocsGPT web application. - -5. **Stopping DocsGPT:** - - To stop DocsGPT run the Docker Compose down command displayed at the end of the setup script's execution. - -**Important for Windows:** Ensure Docker Desktop is installed and running correctly on your Windows system before proceeding. The script will attempt to start Docker if it's not running, but you may need to start it manually if there are issues. - -**Alternative Method:** -If you prefer a more manual approach, you can follow our [Docker Deployment documentation](/Deploying/Docker-Deploying) for detailed instructions on setting up DocsGPT on Windows using Docker commands directly. +**Alternative Method:** To run the pre-built images with Docker Compose yourself, follow the [Docker Deployment documentation](/Deploying/Docker-Deploying). ## Advanced Configuration From bd352812593a9388a0f798cd5c5832c34dd2dcfc Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:20:02 +0100 Subject: [PATCH 018/130] fix: plain if in the installer's uv lookup (shellcheck SC2015) --- deployment/install.sh | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/deployment/install.sh b/deployment/install.sh index 95ed1401..509a04a2 100755 --- a/deployment/install.sh +++ b/deployment/install.sh @@ -96,7 +96,9 @@ main() { # uv installs and upgrades the package, and brings Python 3.12 when the system has none. local uv="" candidate found for candidate in "$(command -v uv 2>/dev/null || true)" "$HOME/.local/bin/uv" "$HOME/.cargo/bin/uv"; do - [ -n "$candidate" ] && [ -x "$candidate" ] || continue + if [ -z "$candidate" ] || [ ! -x "$candidate" ]; then + continue + fi found="$("$candidate" --version 2>/dev/null | awk '{print $2}')" || continue if [ -n "$found" ] && version_ge "$found" "$UV_MIN_VERSION"; then uv="$candidate" From 49823a685992c064b4a77c9272418004d96bae04 Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 15 Sep 2026 23:29:19 +0100 Subject: [PATCH 019/130] fix: installer review follow-ups - install.sh saves the get.docker.com and uv installers to a file and runs them only after the download finished, so a cut-off transfer runs nothing. - Neither installer prints DOCSGPT_PACKAGE, which may be a URL with credentials. - The CI step assigns the wheel path before exporting it, so a missing wheel fails instead of installing from PyPI. - Docker-Deploying shows one code block per platform; Quickstart names the /opt/docsgpt home used for root on Linux. --- .github/workflows/docker-image-verify.yml | 4 +++- deployment/install.ps1 | 2 +- deployment/install.sh | 20 +++++++++++++++++--- docs/content/Deploying/Docker-Deploying.mdx | 11 +++++++++-- docs/content/quickstart.mdx | 2 +- 5 files changed, 31 insertions(+), 8 deletions(-) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index a2825d52..9a9ad057 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -108,7 +108,9 @@ jobs: # Same Compose project name as the step above; stop that stack first. docker compose -f deployment/docker-compose-standalone.yaml down -v pipx run build --wheel --outdir "$RUNNER_TEMP/dist" - export DOCSGPT_PACKAGE="$(ls "$RUNNER_TEMP"/dist/docsgpt-*.whl)" + # Assigned before export, so a missing wheel fails here instead of installing from PyPI. + DOCSGPT_PACKAGE="$(ls "$RUNNER_TEMP"/dist/docsgpt-*.whl)" + export DOCSGPT_PACKAGE # No uv is set up beforehand, so the installer's pinned uv download runs too. # Without a terminal the installer passes --yes to docsgpt up. bash deployment/install.sh --image-tag verify "$script"; then + rm -f "$script" + die "could not download $url" + fi + "$@" "$script" || status=$? + rm -f "$script" + return "$status" + } # version_ge A B: A >= B for dotted version numbers. version_ge() { local -a left right @@ -82,7 +96,7 @@ main() { sudo="sudo" fi say "Installing Docker" - download https://get.docker.com | $sudo sh + run_downloaded https://get.docker.com $sudo sh $sudo systemctl enable --now docker >/dev/null 2>&1 || true if [ -n "$sudo" ]; then $sudo usermod -aG docker "$(id -un)" @@ -108,7 +122,7 @@ main() { if [ -z "$uv" ]; then local uv_dir="${XDG_BIN_HOME:-$HOME/.local/bin}" say "Installing uv $UV_VERSION into $uv_dir" - download "https://astral.sh/uv/$UV_VERSION/install.sh" | env UV_INSTALL_DIR="$uv_dir" UV_NO_MODIFY_PATH=1 UV_PRINT_QUIET=1 sh + run_downloaded "https://astral.sh/uv/$UV_VERSION/install.sh" env UV_INSTALL_DIR="$uv_dir" UV_NO_MODIFY_PATH=1 UV_PRINT_QUIET=1 sh uv="$uv_dir/uv" [ -x "$uv" ] || die "uv did not install into $uv_dir" fi @@ -117,7 +131,7 @@ main() { die "set DOCSGPT_VERSION or DOCSGPT_PACKAGE, not both" fi if [ -n "${DOCSGPT_PACKAGE:-}" ]; then - say "Installing docsgpt from $DOCSGPT_PACKAGE" + say "Installing docsgpt from DOCSGPT_PACKAGE" "$uv" tool install --reinstall --python 3.12 "$DOCSGPT_PACKAGE" elif [ -n "${DOCSGPT_VERSION:-}" ]; then say "Installing docsgpt $DOCSGPT_VERSION" diff --git a/docs/content/Deploying/Docker-Deploying.mdx b/docs/content/Deploying/Docker-Deploying.mdx index 400c0207..b4c5f87f 100644 --- a/docs/content/Deploying/Docker-Deploying.mdx +++ b/docs/content/Deploying/Docker-Deploying.mdx @@ -24,9 +24,16 @@ you. It needs Docker with Compose 2.24 or newer. The installer gets [uv](https://docs.astral.sh/uv/), installs the package with it and runs `docsgpt up`: +macOS and Linux: + ```bash -curl -fsSL https://docs.ac/install | bash # macOS and Linux -irm https://docs.ac/install.ps1 | iex # Windows (PowerShell) +curl -fsSL https://docs.ac/install | bash +``` + +Windows (PowerShell): + +```powershell +irm https://docs.ac/install.ps1 | iex ``` Both scripts are attached to every [release](https://github.com/arc53/DocsGPT/releases) diff --git a/docs/content/quickstart.mdx b/docs/content/quickstart.mdx index 63917cfb..3b134fe6 100644 --- a/docs/content/quickstart.mdx +++ b/docs/content/quickstart.mdx @@ -34,7 +34,7 @@ The installer: * **Who should reach DocsGPT:** only this computer; other machines on your network (plain HTTP, with an access token); or a domain name with HTTPS (a certificate from Let's Encrypt, with an access token). * **Which model provider:** the DocsGPT public API (no key needed), OpenAI, Anthropic, Google Gemini, OpenRouter, Groq, or an OpenAI-compatible server such as Ollama or vLLM. -It then starts DocsGPT and prints its address, [http://localhost:7091](http://localhost:7091) for a local install. Settings and generated secrets are in `~/.docsgpt/server/.env`. +It then starts DocsGPT and prints its address, [http://localhost:7091](http://localhost:7091) for a local install. Settings and generated secrets are in `~/.docsgpt/server/.env`, or `/opt/docsgpt/.env` when the installer runs as root on Linux; `docsgpt up --dir` or `DOCSGPT_HOME` choose another folder. To read the script before running it, download it first: `curl -fsSL https://docs.ac/install -o install.sh`, then `bash install.sh`. On Windows: `irm https://docs.ac/install.ps1 -OutFile install.ps1`, then `.\install.ps1`. From 6afce45913c0dd459d53094bcbc1535f6e2d3b48 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 00:11:42 +0100 Subject: [PATCH 020/130] ci: the installer check requires non-empty secrets --- .github/workflows/docker-image-verify.yml | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index 9a9ad057..fb8d3bb2 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -118,9 +118,10 @@ jobs: stack="$HOME/.docsgpt/server" "$docsgpt" status curl -fsS http://127.0.0.1:7091/ | grep -q 'src="/config.js"' - grep -q '^POSTGRES_PASSWORD=' "$stack/.env" + # The secrets must have values: an empty one would fall back to a default silently. + grep -Eq '^POSTGRES_PASSWORD=.+$' "$stack/.env" # Running the installer again upgrades in place and keeps the generated secrets. - before=$(grep '^JWT_SECRET_KEY=' "$stack/.env") + before=$(grep -E '^JWT_SECRET_KEY=.+$' "$stack/.env") bash deployment/install.sh --image-tag verify Date: Wed, 16 Sep 2026 00:27:42 +0100 Subject: [PATCH 022/130] fix: the Windows installer fails when docsgpt up fails --- deployment/install.ps1 | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/deployment/install.ps1 b/deployment/install.ps1 index 2eb2b71f..e8ea9a96 100644 --- a/deployment/install.ps1 +++ b/deployment/install.ps1 @@ -87,7 +87,8 @@ function Install-DocsGPT { & $docsgpt up @UpArguments if ($LASTEXITCODE -ne 0) { - Write-Error "docsgpt up exited with code $LASTEXITCODE. Run it again after fixing the problem above: docsgpt up" -ErrorAction Continue + # throw, not Write-Error: the caller (and any automation) must see this fail. + throw "docsgpt up exited with code $LASTEXITCODE. Fix the problem above and run it again: docsgpt up" } } From 0e1963552ce9bec19ebc06555c440ccc2dee79fc Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 00:35:35 +0100 Subject: [PATCH 023/130] ci: each generated secret must appear exactly once --- .github/workflows/docker-image-verify.yml | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/.github/workflows/docker-image-verify.yml b/.github/workflows/docker-image-verify.yml index 468c0dc8..552764c0 100644 --- a/.github/workflows/docker-image-verify.yml +++ b/.github/workflows/docker-image-verify.yml @@ -118,11 +118,18 @@ jobs: stack="$HOME/.docsgpt/server" "$docsgpt" status curl -fsS http://127.0.0.1:7091/ | grep -q 'src="/config.js"' - # Both secrets must have values: an empty one falls back to a default silently. + # Each secret must appear exactly once with a value: a missing or empty one + # falls back to a default silently. + check_secrets() { + for key in POSTGRES_PASSWORD JWT_SECRET_KEY; do + [ "$(grep -Ec "^$key=.+$" "$stack/.env")" -eq 1 ] + done + } + check_secrets secrets=$(grep -E '^(POSTGRES_PASSWORD|JWT_SECRET_KEY)=.+$' "$stack/.env" | sort) - [ "$(printf '%s\n' "$secrets" | wc -l)" -eq 2 ] # Running the installer again upgrades in place and keeps both secrets. bash deployment/install.sh --image-tag verify Date: Wed, 16 Sep 2026 00:44:06 +0100 Subject: [PATCH 024/130] fix: check the uv installer against a pinned sha256 before running it Both installers download the pinned uv installer to a file and run it only when its sha256 matches the value pinned next to UV_VERSION; bumping the version means bumping the hash. Astral publishes checksums for the uv binaries but not for the installer scripts, so the hash is pinned here. get.docker.com is still only downloaded in full before running: its content changes over time and it publishes no checksum. --- deployment/install.ps1 | 20 +++++++++++++++++--- deployment/install.sh | 35 +++++++++++++++++++++++++++++------ 2 files changed, 46 insertions(+), 9 deletions(-) diff --git a/deployment/install.ps1 b/deployment/install.ps1 index e8ea9a96..f62edc70 100644 --- a/deployment/install.ps1 +++ b/deployment/install.ps1 @@ -23,6 +23,9 @@ function Install-DocsGPT { $ErrorActionPreference = 'Stop' $UvVersion = '0.12.15' $UvMinVersion = [version]'0.8.0' + # sha256 of https://astral.sh/uv/$UvVersion/install.ps1, checked before it runs. Bump it with + # $UvVersion: (Invoke-WebRequest "https://astral.sh/uv//install.ps1").Content | … + $UvInstallerSha256 = '63f2d7e2ccc347cc018127b22f56f569080e5730a390dcb702b654daf6822193' function Say([string]$Message) { Write-Host "==> $Message" } @@ -50,9 +53,20 @@ function Install-DocsGPT { $env:UV_INSTALL_DIR = $uvDir $env:UV_NO_MODIFY_PATH = '1' $env:UV_PRINT_QUIET = '1' - # A child PowerShell, so nothing the uv installer does can end this session. - $shell = (Get-Process -Id $PID).Path - & $shell -NoProfile -ExecutionPolicy Bypass -Command "irm https://astral.sh/uv/$UvVersion/install.ps1 | iex" + # Downloaded and checked before it runs, and run in a child PowerShell so nothing the + # uv installer does can end this session. + $installer = Join-Path ([System.IO.Path]::GetTempPath()) "uv-installer-$UvVersion.ps1" + try { + Invoke-WebRequest -Uri "https://astral.sh/uv/$UvVersion/install.ps1" -OutFile $installer -UseBasicParsing + $actual = (Get-FileHash -Path $installer -Algorithm SHA256).Hash.ToLower() + if ($actual -ne $UvInstallerSha256) { + throw "The uv installer does not match its pinned sha256: got $actual, expected $UvInstallerSha256. Refusing to run it." + } + $shell = (Get-Process -Id $PID).Path + & $shell -NoProfile -ExecutionPolicy Bypass -File $installer + } finally { + Remove-Item $installer -Force -ErrorAction SilentlyContinue + } $uv = Join-Path $uvDir 'uv.exe' if (-not (Test-Path $uv)) { throw "uv did not install into $uvDir" } } diff --git a/deployment/install.sh b/deployment/install.sh index 018bb761..c7b81dee 100755 --- a/deployment/install.sh +++ b/deployment/install.sh @@ -20,6 +20,9 @@ UV_VERSION="0.12.15" UV_MIN_VERSION="0.8.0" +# sha256 of https://astral.sh/uv/$UV_VERSION/install.sh, checked before it runs. Bump it with +# UV_VERSION: curl -fsSL https://astral.sh/uv//install.sh | shasum -a 256 +UV_INSTALLER_SHA256="716a1d6844740756c68770fcec2f79c2013fb9b03869a113f61e15f6f482a6a1" main() { set -euo pipefail @@ -47,16 +50,33 @@ main() { die "curl or wget is needed to download $1" fi } - # run_downloaded URL COMMAND...: save URL to a file, then run COMMAND with the file as its last - # argument. A transfer cut short fails before anything runs. + sha256_of() { + if has shasum; then + shasum -a 256 "$1" | awk '{print $1}' + elif has sha256sum; then + sha256sum "$1" | awk '{print $1}' + else + die "neither shasum nor sha256sum is available to check $1" + fi + } + # run_downloaded URL SHA256 COMMAND...: save URL to a file, check it against SHA256 ("-" to + # skip), then run COMMAND with the file as its last argument. A transfer cut short, or content + # that does not match, fails before anything runs. run_downloaded() { - local url="$1" script status=0 - shift + local url="$1" expected="$2" script status=0 actual + shift 2 script="$(mktemp)" if ! download "$url" >"$script"; then rm -f "$script" die "could not download $url" fi + if [ "$expected" != "-" ]; then + actual="$(sha256_of "$script")" + if [ "$actual" != "$expected" ]; then + rm -f "$script" + die "$url does not match its pinned sha256: got $actual, expected $expected. Refusing to run it." + fi + fi "$@" "$script" || status=$? rm -f "$script" return "$status" @@ -96,7 +116,9 @@ main() { sudo="sudo" fi say "Installing Docker" - run_downloaded https://get.docker.com $sudo sh + # Docker's script changes over time and publishes no checksum, so it is only downloaded + # in full before it runs. + run_downloaded https://get.docker.com - $sudo sh $sudo systemctl enable --now docker >/dev/null 2>&1 || true if [ -n "$sudo" ]; then $sudo usermod -aG docker "$(id -un)" @@ -122,7 +144,8 @@ main() { if [ -z "$uv" ]; then local uv_dir="${XDG_BIN_HOME:-$HOME/.local/bin}" say "Installing uv $UV_VERSION into $uv_dir" - run_downloaded "https://astral.sh/uv/$UV_VERSION/install.sh" env UV_INSTALL_DIR="$uv_dir" UV_NO_MODIFY_PATH=1 UV_PRINT_QUIET=1 sh + run_downloaded "https://astral.sh/uv/$UV_VERSION/install.sh" "$UV_INSTALLER_SHA256" \ + env UV_INSTALL_DIR="$uv_dir" UV_NO_MODIFY_PATH=1 UV_PRINT_QUIET=1 sh uv="$uv_dir/uv" [ -x "$uv" ] || die "uv did not install into $uv_dir" fi From d993aaced0cc11d7fa3fb1428bf65a7536c45332 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 00:57:43 +0100 Subject: [PATCH 025/130] fix: fail on a nonzero uv installer exit; POSIX quoting for the sg handoff The Windows installer only checked that uv.exe exists after running the uv installer, so a failed install that left an older uv.exe behind was accepted; it now fails on a nonzero exit code. sg runs its command with /bin/sh, which need not be bash, so the handoff after installing Docker quotes each argument as POSIX single quotes instead of with bash's printf %q. --- deployment/install.ps1 | 4 ++++ deployment/install.sh | 13 +++++++++++-- 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/deployment/install.ps1 b/deployment/install.ps1 index f62edc70..5b900eb3 100644 --- a/deployment/install.ps1 +++ b/deployment/install.ps1 @@ -64,6 +64,10 @@ function Install-DocsGPT { } $shell = (Get-Process -Id $PID).Path & $shell -NoProfile -ExecutionPolicy Bypass -File $installer + if ($LASTEXITCODE -ne 0) { + # Without this, a uv.exe left behind by an older install would be accepted below. + throw "The uv installer exited with code $LASTEXITCODE." + } } finally { Remove-Item $installer -Force -ErrorAction SilentlyContinue } diff --git a/deployment/install.sh b/deployment/install.sh index c7b81dee..2eac7ee6 100755 --- a/deployment/install.sh +++ b/deployment/install.sh @@ -50,6 +50,14 @@ main() { die "curl or wget is needed to download $1" fi } + # POSIX single-quote encoding: what needs quoting is decided here, not by the shell that runs it. + shell_quote() { + local arg out="" + for arg in "$@"; do + out="$out'$(printf '%s' "$arg" | sed "s/'/'\\\\''/g")' " + done + printf '%s' "$out" + } sha256_of() { if has shasum; then shasum -a 256 "$1" | awk '{print $1}' @@ -182,9 +190,10 @@ main() { if [ "$docker_group_pending" = 1 ]; then if has sg; then - # The docker group applies to new logins; sg gives it to this command now. + # The docker group applies to new logins; sg gives it to this command now. sg runs the + # command with /bin/sh, which need not be bash, so quote for POSIX sh rather than with %q. local command - command="$(printf '%q ' "$docsgpt" up "$@")" + command="$(shell_quote "$docsgpt" up "$@")" if have_tty; then exec sg docker -c "$command Date: Wed, 16 Sep 2026 14:23:26 +0100 Subject: [PATCH 026/130] chore: 0.21.0 --- docs/content/changelog.mdx | 2 +- docsgpt/version.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/content/changelog.mdx b/docs/content/changelog.mdx index 13cdc752..6e59a1f5 100644 --- a/docs/content/changelog.mdx +++ b/docs/content/changelog.mdx @@ -11,7 +11,7 @@ The notable changes in each release. Every release on GitHub also carries [auto-generated notes](https://github.com/arc53/DocsGPT/releases) listing every merged pull request, and [Upgrading](/upgrading) covers the steps an existing deployment has to take. -## Unreleased +## 0.21.0 ### Install with one command diff --git a/docsgpt/version.py b/docsgpt/version.py index c6e10fde..c327b39b 100644 --- a/docsgpt/version.py +++ b/docsgpt/version.py @@ -2,7 +2,7 @@ from __future__ import annotations -__version__ = "0.20.0" +__version__ = "0.21.0" def get_version() -> str: From ce34d473d3285648bc2c1661053809032e18bbcb Mon Sep 17 00:00:00 2001 From: ManishMadan2882 Date: Wed, 16 Sep 2026 19:33:47 +0530 Subject: [PATCH 027/130] fix(frontend): use icons that match their actions and read in dark mode --- .../src/agents/workflow/WorkflowBuilder.tsx | 2 +- frontend/src/components/ActionButtons.tsx | 8 ++----- .../message-input/AttachFileButton.tsx | 2 +- .../message-input/SourcesTrigger.tsx | 2 +- .../components/message-input/ToolsTrigger.tsx | 2 +- .../src/conversation/ConversationTile.tsx | 23 +++++++++++-------- frontend/src/locale/de.json | 1 + frontend/src/locale/en.json | 1 + frontend/src/locale/es.json | 1 + frontend/src/locale/jp.json | 1 + frontend/src/locale/ru.json | 1 + frontend/src/locale/zh-TW.json | 1 + frontend/src/locale/zh.json | 1 + frontend/src/modals/MoveToFolderModal.tsx | 2 +- frontend/src/settings/Logs.tsx | 2 +- frontend/src/settings/ToolConfig.tsx | 4 ++-- .../settings/components/RetrievalOptions.tsx | 2 +- frontend/src/teams/TeamSwitcher.tsx | 6 ++--- frontend/src/upload/Upload.tsx | 2 +- 19 files changed, 36 insertions(+), 28 deletions(-) diff --git a/frontend/src/agents/workflow/WorkflowBuilder.tsx b/frontend/src/agents/workflow/WorkflowBuilder.tsx index 6d58d662..d71867f6 100644 --- a/frontend/src/agents/workflow/WorkflowBuilder.tsx +++ b/frontend/src/agents/workflow/WorkflowBuilder.tsx @@ -1587,7 +1587,7 @@ function WorkflowBuilderInner() { size="icon-sm" onClick={() => setShowWorkflowSettings(!showWorkflowSettings)} className="text-muted-foreground hover:bg-accent hover:text-foreground size-auto p-1" - aria-label="Workflow settings" + aria-label="Edit workflow details" title={ workflowDescription ? `${workflowName || 'New Workflow'} — ${workflowDescription}` diff --git a/frontend/src/components/ActionButtons.tsx b/frontend/src/components/ActionButtons.tsx index 8520315f..cc161017 100644 --- a/frontend/src/components/ActionButtons.tsx +++ b/frontend/src/components/ActionButtons.tsx @@ -1,4 +1,4 @@ -import { ExternalLink, Plus } from 'lucide-react'; +import { Plus, Share } from 'lucide-react'; import { useTranslation } from 'react-i18next'; import { useSelector } from 'react-redux'; import { ShareConversationModal } from '../modals/ShareConversationModal'; @@ -73,11 +73,7 @@ export default function ActionButtons({ onClick={() => setShareModalState(true)} className="text-muted-foreground hover:text-foreground rounded-full" > - + {isShareModalOpen && ( {t('conversation.attachments.attach')} diff --git a/frontend/src/components/message-input/SourcesTrigger.tsx b/frontend/src/components/message-input/SourcesTrigger.tsx index 7feee685..f8edb1fc 100644 --- a/frontend/src/components/message-input/SourcesTrigger.tsx +++ b/frontend/src/components/message-input/SourcesTrigger.tsx @@ -61,7 +61,7 @@ export default function SourcesTrigger({ Sources {selectedDocs && selectedDocs.length > 0 diff --git a/frontend/src/components/message-input/ToolsTrigger.tsx b/frontend/src/components/message-input/ToolsTrigger.tsx index e9f16c69..6fb2a49a 100644 --- a/frontend/src/components/message-input/ToolsTrigger.tsx +++ b/frontend/src/components/message-input/ToolsTrigger.tsx @@ -62,7 +62,7 @@ export default function ToolsTrigger({ Tools {t('settings.tools.label')} diff --git a/frontend/src/conversation/ConversationTile.tsx b/frontend/src/conversation/ConversationTile.tsx index f5ead293..cfc5d87d 100644 --- a/frontend/src/conversation/ConversationTile.tsx +++ b/frontend/src/conversation/ConversationTile.tsx @@ -1,4 +1,4 @@ -import { ExternalLink, X } from 'lucide-react'; +import { Share, X } from 'lucide-react'; import { SyntheticEvent, useCallback, @@ -135,7 +135,7 @@ export default function ConversationTile({ const menuOptions: ConversationMenuOption[] = [ { - icon: , + icon: , label: t('convTile.share'), onClick: (event: SyntheticEvent) => { event.stopPropagation(); @@ -216,11 +216,13 @@ export default function ConversationTile({
{isEdit ? (
- Edit { event.stopPropagation(); handleSaveConversation({ @@ -228,12 +230,15 @@ export default function ConversationTile({ name: conversationName, }); }} - /> + > + + From fe68fec69e06b149bf31838b5a07dedef2586819 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 21:33:33 +0100 Subject: [PATCH 028/130] feat: docsgpt backup and docsgpt restore `docsgpt backup` writes one archive holding a pg_dump of the database, a tar of each data volume and a manifest of what it came from; `docsgpt restore` puts it back over an install. The settings file is left out unless --with-settings asks for it, since it holds the install's secrets, and a backup taken with a newer DocsGPT is refused without --force. The volume tars go through the image the install already runs, so a backup pulls nothing extra, and compose calls can now redirect stdout and stdin so the dump never passes through this process. --- docs/content/Deploying/Docker-Deploying.mdx | 30 +++++ docs/content/changelog.mdx | 10 ++ docsgpt/cli.py | 12 ++ docsgpt/deploy/backup.py | 107 ++++++++++++++++ docsgpt/deploy/commands.py | 115 +++++++++++++++++ docsgpt/deploy/docker.py | 36 +++++- tests/deploy/test_backup.py | 131 ++++++++++++++++++++ tests/deploy/test_commands.py | 25 +++- tests/deploy/test_docker.py | 30 ++++- 9 files changed, 489 insertions(+), 7 deletions(-) create mode 100644 docsgpt/deploy/backup.py create mode 100644 tests/deploy/test_backup.py diff --git a/docs/content/Deploying/Docker-Deploying.mdx b/docs/content/Deploying/Docker-Deploying.mdx index b4c5f87f..847c7453 100644 --- a/docs/content/Deploying/Docker-Deploying.mdx +++ b/docs/content/Deploying/Docker-Deploying.mdx @@ -79,6 +79,36 @@ the questions again. | `docsgpt down` | Stop the stack; data and settings stay | | `docsgpt uninstall [--purge]` | Remove the containers; `--purge` also deletes the settings and all data | +### Backups + +`docsgpt backup` writes one archive holding a dump of the database and a tar of +each data volume (`indexes`, `inputs`, `vectors`): + +```bash +docsgpt backup # into /backups +docsgpt backup --out /mnt/backups # somewhere else, e.g. a mounted disk +``` + +The archive does **not** include `.env`, because that file holds the install's +secrets. `docsgpt backup --with-settings` puts it in, for when the archive +itself is stored somewhere private. Keep `.env` safe separately otherwise: the +database password in it is what an existing Postgres volume expects. + +Restoring replaces the data in an install: + +```bash +docsgpt restore ~/.docsgpt/server/backups/docsgpt-20260916-120000.tar.gz +``` + +It asks first, then stops the stack, puts the volumes and the database back, and +starts DocsGPT again. `--yes` skips the question for scripts. A backup taken +with a newer DocsGPT is refused, since its data may not fit this version's +schema; upgrade first, or pass `--force` if you know the two match. + +The Postgres data directory itself is not archived: the dump is the database +backup, and copying a directory Postgres is writing to would capture a torn +copy. Caddy's certificates are not archived either, as it obtains them again. + More `docsgpt up` options: `--port`, `--docling` (the image with the docling parser engine and OCR), `--image-tag develop` (follow the `main` branch) and `--adopt` (manage a stack you started from the standalone Compose file in diff --git a/docs/content/changelog.mdx b/docs/content/changelog.mdx index 6e59a1f5..2c5f9501 100644 --- a/docs/content/changelog.mdx +++ b/docs/content/changelog.mdx @@ -11,6 +11,16 @@ The notable changes in each release. Every release on GitHub also carries [auto-generated notes](https://github.com/arc53/DocsGPT/releases) listing every merged pull request, and [Upgrading](/upgrading) covers the steps an existing deployment has to take. +## Unreleased + +### Back up and restore an install + +`docsgpt backup` writes a dump of the database and a tar of each data volume into one archive, and +`docsgpt restore ` puts them back. The settings file is left out unless +`--with-settings` asks for it, since it holds the install's secrets, and a backup from a newer +DocsGPT is refused unless you pass `--force`. See +[Backups](/Deploying/Docker-Deploying#backups). + ## 0.21.0 ### Install with one command diff --git a/docsgpt/cli.py b/docsgpt/cli.py index 873cd099..e5a5af2a 100644 --- a/docsgpt/cli.py +++ b/docsgpt/cli.py @@ -230,6 +230,18 @@ def _add_deploy_commands(commands) -> None: uninstall.add_argument("-y", "--yes", action="store_true", help="do not ask for confirmation") uninstall.add_argument("--purge", action="store_true", help="also delete the settings and all data") + backup = stack_command("backup", "backup", "write a backup of the database and the uploaded data") + backup.add_argument("--out", help="directory for the archive (default: /backups)") + backup.add_argument("--with-settings", action="store_true", + help="include .env in the archive; it holds this install's secrets") + + restore = stack_command("restore", "restore", "restore a backup over this install") + restore.add_argument("archive", help="the .tar.gz written by `docsgpt backup`") + restore.add_argument("-y", "--yes", action="store_true", help="do not ask for confirmation") + restore.add_argument("--force", action="store_true", help="restore a backup taken with a newer DocsGPT") + restore.add_argument("--timeout", type=int, default=300, + help="seconds to wait for the API afterwards (default: 300)") + env = stack_command("env", "env", "show, get or set the stack's settings") env_actions = env.add_subparsers(dest="env_action", metavar="") get = env_actions.add_parser("get", help="print one setting") diff --git a/docsgpt/deploy/backup.py b/docsgpt/deploy/backup.py new file mode 100644 index 00000000..5b041c23 --- /dev/null +++ b/docsgpt/deploy/backup.py @@ -0,0 +1,107 @@ +"""The archive ``docsgpt backup`` writes and ``docsgpt restore`` reads. + +One gzipped tar holds a SQL dump of the database, a tar per data volume, and a +manifest saying which version and image the backup came from. The settings file +is left out unless it is asked for: it holds the secrets. + +``postgres_data`` is not tarred, because the dump is the database backup and a +copy of a running data directory would be a torn one. Caddy's volumes are left +out too: they hold certificates it obtains again on the next start. +""" + +from __future__ import annotations + +import io +import json +import tarfile +import time +from collections.abc import Mapping +from datetime import datetime +from pathlib import Path +from typing import Optional + +from docsgpt.deploy.docker import DeployError + +FORMAT = 1 +MANIFEST = "manifest.json" +DUMP = "database.sql" +SETTINGS = "settings.env" +VOLUME_DIR = "volumes" +DATA_VOLUMES = ("indexes", "inputs", "vectors") + + +def archive_name(when: datetime) -> str: + """The file name for a backup taken at ``when``.""" + return f"docsgpt-{when.strftime('%Y%m%d-%H%M%S')}.tar.gz" + + +def volume_member(name: str) -> str: + """Where a volume's tar sits inside the archive.""" + return f"{VOLUME_DIR}/{name}.tar" + + +def write_archive( + path: Path, + *, + dump: Path, + volume_tars: Mapping[str, Path], + manifest: Mapping[str, object], + settings: Optional[Path] = None, +) -> None: + """Write the backup archive; the manifest is added last so a truncated file has none.""" + path.parent.mkdir(parents=True, exist_ok=True) + with tarfile.open(path, "w:gz") as archive: + archive.add(dump, arcname=DUMP) + for name, tar_path in sorted(volume_tars.items()): + archive.add(tar_path, arcname=volume_member(name)) + if settings is not None: + archive.add(settings, arcname=SETTINGS) + body = json.dumps(dict(manifest), indent=2).encode("utf-8") + info = tarfile.TarInfo(MANIFEST) + info.size = len(body) + info.mtime = int(time.time()) + info.mode = 0o600 + archive.addfile(info, io.BytesIO(body)) + + +def read_manifest(path: Path) -> dict: + """The archive's manifest, or a DeployError naming what is wrong with the file.""" + if not path.is_file(): + raise DeployError(f"{path} does not exist") + try: + with tarfile.open(path, "r:gz") as archive: + member = archive.extractfile(MANIFEST) + if member is None: + raise KeyError(MANIFEST) + manifest = json.loads(member.read().decode("utf-8")) + except (tarfile.TarError, KeyError, ValueError, OSError) as exc: + raise DeployError(f"{path} is not a DocsGPT backup: {exc}") from exc + if not isinstance(manifest, dict) or "volumes" not in manifest: + raise DeployError(f"{path} is not a DocsGPT backup: its manifest is missing what to restore") + return manifest + + +def extract(path: Path, destination: Path) -> None: + """Unpack the archive into ``destination`` (data filter: no paths outside it, no devices).""" + with tarfile.open(path, "r:gz") as archive: + archive.extractall(destination, filter="data") + + +def _parts(version: str) -> tuple[int, ...]: + numbers = [] + for chunk in str(version).split("."): + digits = "".join(character for character in chunk if character.isdigit()) + numbers.append(int(digits) if digits else 0) + return tuple(numbers) + + +def check_version(manifest: Mapping[str, object], current: str, force: bool) -> None: + """Refuse a backup from a newer DocsGPT: its data may not fit this version's schema.""" + taken_with = str(manifest.get("version") or "") + if force or not taken_with: + return + if _parts(taken_with) > _parts(current): + raise DeployError( + f"this backup is from DocsGPT {taken_with}, newer than the installed {current}. " + "Upgrade first with `docsgpt upgrade`, or pass --force to restore it anyway." + ) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index eaa38f44..fb9c5e25 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -9,6 +9,7 @@ import re import shutil import subprocess import sys +import tempfile import webbrowser from collections.abc import Callable, Mapping from dataclasses import dataclass @@ -16,6 +17,7 @@ from datetime import datetime, timezone from pathlib import Path from typing import Any, Optional +from docsgpt.deploy import backup as backup_format from docsgpt.deploy import envfile, stack from docsgpt.deploy.docker import DeployError, Docker, lan_ip, wait_healthy @@ -350,6 +352,119 @@ def env(args, context: Optional[Context] = None) -> int: return 0 +def _stack_image(env: Mapping[str, str]) -> str: + """The image this install runs; the volume tars go through it, so nothing extra is pulled.""" + tag = env.get("DOCSGPT_IMAGE_TAG") or "latest" + return f"arc53/docsgpt:{tag}{env.get('DOCSGPT_IMAGE_VARIANT', '')}" + + +def backup(args, context: Optional[Context] = None) -> int: + """Write a dump of the database and a tar of each data volume into one archive.""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + env = _installed(directory) + if env is None: + return 1 + + out_dir = Path(args.out).expanduser() if args.out else directory / "backups" + taken_at = datetime.now(timezone.utc) + target = out_dir / backup_format.archive_name(taken_at) + image = _stack_image(env) + + with tempfile.TemporaryDirectory() as workspace: + work = Path(workspace) + dump = work / backup_format.DUMP + print("Dumping the database ...") + with dump.open("w", encoding="utf-8") as handle: + context.docker.compose( + directory, "exec", "-T", "postgres", + "pg_dump", "--clean", "--if-exists", "-U", "docsgpt", "-d", "docsgpt", + stdout=handle, + ) + volume_tars = {} + for name in backup_format.DATA_VOLUMES: + print(f"Archiving the {name} volume ...") + tar_path = work / f"{name}.tar" + context.docker.export_volume(f"{PROJECT}_{name}", tar_path, image) + volume_tars[name] = tar_path + manifest = { + "format": backup_format.FORMAT, + "created_at": taken_at.isoformat(timespec="seconds"), + "version": context.version, + "image_tag": env.get("DOCSGPT_IMAGE_TAG", ""), + "volumes": list(backup_format.DATA_VOLUMES), + "settings_included": bool(args.with_settings), + } + settings = directory / ".env" if args.with_settings else None + backup_format.write_archive(target, dump=dump, volume_tars=volume_tars, manifest=manifest, settings=settings) + + size = target.stat().st_size / 1_000_000 + print(f"\nBackup written to {target} ({size:.1f} MB)") + if args.with_settings: + print("It contains .env, so it holds this install's secrets: keep it somewhere private.") + else: + print("Settings are not in it; `docsgpt backup --with-settings` includes .env, secrets and all.") + print(f"Restore it with: docsgpt restore {target}") + return 0 + + +def restore(args, context: Optional[Context] = None) -> int: + """Put a backup's database and data volumes back over this install.""" + context = context or Context.default(args) + archive = Path(args.archive).expanduser() + manifest = backup_format.read_manifest(archive) + backup_format.check_version(manifest, context.version, args.force) + + directory = stack.stack_dir(args.dir) + env = _installed(directory) + if env is None: + return 1 + + taken_at = manifest.get("created_at", "an unknown time") + if not args.yes: + if not context.interactive: + raise DeployError("restore needs --yes when there is no terminal to confirm on") + question = f"Replace the data in {directory} with the backup from {taken_at}? This cannot be undone." + if not context.prompter.confirm(question, default=False): + print("Nothing restored.") + return 1 + + image = _stack_image(env) + print("Stopping the stack ...") + context.docker.compose(directory, *_EVERY_PROFILE, "down") + + with tempfile.TemporaryDirectory() as workspace: + work = Path(workspace) + backup_format.extract(archive, work) + for name in manifest.get("volumes", []): + tar_path = work / backup_format.volume_member(name) + if not tar_path.is_file(): + raise DeployError(f"{archive} is missing the {name} volume it says it contains") + print(f"Restoring the {name} volume ...") + context.docker.import_volume(f"{PROJECT}_{name}", tar_path, image) + + dump = work / backup_format.DUMP + if not dump.is_file(): + raise DeployError(f"{archive} is missing its database dump") + print("Starting the database ...") + context.docker.compose(directory, "up", "-d", "--wait", "postgres") + print("Restoring the database ...") + with dump.open("r", encoding="utf-8") as handle: + context.docker.compose( + directory, "exec", "-T", "postgres", + "psql", "--quiet", "-U", "docsgpt", "-d", "docsgpt", + stdin=handle, + ) + + print("Starting DocsGPT ...") + context.docker.compose(directory, "up", "-d", "--remove-orphans") + if not context.wait(stack.health_url(env), args.timeout): + print("DocsGPT did not answer after the restore. See `docsgpt logs backend`.", file=sys.stderr) + return 1 + print(f"\nRestored the backup from {taken_at}. DocsGPT is running at {stack.url(env, context.lan_ip())}") + return 0 + + def _uninstall_hint(installer: str) -> str: return {"uv": "uv tool uninstall docsgpt", "pipx": "pipx uninstall docsgpt"}.get(installer, "pip uninstall docsgpt") diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py index 795aaeb9..69cbe187 100644 --- a/docsgpt/deploy/docker.py +++ b/docsgpt/deploy/docker.py @@ -22,10 +22,16 @@ class DeployError(Exception): """A problem the user can act on; the command prints it without a traceback.""" -def run(args: Sequence[str], *, cwd: Optional[Path] = None, capture: bool = False, check: bool = True): - """Run a command, streaming its output unless ``capture``; with ``check`` a failure raises DeployError.""" +def run(args: Sequence[str], *, cwd: Optional[Path] = None, capture: bool = False, check: bool = True, + stdout=None, stdin=None): + """Run a command, streaming its output unless ``capture`` or a file is given for ``stdout``. + + ``stdout`` and ``stdin`` take open files, so a dump goes straight to disk and + back in again without passing through this process. + """ + streams = {"capture_output": True} if capture else {"stdout": stdout, "stdin": stdin} try: - result = subprocess.run(list(args), cwd=cwd, text=True, capture_output=capture, check=False) + result = subprocess.run(list(args), cwd=cwd, text=True, check=False, **streams) except FileNotFoundError as exc: raise DeployError(f"{args[0]} is not installed or not on PATH") from exc if check and result.returncode != 0: @@ -94,9 +100,29 @@ class Docker: raise DeployError("Docker is not running. Start it with `sudo systemctl start docker` and run this again.") raise DeployError("Docker is not running. Start Docker Desktop and run this again.") - def compose(self, directory: Path, *args: str, capture: bool = False, check: bool = True): + def compose(self, directory: Path, *args: str, capture: bool = False, check: bool = True, + stdout=None, stdin=None): """``docker compose `` in ``directory``, which holds the Compose file and its ``.env``.""" - return self._run(["docker", "compose", *args], cwd=directory, capture=capture, check=check) + return self._run(["docker", "compose", *args], cwd=directory, capture=capture, check=check, + stdout=stdout, stdin=stdin) + + def export_volume(self, volume: str, dest: Path, image: str) -> None: + """Write ``volume`` to ``dest`` as a tar, through an image the install already has.""" + dest.parent.mkdir(parents=True, exist_ok=True) + with dest.open("wb") as handle: + self._run( + ["docker", "run", "--rm", "-v", f"{volume}:/data:ro", image, "tar", "cf", "-", "-C", "/data", "."], + stdout=handle, + ) + + def import_volume(self, volume: str, source: Path, image: str) -> None: + """Replace ``volume``'s contents with the tar at ``source``; the volume is created when missing.""" + with source.open("rb") as handle: + self._run( + ["docker", "run", "--rm", "-i", "-v", f"{volume}:/data", image, + "sh", "-c", "find /data -mindepth 1 -delete && tar xf - -C /data"], + stdin=handle, + ) def volume_exists(self, name: str) -> bool: return self._run(["docker", "volume", "inspect", name], capture=True, check=False).returncode == 0 diff --git a/tests/deploy/test_backup.py b/tests/deploy/test_backup.py new file mode 100644 index 00000000..be43bfcf --- /dev/null +++ b/tests/deploy/test_backup.py @@ -0,0 +1,131 @@ +"""`docsgpt backup` and `docsgpt restore`, against a fake Docker.""" + +import io +import json +import tarfile + +import pytest + +from docsgpt.deploy import backup as backup_module +from docsgpt.deploy.docker import DeployError + +from .test_commands import FakeDocker, FakePrompter, _context, _run + + +def _installed(tmp_path, *extra): + assert _run(["up", "--yes", "--dir", str(tmp_path), *extra], _context()) == 0 + + +def _archive_names(path): + with tarfile.open(path, "r:gz") as archive: + return sorted(member.name for member in archive.getmembers() if member.isfile()) + + +def _manifest(path): + with tarfile.open(path, "r:gz") as archive: + return json.loads(archive.extractfile(backup_module.MANIFEST).read()) + + +def _copy_with_version(archive, version, dest): + """The same archive, with another DocsGPT version written into its manifest.""" + with tarfile.open(archive, "r:gz") as source, tarfile.open(dest, "w:gz") as out: + for member in source.getmembers(): + body = source.extractfile(member).read() if member.isfile() else None + if member.name == backup_module.MANIFEST: + manifest = json.loads(body.decode()) + manifest["version"] = version + body = json.dumps(manifest).encode() + member.size = len(body) + out.addfile(member, io.BytesIO(body) if body is not None else None) + return dest + + +class TestBackup: + def test_writes_a_dump_a_tar_per_volume_and_a_manifest(self, tmp_path, capsys): + _installed(tmp_path) + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + out = tmp_path / "backups" + assert _run(["backup", "--dir", str(tmp_path), "--out", str(out)], _context(docker)) == 0 + + archives = list(out.glob("docsgpt-*.tar.gz")) + assert len(archives) == 1, archives + names = _archive_names(archives[0]) + assert backup_module.DUMP in names + for volume in ("indexes", "inputs", "vectors"): + assert f"volumes/{volume}.tar" in names + manifest = _manifest(archives[0]) + assert manifest["version"] == "0.21.0" + assert manifest["image_tag"] == "0.21.0" + assert manifest["volumes"] == ["indexes", "inputs", "vectors"] + assert manifest["settings_included"] is False + assert "Backup written to" in capsys.readouterr().out + assert [op for op in docker.volume_ops if op[0] == "export"], docker.volume_ops + + def test_the_settings_file_is_left_out_unless_asked_for(self, tmp_path): + _installed(tmp_path) + out = tmp_path / "backups" + assert _run(["backup", "--dir", str(tmp_path), "--out", str(out)], _context()) == 0 + assert backup_module.SETTINGS not in _archive_names(next(out.glob("*.tar.gz"))) + + with_settings = tmp_path / "with-settings" + argv = ["backup", "--dir", str(tmp_path), "--out", str(with_settings), "--with-settings"] + assert _run(argv, _context()) == 0 + assert backup_module.SETTINGS in _archive_names(next(with_settings.glob("*.tar.gz"))) + + def test_the_database_is_dumped_from_the_running_container(self, tmp_path): + _installed(tmp_path) + docker = FakeDocker() + assert _run(["backup", "--dir", str(tmp_path), "--out", str(tmp_path / "b")], _context(docker)) == 0 + dumps = [args for _, args in docker.calls if "pg_dump" in " ".join(args)] + assert dumps, docker.calls + assert dumps[0][:3] == ["exec", "-T", "postgres"] + + def test_without_an_install(self, tmp_path, capsys): + assert _run(["backup", "--dir", str(tmp_path)], _context()) == 1 + assert "docsgpt up" in capsys.readouterr().err + + +class TestRestore: + def _backup(self, tmp_path, *extra): + _installed(tmp_path) + out = tmp_path / "backups" + assert _run(["backup", "--dir", str(tmp_path), "--out", str(out), *extra], _context()) == 0 + return next(out.glob("*.tar.gz")) + + def test_restores_the_volumes_and_the_database(self, tmp_path): + archive = self._backup(tmp_path) + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + argv = ["restore", str(archive), "--dir", str(tmp_path), "--yes"] + assert _run(argv, _context(docker)) == 0 + joined = [" ".join(args) for _, args in docker.calls] + assert any(call.startswith("--profile https down") for call in joined), joined + assert any("psql" in call for call in joined), joined + assert any(call.startswith("up -d") for call in joined), joined + + def test_asks_before_replacing_data(self, tmp_path): + archive = self._backup(tmp_path) + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + prompter = FakePrompter([False]) + assert _run(["restore", str(archive), "--dir", str(tmp_path)], _context(docker, prompter, interactive=True)) == 1 + assert docker.calls == [] + + def test_refuses_an_archive_that_is_not_a_docsgpt_backup(self, tmp_path): + stray = tmp_path / "stray.tar.gz" + with tarfile.open(stray, "w:gz") as archive: + note = tmp_path / "note.txt" + note.write_text("not a backup") + archive.add(note, arcname="note.txt") + with pytest.raises(DeployError, match="not a DocsGPT backup"): + _run(["restore", str(stray), "--dir", str(tmp_path), "--yes"], _context()) + + def test_a_missing_archive(self, tmp_path): + with pytest.raises(DeployError, match="does not exist"): + _run(["restore", str(tmp_path / "nope.tar.gz"), "--dir", str(tmp_path), "--yes"], _context()) + + def test_a_newer_backup_is_refused_without_force(self, tmp_path): + """Restoring a 0.22 backup into 0.21 would hand an older schema newer data.""" + newer = _copy_with_version(self._backup(tmp_path), "0.22.0", tmp_path / "newer.tar.gz") + with pytest.raises(DeployError, match="newer"): + _run(["restore", str(newer), "--dir", str(tmp_path), "--yes"], _context(FakeDocker())) + argv = ["restore", str(newer), "--dir", str(tmp_path), "--yes", "--force"] + assert _run(argv, _context(FakeDocker())) == 0 diff --git a/tests/deploy/test_commands.py b/tests/deploy/test_commands.py index 30faa350..b73fedb2 100644 --- a/tests/deploy/test_commands.py +++ b/tests/deploy/test_commands.py @@ -1,8 +1,10 @@ """`docsgpt up` and the commands that manage the stack, against a fake Docker.""" +import io import json import subprocess import sys +import tarfile from pathlib import Path import pytest @@ -19,17 +21,38 @@ class FakeDocker: self.volumes = set(volumes) self.dirs = {Path(d) for d in project_dirs} self.calls = [] + self.volume_ops = [] self.preflights = 0 def preflight(self, interactive=False): self.preflights += 1 - def compose(self, directory, *args, capture=False, check=True): + def compose(self, directory, *args, capture=False, check=True, stdout=None, stdin=None): self.calls.append((Path(directory), list(args))) if "down" in args and "-v" in args: self.volumes.clear() + if stdout is not None: + stdout.write("-- fake pg_dump\n") + if stdin is not None: + self.restored_sql = stdin.read() return subprocess.CompletedProcess(["docker", "compose", *args], 0, stdout="", stderr="") + def export_volume(self, volume, dest, image): + """Write a small tar, as the real one does with `docker run ... tar cf -`.""" + self.volume_ops.append(("export", volume, image)) + dest = Path(dest) + dest.parent.mkdir(parents=True, exist_ok=True) + with tarfile.open(dest, "w") as tar: + body = f"{volume} contents".encode() + info = tarfile.TarInfo(f"{volume}.marker") + info.size = len(body) + tar.addfile(info, io.BytesIO(body)) + + def import_volume(self, volume, source, image): + self.volume_ops.append(("import", volume, image)) + self.volumes.add(volume) + assert Path(source).is_file(), source + def volume_exists(self, name): return name in self.volumes diff --git a/tests/deploy/test_docker.py b/tests/deploy/test_docker.py index 8e42556e..0ae2dfe4 100644 --- a/tests/deploy/test_docker.py +++ b/tests/deploy/test_docker.py @@ -15,9 +15,11 @@ class FakeRunner: def __init__(self, answers=None): self.answers = list((answers or {}).items()) self.calls = [] + self.streams = [] - def __call__(self, args, *, cwd=None, capture=False, check=True): + def __call__(self, args, *, cwd=None, capture=False, check=True, stdout=None, stdin=None): self.calls.append((list(args), cwd)) + self.streams.append((stdout, stdin)) for prefix, answer in self.answers: if list(args[: len(prefix)]) == list(prefix): if callable(answer): @@ -105,6 +107,32 @@ class TestQueries: assert "label=com.docker.compose.project=docsgpt" in args +class TestVolumes: + def test_export_writes_the_volume_through_the_stack_image(self, tmp_path): + runner = FakeRunner() + dest = tmp_path / "volumes" / "inputs.tar" + _docker(runner).export_volume("docsgpt_inputs", dest, "arc53/docsgpt:0.21.0") + args, _ = runner.calls[0] + assert args[:4] == ["docker", "run", "--rm", "-v"] + assert args[4] == "docsgpt_inputs:/data:ro" + assert args[5] == "arc53/docsgpt:0.21.0" + assert args[6:] == ["tar", "cf", "-", "-C", "/data", "."] + assert dest.is_file(), "the tar is written to the destination" + assert runner.streams[0][0] is not None, "stdout goes to the file, not through this process" + + def test_import_replaces_the_volume_contents(self, tmp_path): + source = tmp_path / "inputs.tar" + source.write_bytes(b"tar") + runner = FakeRunner() + _docker(runner).import_volume("docsgpt_inputs", source, "arc53/docsgpt:0.21.0") + args, _ = runner.calls[0] + assert "docsgpt_inputs:/data" in args + assert args[-2] == "-c" + assert "tar xf - -C /data" in args[-1] + assert "find /data -mindepth 1 -delete" in args[-1], "old contents go first" + assert runner.streams[0][1] is not None, "the tar is fed in on stdin" + + class TestWaitHealthy: def test_succeeds_once_the_api_answers(self): attempts = {"n": 0} From f23a32d9c5a3e27939a45038d2653ff4878fc2ab Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 21:51:43 +0100 Subject: [PATCH 029/130] feat: docsgpt up --native Run DocsGPT without Docker: the API and the worker each become a service on the machine itself, a launchd agent on macOS and a systemd user unit on Linux, pointed at a PostgreSQL and a Redis that already run. `docsgpt up --native --postgres-uri ... --redis-url ...` writes the same .env a Docker install uses, applies the migrations and starts both services. status, logs, down and uninstall work on a native install the same way they do on a Docker one, and never touch the database or Redis: they were the user's to begin with. One Redis URL covers the broker, the result backend and the cache on three consecutive databases, starting at the one the URL names, so a Redis that already holds something else can be shared. Windows has neither service manager, so native mode refuses it and says what to do instead. --- .../config/vocabularies/DocsGPT/accept.txt | 6 +- docs/content/Deploying/Pip-Install.mdx | 29 ++ docs/content/changelog.mdx | 8 + docsgpt/cli.py | 4 + docsgpt/deploy/commands.py | 180 ++++++++++- docsgpt/deploy/native.py | 218 ++++++++++++++ tests/deploy/test_native.py | 284 ++++++++++++++++++ 7 files changed, 722 insertions(+), 7 deletions(-) create mode 100644 docsgpt/deploy/native.py create mode 100644 tests/deploy/test_native.py diff --git a/.github/styles/config/vocabularies/DocsGPT/accept.txt b/.github/styles/config/vocabularies/DocsGPT/accept.txt index 7605df49..48520858 100644 --- a/.github/styles/config/vocabularies/DocsGPT/accept.txt +++ b/.github/styles/config/vocabularies/DocsGPT/accept.txt @@ -3,9 +3,9 @@ Anthropic's api APIs Atlassian -automations autoescaping Autoescaping +automations backfill backfills bool @@ -21,9 +21,9 @@ diarization Docling docsgpt docstrings +enqueues Entra env -enqueues EOL ESLint feedbacks @@ -35,6 +35,7 @@ hardcoding Idempotency JSONPath kubectl +launchd Lightsail llama_cpp llm @@ -70,6 +71,7 @@ SGLang Shareability Signup Supabase +systemd UIs uncomment URl diff --git a/docs/content/Deploying/Pip-Install.mdx b/docs/content/Deploying/Pip-Install.mdx index 1b946687..04574490 100644 --- a/docs/content/Deploying/Pip-Install.mdx +++ b/docs/content/Deploying/Pip-Install.mdx @@ -99,6 +99,35 @@ Other commands: - `docsgpt verify-offline`: check that a prepared install starts with networking off. - `docsgpt reembed`: re-embed every index after changing `EMBEDDINGS_NAME` (see [Upgrading](/upgrading)). +## Run it as services, without Docker + +`docsgpt up --native` runs the API and the worker as services on the machine itself: launchd agents on macOS, systemd user units on Linux. It does not start PostgreSQL or Redis; point it at ones you already run. + +```bash +docsgpt up --native \ + --postgres-uri postgresql://docsgpt:@localhost:5432/docsgpt \ + --redis-url redis://localhost:6379 +``` + +Without a terminal both flags are required; with one, it asks for them and for the model provider. It writes the same `.env` a Docker install uses (minus the image settings), generates `INTERNAL_KEY` and `JWT_SECRET_KEY` on the first run, applies the migrations, then starts `docsgpt-api` and `docsgpt-worker` and waits for the API to answer. + +One Redis URL covers all three uses: the Celery broker, its result backend and the cache go on databases 0, 1 and 2 of it. Name a database in the URL and the three start there instead, so `redis://localhost:6379/5` puts them on 5, 6 and 7 — that is how you share a Redis that already holds something else. + +The same commands manage it: + +| Command | In native mode | +| --- | --- | +| `docsgpt status` | Which services run, the address, and whether the API answers | +| `docsgpt logs [api\|worker]` | The service log files under `/logs` | +| `docsgpt down` | Stops both services; settings stay | +| `docsgpt uninstall [--purge]` | Removes the services; `--purge` also deletes the stack directory | + +`uninstall` never touches the database or Redis: they were yours to begin with. + + + Windows has neither launchd nor systemd, so native mode is macOS and Linux only. On Windows, run DocsGPT on Docker with `docsgpt up`, or start `docsgpt api` and `docsgpt worker` yourself. + + ## Upgrade ```bash diff --git a/docs/content/changelog.mdx b/docs/content/changelog.mdx index 2c5f9501..107685e7 100644 --- a/docs/content/changelog.mdx +++ b/docs/content/changelog.mdx @@ -13,6 +13,14 @@ request, and [Upgrading](/upgrading) covers the steps an existing deployment has ## Unreleased +### Run DocsGPT without Docker + +`docsgpt up --native` runs the API and the worker as services on the machine itself, launchd on +macOS and systemd user units on Linux, against a PostgreSQL and Redis you already have +(`--postgres-uri`, `--redis-url`). `status`, `logs`, `down` and `uninstall` work on such an install +the same way they do on a Docker one, and never touch the database or Redis. See +[Run it as services, without Docker](/Deploying/Pip-Install#run-it-as-services-without-docker). + ### Back up and restore an install `docsgpt backup` writes a dump of the database and a tar of each data volume into one archive, and diff --git a/docsgpt/cli.py b/docsgpt/cli.py index e5a5af2a..1ca2a187 100644 --- a/docsgpt/cli.py +++ b/docsgpt/cli.py @@ -206,6 +206,10 @@ def _add_deploy_commands(commands) -> None: help="go back to the default image") up.set_defaults(docling=None) up.add_argument("--image-tag", help="image tag to run instead of this package's version, e.g. develop") + up.add_argument("--native", action="store_true", + help="run the API and worker as services on this machine instead of on Docker") + up.add_argument("--postgres-uri", help="native mode: the PostgreSQL DocsGPT should use") + up.add_argument("--redis-url", help="native mode: the Redis for the queue and the cache (default: localhost:6379)") up.add_argument("-y", "--yes", action="store_true", help="ask nothing: use the flags, then the defaults") up.add_argument("--reconfigure", action="store_true", help="ask the setup questions again") up.add_argument("--adopt", action="store_true", help="take over a DocsGPT stack started from another folder") diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index fb9c5e25..26c287fc 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -6,6 +6,7 @@ import getpass import json import os import re +import secrets import shutil import subprocess import sys @@ -18,7 +19,7 @@ from pathlib import Path from typing import Any, Optional from docsgpt.deploy import backup as backup_format -from docsgpt.deploy import envfile, stack +from docsgpt.deploy import envfile, native, stack from docsgpt.deploy.docker import DeployError, Docker, lan_ip, wait_healthy PROJECT = "docsgpt" @@ -78,9 +79,10 @@ def detect_installer() -> str: return "pip" -def _run_command(args: list[str]) -> int: +def _run_command(args: list[str], env: Optional[Mapping[str, str]] = None) -> int: + """Run ``args``; ``env`` adds to this process's environment rather than replacing it.""" try: - return subprocess.call(args) + return subprocess.call(args, env={**os.environ, **env} if env else None) except FileNotFoundError as exc: raise DeployError(f"{args[0]} is not on PATH") from exc @@ -106,8 +108,9 @@ class Context: wait: Callable[[str, float], bool] = wait_healthy open_browser: Callable[[str], Any] = webbrowser.open installer: Callable[[], str] = detect_installer - run: Callable[[list[str]], int] = _run_command + run: Callable[..., int] = _run_command exec_up: Callable[[list[str]], int] = _exec_command + services: Any = None @classmethod def default(cls, args) -> "Context": @@ -116,10 +119,33 @@ class Context: interactive = sys.stdin.isatty() and not getattr(args, "yes", False) return cls(docker=Docker(), prompter=Prompter(), interactive=interactive, version=__version__) + def service_manager(self): + """The launchd or systemd wrapper, made on first use so Docker installs never touch it.""" + if self.services is None: + self.services = native.services_for_platform() + return self.services + + +def _record(directory: Path) -> dict: + """What ``install.json`` says about this install (empty when there is none).""" + path = directory / stack.RECORD_FILE + if not path.is_file(): + return {} + try: + return json.loads(path.read_text(encoding="utf-8")) + except ValueError: + return {} + + +def _mode(directory: Path) -> str: + """``native`` or ``docker``, from the install record.""" + return "native" if _record(directory).get("mode") == "native" else "docker" + def _installed(directory: Path) -> Optional[dict[str, str]]: """The stack's settings, or None (with a message) when there is no install in ``directory``.""" - if not (directory / stack.COMPOSE_FILE).is_file(): + installed = (directory / stack.COMPOSE_FILE).is_file() or _mode(directory) == "native" + if not installed: print(f"No DocsGPT install in {directory}. Run `docsgpt up` first, or pass --dir.", file=sys.stderr) return None return envfile.read(directory / ".env") @@ -174,10 +200,110 @@ def _choose_provider(args, context: Context, existing: Mapping[str, str], ask: b raise DeployError(str(exc)) from exc +def _redis_urls(base: str) -> dict[str, str]: + """Celery's broker and result backend, and the cache, on three consecutive databases of one Redis. + + They start at database 0, or at the one the URL names: ``redis://host:6379/5`` puts them on 5, 6 + and 7, which is how one Redis is shared with something that already uses the first databases. + """ + trimmed = base.rstrip("/") + first = 0 + head, _, tail = trimmed.rpartition("/") + if tail.isdigit(): + trimmed, first = head, int(tail) + return { + "CELERY_BROKER_URL": f"{trimmed}/{first}", + "CELERY_RESULT_BACKEND": f"{trimmed}/{first + 1}", + "CACHE_REDIS_URL": f"{trimmed}/{first + 2}", + } + + +def _native_up(args, context: Context, directory: Path) -> int: + """Run the API and the worker as services on this machine, against an existing Postgres and Redis.""" + services = context.service_manager() + env_path = directory / ".env" + existing = envfile.read(env_path) + record = _record(directory) + configured = bool(record) + + postgres = args.postgres_uri or existing.get("POSTGRES_URI") + redis = args.redis_url or "" + ask = context.interactive and (not configured or args.reconfigure) + if not postgres and ask: + postgres = context.prompter.text("PostgreSQL URL (postgresql://user:password@host:5432/docsgpt)") + if not redis and ask: + redis = context.prompter.text("Redis URL", default="redis://localhost:6379") + if not postgres: + raise DeployError( + "native mode needs a database: pass --postgres-uri postgresql://user:password@host:5432/docsgpt " + "(Redis defaults to redis://localhost:6379)." + ) + + # An install keeps the port it was given: a later `up` with no --port must not move it back to the default. + port = int(args.port or existing.get("DOCSGPT_PORT") or stack.DEFAULT_PORT) + updates: dict[str, Optional[str]] = { + "POSTGRES_URI": postgres, + "API_URL": f"http://127.0.0.1:{port}", + "DOCSGPT_PORT": str(port), + } + if redis or "CELERY_BROKER_URL" not in existing: + updates.update(_redis_urls(redis or "redis://localhost:6379")) + for key in ("INTERNAL_KEY", "JWT_SECRET_KEY"): + if not existing.get(key): + updates[key] = secrets.token_hex(32) + if "VITE_API_STREAMING" not in existing: + updates["VITE_API_STREAMING"] = "true" + provider = _choose_provider(args, context, existing, ask) + if provider: + updates.update(provider) + elif "LLM_PROVIDER" not in existing: + updates.update(stack.provider_settings("docsgpt")) + + directory.mkdir(parents=True, exist_ok=True) + (directory / "logs").mkdir(exist_ok=True) + envfile.update(env_path, updates) + + executable = shutil.which("docsgpt") or sys.argv[0] + print("Applying database migrations ...") + # The child reads the stack's settings, not a .env in whatever directory this was run from. + stack_env = {"DOCSGPT_HOME": str(directory), "DOCSGPT_ENV_FILE": str(env_path)} + if context.run([executable, "migrate"], stack_env) != 0: + raise DeployError("`docsgpt migrate` failed; check the database URL and that the server is reachable.") + + for unit in native.units_for(directory, executable, port, directory): + services.install(unit) + for name in native.SERVICES: + print(f"Starting {name} ...") + services.start(name) + + now = datetime.now(timezone.utc).isoformat(timespec="seconds") + record = record or {"installed_at": now} + record.update(version=context.version, mode="native", updated_at=now) + (directory / stack.RECORD_FILE).write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") + + health = f"http://127.0.0.1:{port}/api/health" + print("Waiting for DocsGPT to answer ...") + if not context.wait(health, args.timeout): + print( + f"DocsGPT did not answer at {health} within {args.timeout} seconds. " + "See what happened with `docsgpt logs`.", + file=sys.stderr, + ) + return 1 + + print(f"\nDocsGPT is running at http://localhost:{port}") + print(f"Services: {', '.join(native.SERVICES)} under {services.name}") + print(f"Settings: {env_path}") + print("Manage it with: docsgpt status | logs | down | uninstall") + return 0 + + def up(args, context: Optional[Context] = None) -> int: """Install or update the stack in its directory and start it.""" context = context or Context.default(args) directory = stack.stack_dir(args.dir) + if getattr(args, "native", False) or _mode(directory) == "native": + return _native_up(args, context, directory) env_path = directory / ".env" record_path = directory / stack.RECORD_FILE existing = envfile.read(env_path) @@ -269,6 +395,12 @@ def down(args, context: Optional[Context] = None) -> int: directory = stack.stack_dir(args.dir) if _installed(directory) is None: return 1 + if _mode(directory) == "native": + services = context.service_manager() + # The worker goes first: it talks to the API, not the other way round. + for name in reversed(native.SERVICES): + services.stop(name) + return 0 context.docker.compose(directory, *_EVERY_PROFILE, "down") return 0 @@ -280,6 +412,16 @@ def status(args, context: Optional[Context] = None) -> int: env = _installed(directory) if env is None: return 1 + if _mode(directory) == "native": + services = context.service_manager() + print(f"DocsGPT {_record(directory).get('version', 'unknown')} in {directory} (native, {services.name})") + print(f"Address: {stack.url(env, context.lan_ip())}") + for name in native.SERVICES: + print(f" {name}: {'running' if services.is_running(name) else 'stopped'}") + healthy = context.wait(stack.health_url(env), 0) + print("API: answering" if healthy else "API: not answering (see `docsgpt logs`)") + return 0 if healthy else 1 + tag = env.get("DOCSGPT_IMAGE_TAG", "unknown") + env.get("DOCSGPT_IMAGE_VARIANT", "") print(f"DocsGPT {tag} in {directory}") print(f"Address: {stack.url(env, context.lan_ip())} ({stack.exposure(env)})") @@ -295,6 +437,20 @@ def logs(args, context: Optional[Context] = None) -> int: directory = stack.stack_dir(args.dir) if _installed(directory) is None: return 1 + if _mode(directory) == "native": + logs_dir = directory / "logs" + wanted = args.services or [name.removeprefix("docsgpt-") for name in native.SERVICES] + for service in wanted: + path = logs_dir / f"{service}.log" + print(f"=== {path}") + if path.is_file(): + lines = path.read_text(encoding="utf-8", errors="replace").splitlines() + print("\n".join(lines[-args.tail:] if args.tail else lines)) + else: + print("(nothing logged yet)") + if args.follow: + print("Following is not supported in native mode; use `tail -f` on the files above.", file=sys.stderr) + return 0 options = (["--follow"] if args.follow else []) + (["--tail", str(args.tail)] if args.tail else []) return context.docker.compose(directory, "logs", *options, *args.services, check=False).returncode @@ -499,6 +655,20 @@ def uninstall(args, context: Optional[Context] = None) -> int: if not context.prompter.confirm(f"Remove the DocsGPT {what} in {directory}?", default=False): print("Nothing removed.") return 1 + if _mode(directory) == "native": + services = context.service_manager() + for name in reversed(native.SERVICES): + services.remove(name) + print(f"Removed the {', '.join(native.SERVICES)} services.") + if args.purge: + shutil.rmtree(directory) + print(f"Removed {directory}. The database and Redis it used are untouched.") + else: + (directory / stack.RECORD_FILE).unlink(missing_ok=True) + print(f"Settings stay in {directory / '.env'}; the database and Redis are untouched.") + print(f"To remove the docsgpt command too: {_uninstall_hint(context.installer())}") + return 0 + context.docker.compose(directory, *_EVERY_PROFILE, "down", "--remove-orphans", *(["-v"] if args.purge else [])) if args.purge: shutil.rmtree(directory) diff --git a/docsgpt/deploy/native.py b/docsgpt/deploy/native.py new file mode 100644 index 00000000..0ef0cf19 --- /dev/null +++ b/docsgpt/deploy/native.py @@ -0,0 +1,218 @@ +"""Running DocsGPT without Docker, supervised by the system's own service manager. + +The API and the worker each become one service: a launchd agent on macOS, a +systemd user unit on Linux. Both run the ``docsgpt`` command of this +installation with the stack directory as their data home, so a native install +keeps its settings in the same ``.env`` a Docker install would. + +Postgres and Redis are not started here; a native install points at ones that +already run (``--postgres-uri``, ``--redis-url``). +""" + +from __future__ import annotations + +import os +import plistlib +import shlex +import subprocess +import sys +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Optional + +from docsgpt.deploy.docker import DeployError + +API_SERVICE = "docsgpt-api" +WORKER_SERVICE = "docsgpt-worker" +SERVICES = (API_SERVICE, WORKER_SERVICE) +LABEL_PREFIX = "cloud.docsgpt" + + +@dataclass +class Unit: + """One supervised process, in the terms every service manager needs.""" + + name: str + arguments: list[str] + environment: dict[str, str] + working_directory: str + log_file: str + + +def label_for(name: str) -> str: + """The launchd label for a service name (``docsgpt-api`` -> ``cloud.docsgpt.api``).""" + return f"{LABEL_PREFIX}.{name.removeprefix('docsgpt-')}" + + +def launchd_plist(unit: Unit, label: Optional[str] = None) -> str: + """The launchd agent for ``unit``: kept alive, with its output in the stack's log file.""" + body = { + "Label": label or label_for(unit.name), + "ProgramArguments": list(unit.arguments), + "EnvironmentVariables": dict(unit.environment), + "WorkingDirectory": unit.working_directory, + "StandardOutPath": unit.log_file, + "StandardErrorPath": unit.log_file, + "KeepAlive": True, + "RunAtLoad": True, + "ProcessType": "Background", + } + return plistlib.dumps(body).decode("utf-8") + + +def systemd_unit(unit: Unit) -> str: + """The systemd user unit for ``unit``; arguments are quoted, so a path with spaces survives.""" + environment = "\n".join(f'Environment="{key}={value}"' for key, value in sorted(unit.environment.items())) + command = " ".join(shlex.quote(argument) for argument in unit.arguments) + return f"""[Unit] +Description=DocsGPT ({unit.name}) +After=network-online.target + +[Service] +Type=simple +ExecStart={command} +WorkingDirectory={unit.working_directory} +{environment} +Restart=always +RestartSec=5 +StandardOutput=append:{unit.log_file} +StandardError=append:{unit.log_file} + +[Install] +WantedBy=default.target +""" + + +class LaunchdServices: + """launchd user agents in ~/Library/LaunchAgents (macOS).""" + + name = "launchd" + poll_interval = 0.2 + unload_timeout = 15.0 + + def __init__(self, runner=subprocess.run, home: Optional[Path] = None) -> None: + self._run = runner + self.directory = (home or Path.home()) / "Library" / "LaunchAgents" + + def _plist(self, name: str) -> Path: + return self.directory / f"{label_for(name)}.plist" + + def _target(self, name: str) -> str: + return f"gui/{os.getuid()}/{label_for(name)}" + + def _print(self, name: str): + return self._run(["launchctl", "print", self._target(name)], capture_output=True, text=True, check=False) + + def _bootout(self, name: str) -> None: + """Unload the job and wait for it to go. + + ``bootout`` returns before launchd has finished unloading, and bootstrapping a label that is + still on its way out fails with "Bootstrap failed: 5: Input/output error" — which is what a + restart used to hit. Waiting for the label to stop resolving makes stop and start ordered. + """ + self._run(["launchctl", "bootout", self._target(name)], capture_output=True, text=True, check=False) + deadline = time.monotonic() + self.unload_timeout + while self._print(name).returncode == 0: + if time.monotonic() >= deadline: + raise DeployError(f"{label_for(name)} is still loaded after {self.unload_timeout:.0f}s") + time.sleep(self.poll_interval) + + def install(self, unit: Unit) -> None: + self.directory.mkdir(parents=True, exist_ok=True) + self._plist(unit.name).write_text(launchd_plist(unit), encoding="utf-8") + + def start(self, name: str) -> None: + # bootout first: a reinstall must pick up the new plist rather than the loaded one. + self._bootout(name) + result = self._run( + ["launchctl", "bootstrap", f"gui/{os.getuid()}", str(self._plist(name))], + capture_output=True, text=True, check=False, + ) + if result.returncode != 0: + raise DeployError(f"launchctl could not start {name}: {(result.stderr or '').strip()}") + + def stop(self, name: str) -> None: + self._bootout(name) + + def remove(self, name: str) -> None: + self.stop(name) + self._plist(name).unlink(missing_ok=True) + + def is_running(self, name: str) -> bool: + """A loaded job is not a running one: a crashed service still answers ``launchctl print``.""" + result = self._print(name) + return result.returncode == 0 and "state = running" in (result.stdout or "") + + +class SystemdServices: + """systemd user units in ~/.config/systemd/user (Linux).""" + + name = "systemd" + + def __init__(self, runner=subprocess.run, home: Optional[Path] = None) -> None: + self._run = runner + base = os.environ.get("XDG_CONFIG_HOME") + root = Path(base) if base else (home or Path.home()) / ".config" + self.directory = root / "systemd" / "user" + + def _unit_file(self, name: str) -> Path: + return self.directory / f"{name}.service" + + def _systemctl(self, *args: str, check: bool = False): + result = self._run(["systemctl", "--user", *args], capture_output=True, text=True, check=False) + if check and result.returncode != 0: + raise DeployError(f"systemctl --user {' '.join(args)} failed: {(result.stderr or '').strip()}") + return result + + def install(self, unit: Unit) -> None: + self.directory.mkdir(parents=True, exist_ok=True) + self._unit_file(unit.name).write_text(systemd_unit(unit), encoding="utf-8") + self._systemctl("daemon-reload") + + def start(self, name: str) -> None: + self._systemctl("enable", "--now", f"{name}.service", check=True) + + def stop(self, name: str) -> None: + self._systemctl("stop", f"{name}.service") + + def remove(self, name: str) -> None: + self._systemctl("disable", "--now", f"{name}.service") + self._unit_file(name).unlink(missing_ok=True) + self._systemctl("daemon-reload") + + def is_running(self, name: str) -> bool: + return self._systemctl("is-active", "--quiet", f"{name}.service").returncode == 0 + + +def services_for_platform(platform: str = sys.platform): + """The service manager for this machine, or a DeployError saying what to do instead.""" + if platform == "darwin": + return LaunchdServices() + if platform.startswith("linux"): + return SystemdServices() + raise DeployError( + "native mode supervises services with launchd or systemd, which Windows does not have. " + "Run DocsGPT on Docker with `docsgpt up`, or start `docsgpt api` and `docsgpt worker` yourself." + ) + + +def units_for(stack_directory: Path, executable: str, port: int, home: Path) -> list[Unit]: + """The API and worker services for a native install in ``stack_directory``.""" + environment = {"DOCSGPT_HOME": str(home)} + return [ + Unit( + name=API_SERVICE, + arguments=[executable, "api", "--host", "127.0.0.1", "--port", str(port)], + environment=dict(environment), + working_directory=str(stack_directory), + log_file=str(stack_directory / "logs" / "api.log"), + ), + Unit( + name=WORKER_SERVICE, + arguments=[executable, "worker"], + environment=dict(environment), + working_directory=str(stack_directory), + log_file=str(stack_directory / "logs" / "worker.log"), + ), + ] diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py new file mode 100644 index 00000000..63646d98 --- /dev/null +++ b/tests/deploy/test_native.py @@ -0,0 +1,284 @@ +"""`docsgpt up --native`: run the server without Docker, supervised by the system's service manager.""" + +import json +import plistlib +import types + +import pytest + +from docsgpt.deploy import envfile, native +from docsgpt.deploy.docker import DeployError + +from .test_commands import FakePrompter, _context, _run + + +class FakeServices: + """A stand-in for launchd/systemd: records what would be installed and started.""" + + name = "fake" + + def __init__(self, running=()): + self.units = {} + self.started = [] + self.stopped = [] + self.removed = [] + self.running = set(running) + + def install(self, unit): + self.units[unit.name] = unit + + def start(self, name): + self.started.append(name) + self.running.add(name) + + def stop(self, name): + self.stopped.append(name) + self.running.discard(name) + + def remove(self, name): + self.removed.append(name) + self.running.discard(name) + + def is_running(self, name): + return name in self.running + + +def _native_context(services=None, migrated=None, **overrides): + context = _context(**overrides) + context.services = services or FakeServices() + context.run = lambda args, env=None: (migrated.append(args) if migrated is not None else None) or 0 + return context + + +class TestNativeUp: + def test_writes_settings_for_the_given_postgres_and_redis(self, tmp_path): + context = _native_context() + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://docsgpt:pw@localhost:5432/docsgpt", + "--redis-url", "redis://localhost:6379"] + assert _run(argv, context) == 0 + + env = envfile.read(tmp_path / ".env") + assert env["POSTGRES_URI"] == "postgresql://docsgpt:pw@localhost:5432/docsgpt" + assert env["CELERY_BROKER_URL"] == "redis://localhost:6379/0" + assert env["CELERY_RESULT_BACKEND"] == "redis://localhost:6379/1" + assert env["CACHE_REDIS_URL"] == "redis://localhost:6379/2" + assert env["INTERNAL_KEY"], "the worker needs it to hand indexes to the API" + assert env["API_URL"] == "http://127.0.0.1:7091" + assert "DOCSGPT_IMAGE_TAG" not in env, "nothing here runs an image" + + def test_records_the_mode_so_the_other_commands_can_tell(self, tmp_path): + assert _run(["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/docsgpt"], _native_context()) == 0 + record = json.loads((tmp_path / "install.json").read_text()) + assert record["mode"] == "native" + assert record["version"] == "0.21.0" + + def test_migrates_before_starting_the_services(self, tmp_path): + migrated = [] + services = FakeServices() + context = _native_context(services=services, migrated=migrated) + assert _run(["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/docsgpt"], context) == 0 + assert any("migrate" in " ".join(call) for call in migrated), migrated + assert services.started == ["docsgpt-api", "docsgpt-worker"] + + def test_the_units_run_this_interpreter_with_the_stack_as_its_home(self, tmp_path): + services = FakeServices() + assert _run(["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/docsgpt"], _native_context(services)) == 0 + api = services.units["docsgpt-api"] + worker = services.units["docsgpt-worker"] + assert api.arguments[1:] == ["api", "--host", "127.0.0.1", "--port", "7091"] + assert worker.arguments[1:] == ["worker"] + for unit in (api, worker): + assert unit.environment["DOCSGPT_HOME"] == str(tmp_path) + assert unit.working_directory == str(tmp_path) + assert unit.log_file.endswith(".log") + + def test_the_migration_targets_the_stack_not_the_working_directory(self, tmp_path): + """`docsgpt up --native` run from a checkout must not migrate the checkout's database.""" + calls = [] + context = _native_context() + context.run = lambda args, env=None: calls.append((args, env)) or 0 + assert _run(["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/docsgpt"], context) == 0 + arguments, environment = calls[0] + assert arguments[1] == "migrate" + assert environment["DOCSGPT_HOME"] == str(tmp_path) + assert environment["DOCSGPT_ENV_FILE"] == str(tmp_path / ".env") + + def test_a_redis_database_in_the_url_moves_the_three_up(self, tmp_path): + """Sharing a Redis: `/5` puts the broker, the backend and the cache on 5, 6 and 7.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/docsgpt", "--redis-url", "redis://localhost:6379/5"] + assert _run(argv, _native_context()) == 0 + env = envfile.read(tmp_path / ".env") + assert env["CELERY_BROKER_URL"] == "redis://localhost:6379/5" + assert env["CELERY_RESULT_BACKEND"] == "redis://localhost:6379/6" + assert env["CACHE_REDIS_URL"] == "redis://localhost:6379/7" + + def test_a_later_up_keeps_the_port_the_install_was_given(self, tmp_path): + """`docsgpt up` with no --port must not move an install from its port back to 7091.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--port", "7099", + "--postgres-uri", "postgresql://localhost/docsgpt"] + assert _run(argv, _native_context()) == 0 + + services = FakeServices() + assert _run(["up", "--dir", str(tmp_path), "--yes"], _native_context(services)) == 0 + env = envfile.read(tmp_path / ".env") + assert env["DOCSGPT_PORT"] == "7099" + assert env["API_URL"] == "http://127.0.0.1:7099" + assert services.units["docsgpt-api"].arguments[-1] == "7099" + + def test_a_database_url_is_required(self, tmp_path): + with pytest.raises(DeployError, match="--postgres-uri"): + _run(["up", "--native", "--dir", str(tmp_path), "--yes"], _native_context()) + + def test_asks_for_the_urls_when_there_is_a_terminal(self, tmp_path): + prompter = FakePrompter(["postgresql://localhost/docsgpt", "redis://localhost:6379", "docsgpt"]) + context = _native_context(interactive=True, prompter=prompter) + assert _run(["up", "--native", "--dir", str(tmp_path)], context) == 0 + env = envfile.read(tmp_path / ".env") + assert env["POSTGRES_URI"] == "postgresql://localhost/docsgpt" + assert env["CELERY_BROKER_URL"] == "redis://localhost:6379/0" + assert env["LLM_PROVIDER"] == "docsgpt" + assert prompter.questions[0].startswith("PostgreSQL URL") + + def test_an_unhealthy_start_points_at_the_logs(self, tmp_path, capsys): + context = _native_context(healthy=False) + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/docsgpt"] + assert _run(argv, context) == 1 + assert "docsgpt logs" in capsys.readouterr().err + + +class TestNativeLifecycle: + def _installed(self, tmp_path, services): + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/docsgpt"] + assert _run(argv, _native_context(services)) == 0 + + def test_status_reports_the_native_services(self, tmp_path, capsys): + services = FakeServices() + self._installed(tmp_path, services) + assert _run(["status", "--dir", str(tmp_path)], _native_context(services)) == 0 + out = capsys.readouterr().out + assert "native" in out.lower() + assert "docsgpt-api" in out + + def test_down_stops_them_and_leaves_the_settings(self, tmp_path): + services = FakeServices() + self._installed(tmp_path, services) + assert _run(["down", "--dir", str(tmp_path)], _native_context(services)) == 0 + assert services.stopped == ["docsgpt-worker", "docsgpt-api"], "the worker goes first" + assert (tmp_path / ".env").is_file() + + def test_uninstall_removes_the_services(self, tmp_path): + services = FakeServices() + self._installed(tmp_path, services) + assert _run(["uninstall", "--yes", "--dir", str(tmp_path)], _native_context(services)) == 0 + assert sorted(services.removed) == ["docsgpt-api", "docsgpt-worker"] + + def test_docker_commands_refuse_a_native_install(self, tmp_path, capsys): + self._installed(tmp_path, FakeServices()) + assert _run(["logs", "--dir", str(tmp_path)], _native_context()) == 0, "logs work in both modes" + + +class FakeLaunchctl: + """launchctl, with a bootout that takes `unload_polls` polls to take effect.""" + + def __init__(self, unload_polls=1, state="running"): + self.calls = [] + self.loaded = True + self.unload_polls = unload_polls + self.state = state + + def __call__(self, args, capture_output=False, text=False, check=False): + self.calls.append(args) + verb = args[1] + if verb == "bootout": + self.pending = self.unload_polls + return types.SimpleNamespace(returncode=0, stdout="", stderr="") + if verb == "print": + if not self.loaded: + return types.SimpleNamespace(returncode=113, stdout="", stderr="Could not find service") + if getattr(self, "pending", 0) > 0: + self.pending -= 1 + if self.pending == 0: + self.loaded = False + return types.SimpleNamespace(returncode=0, stdout="\tstate = running\n", stderr="") + return types.SimpleNamespace(returncode=0, stdout=f"\tstate = {self.state}\n", stderr="") + if verb == "bootstrap": + if self.loaded: + return types.SimpleNamespace(returncode=5, stdout="", stderr="Bootstrap failed: 5: Input/output error") + self.loaded = True + return types.SimpleNamespace(returncode=0, stdout="", stderr="") + return types.SimpleNamespace(returncode=0, stdout="", stderr="") + + +class TestLaunchd: + def _services(self, launchctl, tmp_path): + services = native.LaunchdServices(runner=launchctl, home=tmp_path) + services.poll_interval = 0 + return services + + def test_start_waits_for_the_old_job_to_unload_before_bootstrapping(self, tmp_path): + """launchctl bootout returns early; bootstrapping too soon fails with 'Bootstrap failed: 5'.""" + launchctl = FakeLaunchctl(unload_polls=3) + services = self._services(launchctl, tmp_path) + services.directory.mkdir(parents=True, exist_ok=True) + services.start("docsgpt-api") + verbs = [call[1] for call in launchctl.calls] + assert verbs[0] == "bootout" + assert verbs[-1] == "bootstrap" + assert verbs.count("print") >= 1, "it polled until the label stopped resolving" + assert verbs.index("bootstrap") > max(index for index, verb in enumerate(verbs) if verb == "print") + + def test_stop_returns_only_once_the_job_is_gone(self, tmp_path): + launchctl = FakeLaunchctl(unload_polls=2) + self._services(launchctl, tmp_path).stop("docsgpt-api") + assert not launchctl.loaded + + def test_a_job_that_will_not_unload_is_reported(self, tmp_path): + launchctl = FakeLaunchctl(unload_polls=10_000) + services = self._services(launchctl, tmp_path) + services.unload_timeout = 0 + with pytest.raises(DeployError, match="still loaded"): + services.stop("docsgpt-api") + + def test_a_loaded_but_not_running_job_is_not_running(self, tmp_path): + """`launchctl print` answers for a crashed service too, so status must read its state.""" + services = self._services(FakeLaunchctl(state="not running"), tmp_path) + assert services.is_running("docsgpt-api") is False + assert self._services(FakeLaunchctl(state="running"), tmp_path).is_running("docsgpt-api") is True + + +class TestUnitFiles: + def _unit(self, tmp_path): + return native.Unit( + name="docsgpt-api", + arguments=["/opt/venv/bin/docsgpt", "api", "--port", "7091"], + environment={"DOCSGPT_HOME": str(tmp_path)}, + working_directory=str(tmp_path), + log_file=str(tmp_path / "api.log"), + ) + + def test_launchd_plist_is_valid_and_keeps_the_service_alive(self, tmp_path): + body = native.launchd_plist(self._unit(tmp_path), label="cloud.docsgpt.api") + parsed = plistlib.loads(body.encode()) + assert parsed["Label"] == "cloud.docsgpt.api" + assert parsed["ProgramArguments"][1] == "api" + assert parsed["KeepAlive"] is True + assert parsed["EnvironmentVariables"]["DOCSGPT_HOME"] == str(tmp_path) + assert parsed["StandardErrorPath"].endswith("api.log") + + def test_systemd_unit_restarts_and_carries_the_environment(self, tmp_path): + body = native.systemd_unit(self._unit(tmp_path)) + assert "Restart=always" in body + assert f'Environment="DOCSGPT_HOME={tmp_path}"' in body + assert "ExecStart=/opt/venv/bin/docsgpt api --port 7091" in body + assert "WantedBy=default.target" in body + + def test_an_argument_with_spaces_survives_the_systemd_unit(self, tmp_path): + unit = self._unit(tmp_path) + unit.arguments = ["/opt/my venv/bin/docsgpt", "api"] + assert "ExecStart='/opt/my venv/bin/docsgpt' api" in native.systemd_unit(unit) From 447ae72fe25ff91d1c91bcca46d39d909d54ddbd Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 21:55:33 +0100 Subject: [PATCH 030/130] fix: harden docsgpt backup and docsgpt restore From review of #2799: - The archive is created 0600 rather than at the process umask: it holds the install's data, and --with-settings puts .env and its secrets in it. - restore validates everything the manifest declares before the stack is stopped, so a damaged archive fails while DocsGPT is still running rather than after `compose down` has taken it away. - Only the volumes a backup is made of are restored. A hand-made manifest can no longer point import_volume at postgres_data, whose contents it empties. - psql runs with ON_ERROR_STOP=on, so a restore that fails halfway cannot start DocsGPT again and call it a success. - import_volume unpacks into the container's own filesystem first and clears the live volume only once the tar has come out whole, so a corrupt one leaves the volume as it was. - The backend and the worker stop while the archive is made and start again even if the dump fails, so the dump and the volume tars describe the same moment instead of drifting apart as ingestion writes. --- docs/content/Deploying/Docker-Deploying.mdx | 6 ++ docs/content/changelog.mdx | 3 +- docsgpt/deploy/backup.py | 34 ++++++++- docsgpt/deploy/commands.py | 35 +++++++--- docsgpt/deploy/docker.py | 7 +- tests/deploy/test_backup.py | 76 +++++++++++++++++++++ tests/deploy/test_docker.py | 5 +- 7 files changed, 150 insertions(+), 16 deletions(-) diff --git a/docs/content/Deploying/Docker-Deploying.mdx b/docs/content/Deploying/Docker-Deploying.mdx index 847c7453..ac835442 100644 --- a/docs/content/Deploying/Docker-Deploying.mdx +++ b/docs/content/Deploying/Docker-Deploying.mdx @@ -89,6 +89,12 @@ docsgpt backup # into /backups docsgpt backup --out /mnt/backups # somewhere else, e.g. a mounted disk ``` +While the archive is made, the backend and the worker stop and start again, so +the database dump and the files in the volumes describe the same moment; Postgres +itself keeps running. On a small install that pause is seconds, but count on it +if you run `docsgpt backup` from cron. The archive is written readable only by +the user who took it. + The archive does **not** include `.env`, because that file holds the install's secrets. `docsgpt backup --with-settings` puts it in, for when the archive itself is stored somewhere private. Keep `.env` safe separately otherwise: the diff --git a/docs/content/changelog.mdx b/docs/content/changelog.mdx index 2c5f9501..056fb101 100644 --- a/docs/content/changelog.mdx +++ b/docs/content/changelog.mdx @@ -18,7 +18,8 @@ request, and [Upgrading](/upgrading) covers the steps an existing deployment has `docsgpt backup` writes a dump of the database and a tar of each data volume into one archive, and `docsgpt restore ` puts them back. The settings file is left out unless `--with-settings` asks for it, since it holds the install's secrets, and a backup from a newer -DocsGPT is refused unless you pass `--force`. See +DocsGPT is refused unless you pass `--force`. The archive is written readable only by its owner, +and the backend and the worker pause while it is made so the dump and the volume tars match. See [Backups](/Deploying/Docker-Deploying#backups). ## 0.21.0 diff --git a/docsgpt/deploy/backup.py b/docsgpt/deploy/backup.py index 5b041c23..eb9dbfac 100644 --- a/docsgpt/deploy/backup.py +++ b/docsgpt/deploy/backup.py @@ -13,6 +13,7 @@ from __future__ import annotations import io import json +import os import tarfile import time from collections.abc import Mapping @@ -50,7 +51,10 @@ def write_archive( ) -> None: """Write the backup archive; the manifest is added last so a truncated file has none.""" path.parent.mkdir(parents=True, exist_ok=True) - with tarfile.open(path, "w:gz") as archive: + # 0600 from the start: the dump is the install's data, and --with-settings adds its secrets. + descriptor = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600) + os.fchmod(descriptor, 0o600) + with os.fdopen(descriptor, "wb") as handle, tarfile.open(fileobj=handle, mode="w:gz") as archive: archive.add(dump, arcname=DUMP) for name, tar_path in sorted(volume_tars.items()): archive.add(tar_path, arcname=volume_member(name)) @@ -81,6 +85,34 @@ def read_manifest(path: Path) -> dict: return manifest +def validate(path: Path, manifest: Mapping[str, object]) -> list[str]: + """The volumes to restore, once the archive is known to hold everything it declares. + + Checked before the stack is stopped: a damaged or hand-made archive must fail while DocsGPT is + still running, not after ``docker compose down``. Only the volumes a backup is made of are + accepted, so a manifest cannot name ``postgres_data`` and have it emptied on the way in. + """ + volumes = manifest.get("volumes") + if not isinstance(volumes, list) or not all(isinstance(name, str) for name in volumes): + raise DeployError(f"{path} is not a DocsGPT backup: its manifest does not list the volumes it holds") + unsupported = sorted(set(volumes) - set(DATA_VOLUMES)) + if unsupported: + raise DeployError( + f"{path} names volumes that are not part of a backup: {', '.join(unsupported)}. " + f"A DocsGPT backup holds {', '.join(DATA_VOLUMES)}." + ) + required = [DUMP, *(volume_member(name) for name in volumes)] + try: + with tarfile.open(path, "r:gz") as archive: + present = {member.name for member in archive.getmembers() if member.isfile()} + except (tarfile.TarError, OSError) as exc: + raise DeployError(f"{path} is not a DocsGPT backup: {exc}") from exc + missing = [name for name in required if name not in present] + if missing: + raise DeployError(f"{path} is missing {', '.join(missing)}, so there is nothing to restore from") + return list(volumes) + + def extract(path: Path, destination: Path) -> None: """Unpack the archive into ``destination`` (data filter: no paths outside it, no devices).""" with tarfile.open(path, "r:gz") as archive: diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index fb9c5e25..b3bf7b28 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -371,6 +371,26 @@ def backup(args, context: Optional[Context] = None) -> int: target = out_dir / backup_format.archive_name(taken_at) image = _stack_image(env) + print("Pausing the backend and the worker so the database and the files match ...") + context.docker.compose(directory, "stop", "backend", "worker") + try: + _write_backup(args, context, directory, env, target, image, taken_at) + finally: + print("Starting the backend and the worker again ...") + context.docker.compose(directory, "up", "-d", "backend", "worker") + + size = target.stat().st_size / 1_000_000 + print(f"\nBackup written to {target} ({size:.1f} MB)") + if args.with_settings: + print("It contains .env, so it holds this install's secrets: keep it somewhere private.") + else: + print("Settings are not in it; `docsgpt backup --with-settings` includes .env, secrets and all.") + print(f"Restore it with: docsgpt restore {target}") + return 0 + + +def _write_backup(args, context: Context, directory: Path, env, target: Path, image: str, taken_at) -> None: + """Dump the database and the data volumes into ``target``, with the writers stopped.""" with tempfile.TemporaryDirectory() as workspace: work = Path(workspace) dump = work / backup_format.DUMP @@ -398,15 +418,6 @@ def backup(args, context: Optional[Context] = None) -> int: settings = directory / ".env" if args.with_settings else None backup_format.write_archive(target, dump=dump, volume_tars=volume_tars, manifest=manifest, settings=settings) - size = target.stat().st_size / 1_000_000 - print(f"\nBackup written to {target} ({size:.1f} MB)") - if args.with_settings: - print("It contains .env, so it holds this install's secrets: keep it somewhere private.") - else: - print("Settings are not in it; `docsgpt backup --with-settings` includes .env, secrets and all.") - print(f"Restore it with: docsgpt restore {target}") - return 0 - def restore(args, context: Optional[Context] = None) -> int: """Put a backup's database and data volumes back over this install.""" @@ -414,6 +425,8 @@ def restore(args, context: Optional[Context] = None) -> int: archive = Path(args.archive).expanduser() manifest = backup_format.read_manifest(archive) backup_format.check_version(manifest, context.version, args.force) + # Everything the archive declares is checked here, while DocsGPT is still up. + volumes = backup_format.validate(archive, manifest) directory = stack.stack_dir(args.dir) env = _installed(directory) @@ -436,7 +449,7 @@ def restore(args, context: Optional[Context] = None) -> int: with tempfile.TemporaryDirectory() as workspace: work = Path(workspace) backup_format.extract(archive, work) - for name in manifest.get("volumes", []): + for name in volumes: tar_path = work / backup_format.volume_member(name) if not tar_path.is_file(): raise DeployError(f"{archive} is missing the {name} volume it says it contains") @@ -452,7 +465,7 @@ def restore(args, context: Optional[Context] = None) -> int: with dump.open("r", encoding="utf-8") as handle: context.docker.compose( directory, "exec", "-T", "postgres", - "psql", "--quiet", "-U", "docsgpt", "-d", "docsgpt", + "psql", "--quiet", "--set", "ON_ERROR_STOP=on", "-U", "docsgpt", "-d", "docsgpt", stdin=handle, ) diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py index 69cbe187..c44eb241 100644 --- a/docsgpt/deploy/docker.py +++ b/docsgpt/deploy/docker.py @@ -119,8 +119,11 @@ class Docker: """Replace ``volume``'s contents with the tar at ``source``; the volume is created when missing.""" with source.open("rb") as handle: self._run( - ["docker", "run", "--rm", "-i", "-v", f"{volume}:/data", image, - "sh", "-c", "find /data -mindepth 1 -delete && tar xf - -C /data"], + ["docker", "run", "--rm", "-i", "-v", f"{volume}:/data", image, "sh", "-c", + # Unpack into the container's own filesystem first: a truncated or corrupt tar must + # fail before the live volume is touched, not halfway through emptying it. + "set -e; rm -rf /stage; mkdir /stage; tar xf - -C /stage; " + "find /data -mindepth 1 -delete; tar cf - -C /stage . | tar xf - -C /data"], stdin=handle, ) diff --git a/tests/deploy/test_backup.py b/tests/deploy/test_backup.py index be43bfcf..371023a5 100644 --- a/tests/deploy/test_backup.py +++ b/tests/deploy/test_backup.py @@ -2,6 +2,7 @@ import io import json +import stat import tarfile import pytest @@ -26,6 +27,22 @@ def _manifest(path): return json.loads(archive.extractfile(backup_module.MANIFEST).read()) +def _rewrite(archive, dest, *, volumes=None, drop=()): + """The same archive with other volumes in its manifest, or with members left out.""" + with tarfile.open(archive, "r:gz") as source, tarfile.open(dest, "w:gz") as out: + for member in source.getmembers(): + if member.name in drop: + continue + body = source.extractfile(member).read() if member.isfile() else None + if member.name == backup_module.MANIFEST and volumes is not None: + manifest = json.loads(body.decode()) + manifest["volumes"] = volumes + body = json.dumps(manifest).encode() + member.size = len(body) + out.addfile(member, io.BytesIO(body) if body is not None else None) + return dest + + def _copy_with_version(archive, version, dest): """The same archive, with another DocsGPT version written into its manifest.""" with tarfile.open(archive, "r:gz") as source, tarfile.open(dest, "w:gz") as out: @@ -80,6 +97,40 @@ class TestBackup: assert dumps, docker.calls assert dumps[0][:3] == ["exec", "-T", "postgres"] + def test_the_archive_is_private_to_whoever_took_it(self, tmp_path): + """It can hold .env, and the dump is the install's data either way.""" + _installed(tmp_path) + out = tmp_path / "backups" + argv = ["backup", "--dir", str(tmp_path), "--out", str(out), "--with-settings"] + assert _run(argv, _context()) == 0 + archive = next(out.glob("*.tar.gz")) + assert stat.S_IMODE(archive.stat().st_mode) == 0o600 + + def test_the_writers_are_stopped_while_the_archive_is_made(self, tmp_path): + """Ingestion writes files and rows; a dump taken alongside live writes would not match them.""" + _installed(tmp_path) + docker = FakeDocker() + assert _run(["backup", "--dir", str(tmp_path), "--out", str(tmp_path / "b")], _context(docker)) == 0 + joined = [" ".join(args) for _, args in docker.calls] + stopped = next(index for index, call in enumerate(joined) if call.startswith("stop backend worker")) + dumped = next(index for index, call in enumerate(joined) if "pg_dump" in call) + started = next(index for index, call in enumerate(joined) if call.startswith("up -d backend worker")) + assert stopped < dumped < started, joined + + def test_the_writers_come_back_even_when_the_dump_fails(self, tmp_path): + _installed(tmp_path) + + class FailingDump(FakeDocker): + def compose(self, directory, *args, **kwargs): + if "pg_dump" in args: + raise DeployError("pg_dump exploded") + return super().compose(directory, *args, **kwargs) + + docker = FailingDump() + with pytest.raises(DeployError, match="pg_dump exploded"): + _run(["backup", "--dir", str(tmp_path), "--out", str(tmp_path / "b")], _context(docker)) + assert any(" ".join(args).startswith("up -d backend worker") for _, args in docker.calls), docker.calls + def test_without_an_install(self, tmp_path, capsys): assert _run(["backup", "--dir", str(tmp_path)], _context()) == 1 assert "docsgpt up" in capsys.readouterr().err @@ -122,6 +173,31 @@ class TestRestore: with pytest.raises(DeployError, match="does not exist"): _run(["restore", str(tmp_path / "nope.tar.gz"), "--dir", str(tmp_path), "--yes"], _context()) + def test_a_manifest_naming_another_volume_is_refused(self, tmp_path): + """import_volume empties what it is given, so a hand-made manifest must not name postgres_data.""" + crafted = _rewrite(self._backup(tmp_path), tmp_path / "crafted.tar.gz", volumes=["postgres_data"]) + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + with pytest.raises(DeployError, match="not part of a backup"): + _run(["restore", str(crafted), "--dir", str(tmp_path), "--yes"], _context(docker)) + assert docker.calls == [], "nothing was stopped" + assert docker.volume_ops == [], "and nothing was touched" + + def test_a_damaged_archive_fails_before_the_stack_is_stopped(self, tmp_path): + """Discovering a missing member after `compose down` would leave DocsGPT down for nothing.""" + without_dump = _rewrite(self._backup(tmp_path), tmp_path / "nodump.tar.gz", drop=(backup_module.DUMP,)) + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + with pytest.raises(DeployError, match="nothing to restore"): + _run(["restore", str(without_dump), "--dir", str(tmp_path), "--yes"], _context(docker)) + assert docker.calls == [], "DocsGPT is still running" + + def test_psql_stops_at_the_first_failing_statement(self, tmp_path): + """Without it psql runs on after an error and a half-restored database looks like success.""" + archive = self._backup(tmp_path) + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + assert _run(["restore", str(archive), "--dir", str(tmp_path), "--yes"], _context(docker)) == 0 + psql = next(args for _, args in docker.calls if "psql" in args) + assert "ON_ERROR_STOP=on" in psql + def test_a_newer_backup_is_refused_without_force(self, tmp_path): """Restoring a 0.22 backup into 0.21 would hand an older schema newer data.""" newer = _copy_with_version(self._backup(tmp_path), "0.22.0", tmp_path / "newer.tar.gz") diff --git a/tests/deploy/test_docker.py b/tests/deploy/test_docker.py index 0ae2dfe4..61c15dd8 100644 --- a/tests/deploy/test_docker.py +++ b/tests/deploy/test_docker.py @@ -129,7 +129,10 @@ class TestVolumes: assert "docsgpt_inputs:/data" in args assert args[-2] == "-c" assert "tar xf - -C /data" in args[-1] - assert "find /data -mindepth 1 -delete" in args[-1], "old contents go first" + command = args[-1] + assert command.index("tar xf - -C /stage") < command.index("find /data -mindepth 1 -delete"), ( + "the incoming tar is unpacked outside the volume first, so a corrupt one leaves it alone" + ) assert runner.streams[0][1] is not None, "the tar is fed in on stdin" From 434763649629eae7ae0de14a7350bebfa36875a7 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:01:39 +0100 Subject: [PATCH 031/130] fix: stage a restored volume under /tmp, where the image can write The image does not run as root, so the staging directory could not be created at the container root: `mkdir /stage` failed with permission denied and every restore would have failed. It goes under /tmp now, and the command is built as one string instead of concatenated pieces inside the argument list. Checked against a real volume and the published image: a truncated tar fails and leaves the volume exactly as it was, and a whole one restores it. --- docsgpt/deploy/docker.py | 14 +++++++++----- tests/deploy/test_docker.py | 3 ++- 2 files changed, 11 insertions(+), 6 deletions(-) diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py index c44eb241..2b3dcdfa 100644 --- a/docsgpt/deploy/docker.py +++ b/docsgpt/deploy/docker.py @@ -117,13 +117,17 @@ class Docker: def import_volume(self, volume: str, source: Path, image: str) -> None: """Replace ``volume``'s contents with the tar at ``source``; the volume is created when missing.""" + # Unpack into the container's own filesystem first, so a truncated or corrupt tar fails + # before the live volume is touched rather than halfway through emptying it. It goes under + # /tmp because the image does not run as root and cannot write to /. + staging = "/tmp/docsgpt-restore" + script = ( + f"set -e; rm -rf {staging}; mkdir -p {staging}; tar xf - -C {staging}; " + f"find /data -mindepth 1 -delete; tar cf - -C {staging} . | tar xf - -C /data" + ) with source.open("rb") as handle: self._run( - ["docker", "run", "--rm", "-i", "-v", f"{volume}:/data", image, "sh", "-c", - # Unpack into the container's own filesystem first: a truncated or corrupt tar must - # fail before the live volume is touched, not halfway through emptying it. - "set -e; rm -rf /stage; mkdir /stage; tar xf - -C /stage; " - "find /data -mindepth 1 -delete; tar cf - -C /stage . | tar xf - -C /data"], + ["docker", "run", "--rm", "-i", "-v", f"{volume}:/data", image, "sh", "-c", script], stdin=handle, ) diff --git a/tests/deploy/test_docker.py b/tests/deploy/test_docker.py index 61c15dd8..0f1a3e8b 100644 --- a/tests/deploy/test_docker.py +++ b/tests/deploy/test_docker.py @@ -130,7 +130,8 @@ class TestVolumes: assert args[-2] == "-c" assert "tar xf - -C /data" in args[-1] command = args[-1] - assert command.index("tar xf - -C /stage") < command.index("find /data -mindepth 1 -delete"), ( + assert "mkdir -p /tmp/" in command, "the image does not run as root, so staging goes under /tmp" + assert command.index("tar xf - -C /tmp/") < command.index("find /data -mindepth 1 -delete"), ( "the incoming tar is unpacked outside the volume first, so a corrupt one leaves it alone" ) assert runner.streams[0][1] is not None, "the tar is fed in on stdin" From e269743bf878177249fd16b859d1a87affadc327 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:03:48 +0100 Subject: [PATCH 032/130] fix: start DocsGPT again when a restore fails after the stack is down Validating the archive catches a damaged one while DocsGPT is still up, but a well-formed archive can still hold a corrupt volume tar or a dump statement psql refuses, and those only surface once the stack is down. The work after the shutdown now runs inside an error boundary that starts the stack again before the failure is reported, so a failed restore never leaves the install stopped. --- docsgpt/deploy/commands.py | 56 ++++++++++++++++++++++--------------- tests/deploy/test_backup.py | 15 ++++++++++ 2 files changed, 49 insertions(+), 22 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index b3bf7b28..fab64d45 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -419,6 +419,32 @@ def _write_backup(args, context: Context, directory: Path, env, target: Path, im backup_format.write_archive(target, dump=dump, volume_tars=volume_tars, manifest=manifest, settings=settings) +def _restore_data(context: Context, directory: Path, archive: Path, volumes: list, image: str) -> None: + """Put the archive's volumes and database back, with the stack stopped.""" + with tempfile.TemporaryDirectory() as workspace: + work = Path(workspace) + backup_format.extract(archive, work) + for name in volumes: + tar_path = work / backup_format.volume_member(name) + if not tar_path.is_file(): + raise DeployError(f"{archive} is missing the {name} volume it says it contains") + print(f"Restoring the {name} volume ...") + context.docker.import_volume(f"{PROJECT}_{name}", tar_path, image) + + dump = work / backup_format.DUMP + if not dump.is_file(): + raise DeployError(f"{archive} is missing its database dump") + print("Starting the database ...") + context.docker.compose(directory, "up", "-d", "--wait", "postgres") + print("Restoring the database ...") + with dump.open("r", encoding="utf-8") as handle: + context.docker.compose( + directory, "exec", "-T", "postgres", + "psql", "--quiet", "--set", "ON_ERROR_STOP=on", "-U", "docsgpt", "-d", "docsgpt", + stdin=handle, + ) + + def restore(args, context: Optional[Context] = None) -> int: """Put a backup's database and data volumes back over this install.""" context = context or Context.default(args) @@ -446,28 +472,14 @@ def restore(args, context: Optional[Context] = None) -> int: print("Stopping the stack ...") context.docker.compose(directory, *_EVERY_PROFILE, "down") - with tempfile.TemporaryDirectory() as workspace: - work = Path(workspace) - backup_format.extract(archive, work) - for name in volumes: - tar_path = work / backup_format.volume_member(name) - if not tar_path.is_file(): - raise DeployError(f"{archive} is missing the {name} volume it says it contains") - print(f"Restoring the {name} volume ...") - context.docker.import_volume(f"{PROJECT}_{name}", tar_path, image) - - dump = work / backup_format.DUMP - if not dump.is_file(): - raise DeployError(f"{archive} is missing its database dump") - print("Starting the database ...") - context.docker.compose(directory, "up", "-d", "--wait", "postgres") - print("Restoring the database ...") - with dump.open("r", encoding="utf-8") as handle: - context.docker.compose( - directory, "exec", "-T", "postgres", - "psql", "--quiet", "--set", "ON_ERROR_STOP=on", "-U", "docsgpt", "-d", "docsgpt", - stdin=handle, - ) + try: + _restore_data(context, directory, archive, volumes, image) + except BaseException: + # The stack is down by now. A corrupt payload inside an otherwise well-formed archive, or a + # statement psql refuses, must not leave DocsGPT stopped: start it again, then report. + print("The restore failed. Starting DocsGPT again ...", file=sys.stderr) + context.docker.compose(directory, "up", "-d", "--remove-orphans", check=False) + raise print("Starting DocsGPT ...") context.docker.compose(directory, "up", "-d", "--remove-orphans") diff --git a/tests/deploy/test_backup.py b/tests/deploy/test_backup.py index 371023a5..ee6d9889 100644 --- a/tests/deploy/test_backup.py +++ b/tests/deploy/test_backup.py @@ -198,6 +198,21 @@ class TestRestore: psql = next(args for _, args in docker.calls if "psql" in args) assert "ON_ERROR_STOP=on" in psql + def test_a_failure_after_the_stack_is_down_starts_it_again(self, tmp_path, capsys): + """A corrupt payload only shows up once the stack is down; it must not be left stopped.""" + archive = self._backup(tmp_path) + + class FailingImport(FakeDocker): + def import_volume(self, volume, source, image): + raise DeployError("tar: unexpected EOF in archive") + + docker = FailingImport(volumes={"docsgpt_postgres_data"}) + with pytest.raises(DeployError, match="unexpected EOF"): + _run(["restore", str(archive), "--dir", str(tmp_path), "--yes"], _context(docker)) + joined = [" ".join(args) for _, args in docker.calls] + assert any(call.startswith("up -d --remove-orphans") for call in joined), joined + assert "Starting DocsGPT again" in capsys.readouterr().err + def test_a_newer_backup_is_refused_without_force(self, tmp_path): """Restoring a 0.22 backup into 0.21 would hand an older schema newer data.""" newer = _copy_with_version(self._backup(tmp_path), "0.22.0", tmp_path / "newer.tar.gz") From 8cfa3fbd18405bb43ee79688aa88b50b367d9737 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:05:22 +0100 Subject: [PATCH 033/130] fix: let the container pick the staging directory for a restored volume A fixed path under /tmp was both a guess about what the image can write to and a temp-file smell that Bandit flags. The container makes the directory itself with mktemp -d and removes it afterwards. --- docsgpt/deploy/docker.py | 13 +++++++------ tests/deploy/test_docker.py | 4 ++-- 2 files changed, 9 insertions(+), 8 deletions(-) diff --git a/docsgpt/deploy/docker.py b/docsgpt/deploy/docker.py index 2b3dcdfa..cba24375 100644 --- a/docsgpt/deploy/docker.py +++ b/docsgpt/deploy/docker.py @@ -117,13 +117,14 @@ class Docker: def import_volume(self, volume: str, source: Path, image: str) -> None: """Replace ``volume``'s contents with the tar at ``source``; the volume is created when missing.""" - # Unpack into the container's own filesystem first, so a truncated or corrupt tar fails - # before the live volume is touched rather than halfway through emptying it. It goes under - # /tmp because the image does not run as root and cannot write to /. - staging = "/tmp/docsgpt-restore" + # Unpack into a throwaway directory inside the container first, so a truncated or corrupt + # tar fails before the live volume is touched rather than halfway through emptying it. The + # container makes the directory itself: the image does not run as root, and a fixed path + # would be both a guess about what is writable and a temp-file smell. script = ( - f"set -e; rm -rf {staging}; mkdir -p {staging}; tar xf - -C {staging}; " - f"find /data -mindepth 1 -delete; tar cf - -C {staging} . | tar xf - -C /data" + 'set -e; stage=$(mktemp -d); tar xf - -C "$stage"; ' + 'find /data -mindepth 1 -delete; tar cf - -C "$stage" . | tar xf - -C /data; ' + 'rm -rf "$stage"' ) with source.open("rb") as handle: self._run( diff --git a/tests/deploy/test_docker.py b/tests/deploy/test_docker.py index 0f1a3e8b..f42d57c9 100644 --- a/tests/deploy/test_docker.py +++ b/tests/deploy/test_docker.py @@ -130,8 +130,8 @@ class TestVolumes: assert args[-2] == "-c" assert "tar xf - -C /data" in args[-1] command = args[-1] - assert "mkdir -p /tmp/" in command, "the image does not run as root, so staging goes under /tmp" - assert command.index("tar xf - -C /tmp/") < command.index("find /data -mindepth 1 -delete"), ( + assert "mktemp -d" in command, "the container picks the staging directory, not a fixed path" + assert command.index('tar xf - -C "$stage"') < command.index("find /data -mindepth 1 -delete"), ( "the incoming tar is unpacked outside the volume first, so a corrupt one leaves it alone" ) assert runner.streams[0][1] is not None, "the tar is fed in on stdin" From f90c442a41e743477d1ebd0ea94b84ba8a29b638 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:20:40 +0100 Subject: [PATCH 034/130] fix: cover the shutdown calls and check every volume before replacing one Both shutdown calls sat outside the recovery that undoes them: a `compose down` that failed partway left the stack down, and a `compose stop` that failed left the backend and worker stopped. Each now runs inside its own try. A restore also replaced volumes one at a time, checking each tar as it reached it, so a damaged third payload was found with the first two already swapped in. Every declared tar is read through first, and the imports start only once they all come out whole. --- docsgpt/deploy/backup.py | 10 ++++++++ docsgpt/deploy/commands.py | 15 +++++++---- tests/deploy/test_backup.py | 50 +++++++++++++++++++++++++++++++++++-- 3 files changed, 68 insertions(+), 7 deletions(-) diff --git a/docsgpt/deploy/backup.py b/docsgpt/deploy/backup.py index eb9dbfac..e823b9ee 100644 --- a/docsgpt/deploy/backup.py +++ b/docsgpt/deploy/backup.py @@ -113,6 +113,16 @@ def validate(path: Path, manifest: Mapping[str, object]) -> list[str]: return list(volumes) +def check_volume_tar(name: str, path: Path) -> None: + """Read a volume tar through, so a damaged one is found before any volume is replaced.""" + try: + with tarfile.open(path, "r:*") as archive: + for _ in archive: + pass + except (tarfile.TarError, OSError) as exc: + raise DeployError(f"the {name} volume in this backup is damaged: {exc}") from exc + + def extract(path: Path, destination: Path) -> None: """Unpack the archive into ``destination`` (data filter: no paths outside it, no devices).""" with tarfile.open(path, "r:gz") as archive: diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index fab64d45..5aa5fa42 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -371,9 +371,9 @@ def backup(args, context: Optional[Context] = None) -> int: target = out_dir / backup_format.archive_name(taken_at) image = _stack_image(env) - print("Pausing the backend and the worker so the database and the files match ...") - context.docker.compose(directory, "stop", "backend", "worker") try: + print("Pausing the backend and the worker so the database and the files match ...") + context.docker.compose(directory, "stop", "backend", "worker") _write_backup(args, context, directory, env, target, image, taken_at) finally: print("Starting the backend and the worker again ...") @@ -424,10 +424,16 @@ def _restore_data(context: Context, directory: Path, archive: Path, volumes: lis with tempfile.TemporaryDirectory() as workspace: work = Path(workspace) backup_format.extract(archive, work) + # Read every volume tar through before replacing any of them: a damaged third tar must not + # be discovered with the first two already swapped in. + tars = {} for name in volumes: tar_path = work / backup_format.volume_member(name) if not tar_path.is_file(): raise DeployError(f"{archive} is missing the {name} volume it says it contains") + backup_format.check_volume_tar(name, tar_path) + tars[name] = tar_path + for name, tar_path in tars.items(): print(f"Restoring the {name} volume ...") context.docker.import_volume(f"{PROJECT}_{name}", tar_path, image) @@ -469,10 +475,9 @@ def restore(args, context: Optional[Context] = None) -> int: return 1 image = _stack_image(env) - print("Stopping the stack ...") - context.docker.compose(directory, *_EVERY_PROFILE, "down") - try: + print("Stopping the stack ...") + context.docker.compose(directory, *_EVERY_PROFILE, "down") _restore_data(context, directory, archive, volumes, image) except BaseException: # The stack is down by now. A corrupt payload inside an otherwise well-formed archive, or a diff --git a/tests/deploy/test_backup.py b/tests/deploy/test_backup.py index ee6d9889..9c80237a 100644 --- a/tests/deploy/test_backup.py +++ b/tests/deploy/test_backup.py @@ -4,6 +4,7 @@ import io import json import stat import tarfile +from pathlib import Path import pytest @@ -27,13 +28,16 @@ def _manifest(path): return json.loads(archive.extractfile(backup_module.MANIFEST).read()) -def _rewrite(archive, dest, *, volumes=None, drop=()): - """The same archive with other volumes in its manifest, or with members left out.""" +def _rewrite(archive, dest, *, volumes=None, drop=(), corrupt=()): + """The same archive with other volumes in its manifest, or with members left out or damaged.""" with tarfile.open(archive, "r:gz") as source, tarfile.open(dest, "w:gz") as out: for member in source.getmembers(): if member.name in drop: continue body = source.extractfile(member).read() if member.isfile() else None + if member.name in corrupt: + body = body[:60] + member.size = len(body) if member.name == backup_module.MANIFEST and volumes is not None: manifest = json.loads(body.decode()) manifest["volumes"] = volumes @@ -117,6 +121,22 @@ class TestBackup: started = next(index for index, call in enumerate(joined) if call.startswith("up -d backend worker")) assert stopped < dumped < started, joined + def test_the_writers_come_back_even_when_stopping_them_fails(self, tmp_path): + """A stop that fails partway must not leave the backend or the worker down.""" + _installed(tmp_path) + + class FailingStop(FakeDocker): + def compose(self, directory, *args, **kwargs): + if args[:1] == ("stop",): + self.calls.append((Path(directory), list(args))) + raise DeployError("compose stop exploded") + return super().compose(directory, *args, **kwargs) + + docker = FailingStop() + with pytest.raises(DeployError, match="compose stop exploded"): + _run(["backup", "--dir", str(tmp_path), "--out", str(tmp_path / "b")], _context(docker)) + assert any(" ".join(args).startswith("up -d backend worker") for _, args in docker.calls), docker.calls + def test_the_writers_come_back_even_when_the_dump_fails(self, tmp_path): _installed(tmp_path) @@ -213,6 +233,32 @@ class TestRestore: assert any(call.startswith("up -d --remove-orphans") for call in joined), joined assert "Starting DocsGPT again" in capsys.readouterr().err + def test_a_shutdown_that_fails_still_starts_the_stack_again(self, tmp_path): + """`compose down` can fail with containers already stopped; the stack must not stay down.""" + archive = self._backup(tmp_path) + + class FailingDown(FakeDocker): + def compose(self, directory, *args, **kwargs): + if "down" in args: + self.calls.append((Path(directory), list(args))) + raise DeployError("compose down exploded") + return super().compose(directory, *args, **kwargs) + + docker = FailingDown(volumes={"docsgpt_postgres_data"}) + with pytest.raises(DeployError, match="compose down exploded"): + _run(["restore", str(archive), "--dir", str(tmp_path), "--yes"], _context(docker)) + joined = [" ".join(args) for _, args in docker.calls] + assert any(call.startswith("up -d --remove-orphans") for call in joined), joined + + def test_a_damaged_volume_tar_is_found_before_any_volume_is_replaced(self, tmp_path): + """vectors sorts last, so a per-volume check would already have swapped indexes and inputs.""" + archive = _rewrite(self._backup(tmp_path), tmp_path / "bad-volume.tar.gz", + corrupt=(backup_module.volume_member("vectors"),)) + docker = FakeDocker(volumes={"docsgpt_postgres_data"}) + with pytest.raises(DeployError, match="damaged"): + _run(["restore", str(archive), "--dir", str(tmp_path), "--yes"], _context(docker)) + assert [op for op in docker.volume_ops if op[0] == "import"] == [], docker.volume_ops + def test_a_newer_backup_is_refused_without_force(self, tmp_path): """Restoring a 0.22 backup into 0.21 would hand an older schema newer data.""" newer = _copy_with_version(self._backup(tmp_path), "0.22.0", tmp_path / "newer.tar.gz") From f9f52e99f0f170c670735e65c570f75f9e7b2c31 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:35:14 +0100 Subject: [PATCH 035/130] fix: make a native install take effect on systemd and stay out of Docker's way From review of #2800: - systemd `enable --now` starts nothing when the unit is already active, so a second `up --native` kept the old ExecStart and left the API on its previous port. start enables and then restarts, as the launchd path already did by booting the job out first. - An explicit `home` now wins over XDG_CONFIG_HOME, which is what callers pass it for. - The ExecStart program must be a real executable: when `docsgpt` is not on PATH, sys.argv[0] is accepted only if it can be run, and otherwise the failure is raised before any unit is written. - `up --native` over a directory holding a Docker install now refuses and says how to proceed, instead of starting native services beside containers that down, status and uninstall would no longer see. - Docs: without a terminal only --postgres-uri is required, and the Windows fallback names `docsgpt beat`, which the worker cannot embed there. SystemdServices was the least covered part of the module and cannot be run on this machine, so it now has tests for install, start, stop, remove, is_running and a failing systemctl. --- docs/content/Deploying/Pip-Install.mdx | 4 +- docsgpt/deploy/commands.py | 20 +++++- docsgpt/deploy/native.py | 13 +++- tests/deploy/test_native.py | 99 ++++++++++++++++++++++++++ 4 files changed, 130 insertions(+), 6 deletions(-) diff --git a/docs/content/Deploying/Pip-Install.mdx b/docs/content/Deploying/Pip-Install.mdx index 04574490..64ca8898 100644 --- a/docs/content/Deploying/Pip-Install.mdx +++ b/docs/content/Deploying/Pip-Install.mdx @@ -109,7 +109,7 @@ docsgpt up --native \ --redis-url redis://localhost:6379 ``` -Without a terminal both flags are required; with one, it asks for them and for the model provider. It writes the same `.env` a Docker install uses (minus the image settings), generates `INTERNAL_KEY` and `JWT_SECRET_KEY` on the first run, applies the migrations, then starts `docsgpt-api` and `docsgpt-worker` and waits for the API to answer. +Without a terminal only `--postgres-uri` is required, since Redis defaults to `redis://localhost:6379`; with one, it asks for both and for the model provider. It writes the same `.env` a Docker install uses (minus the image settings), generates `INTERNAL_KEY` and `JWT_SECRET_KEY` on the first run, applies the migrations, then starts `docsgpt-api` and `docsgpt-worker` and waits for the API to answer. One Redis URL covers all three uses: the Celery broker, its result backend and the cache go on databases 0, 1 and 2 of it. Name a database in the URL and the three start there instead, so `redis://localhost:6379/5` puts them on 5, 6 and 7 — that is how you share a Redis that already holds something else. @@ -125,7 +125,7 @@ The same commands manage it: `uninstall` never touches the database or Redis: they were yours to begin with. - Windows has neither launchd nor systemd, so native mode is macOS and Linux only. On Windows, run DocsGPT on Docker with `docsgpt up`, or start `docsgpt api` and `docsgpt worker` yourself. + Windows has neither launchd nor systemd, so native mode is macOS and Linux only. On Windows, run DocsGPT on Docker with `docsgpt up`, or start `docsgpt api`, `docsgpt worker` and `docsgpt beat` yourself — the worker cannot run the scheduler in-process there, so `docsgpt beat` has to run alongside it for scheduled tasks to fire. ## Upgrade diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 9bd2db78..24504a7d 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -218,6 +218,17 @@ def _redis_urls(base: str) -> dict[str, str]: } +def _own_executable() -> str: + """The docsgpt command to put in the service files, or a DeployError naming what to install.""" + candidate = Path(sys.argv[0]).resolve() + if candidate.is_file() and os.access(candidate, os.X_OK): + return str(candidate) + raise DeployError( + "the docsgpt command is not on PATH, so the services would have nothing to run. " + "Install it with `uv tool install docsgpt` (or `pip install docsgpt`) and run this again." + ) + + def _native_up(args, context: Context, directory: Path) -> int: """Run the API and the worker as services on this machine, against an existing Postgres and Redis.""" services = context.service_manager() @@ -263,7 +274,7 @@ def _native_up(args, context: Context, directory: Path) -> int: (directory / "logs").mkdir(exist_ok=True) envfile.update(env_path, updates) - executable = shutil.which("docsgpt") or sys.argv[0] + executable = shutil.which("docsgpt") or _own_executable() print("Applying database migrations ...") # The child reads the stack's settings, not a .env in whatever directory this was run from. stack_env = {"DOCSGPT_HOME": str(directory), "DOCSGPT_ENV_FILE": str(env_path)} @@ -303,6 +314,13 @@ def up(args, context: Optional[Context] = None) -> int: context = context or Context.default(args) directory = stack.stack_dir(args.dir) if getattr(args, "native", False) or _mode(directory) == "native": + if _mode(directory) != "native" and (directory / stack.COMPOSE_FILE).is_file(): + raise DeployError( + f"{directory} holds a Docker install. Native services would run beside its containers " + "on the same port, and down, status and uninstall would stop seeing them. Stop it " + "first with `docsgpt down` (and `docsgpt uninstall` to remove it), or pass --dir to " + "put the native install somewhere else." + ) return _native_up(args, context, directory) env_path = directory / ".env" record_path = directory / stack.RECORD_FILE diff --git a/docsgpt/deploy/native.py b/docsgpt/deploy/native.py index 0ef0cf19..2c6376f2 100644 --- a/docsgpt/deploy/native.py +++ b/docsgpt/deploy/native.py @@ -152,8 +152,12 @@ class SystemdServices: def __init__(self, runner=subprocess.run, home: Optional[Path] = None) -> None: self._run = runner - base = os.environ.get("XDG_CONFIG_HOME") - root = Path(base) if base else (home or Path.home()) / ".config" + # An explicit home wins: it is what a caller passes to redirect the unit directory. + if home is not None: + root = home / ".config" + else: + base = os.environ.get("XDG_CONFIG_HOME") + root = Path(base) if base else Path.home() / ".config" self.directory = root / "systemd" / "user" def _unit_file(self, name: str) -> Path: @@ -171,7 +175,10 @@ class SystemdServices: self._systemctl("daemon-reload") def start(self, name: str) -> None: - self._systemctl("enable", "--now", f"{name}.service", check=True) + # restart, not `enable --now`: --now starts nothing when the unit is already active, so a + # reinstalled unit would keep running with the ExecStart and environment it started with. + self._systemctl("enable", f"{name}.service", check=True) + self._systemctl("restart", f"{name}.service", check=True) def stop(self, name: str) -> None: self._systemctl("stop", f"{name}.service") diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 63646d98..630e0097 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -2,6 +2,7 @@ import json import plistlib +import sys import types import pytest @@ -131,6 +132,21 @@ class TestNativeUp: assert env["API_URL"] == "http://127.0.0.1:7099" assert services.units["docsgpt-api"].arguments[-1] == "7099" + def test_it_refuses_to_run_beside_a_docker_install(self, tmp_path): + """Native services on the same port would orphan the containers from down/status/uninstall.""" + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context()) == 0 + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + with pytest.raises(DeployError, match="holds a Docker install"): + _run(argv, _native_context()) + + def test_services_are_not_written_without_a_docsgpt_command(self, tmp_path, monkeypatch): + """A unit pointing at something unexecutable would fail to exec with only a health check to show it.""" + monkeypatch.setattr("docsgpt.deploy.commands.shutil.which", lambda name: None) + monkeypatch.setattr(sys, "argv", [str(tmp_path / "not-a-program")]) + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + with pytest.raises(DeployError, match="not on PATH"): + _run(argv, _native_context()) + def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): _run(["up", "--native", "--dir", str(tmp_path), "--yes"], _native_context()) @@ -252,6 +268,89 @@ class TestLaunchd: assert self._services(FakeLaunchctl(state="running"), tmp_path).is_running("docsgpt-api") is True +class FakeSystemctl: + """systemctl, recording what it was asked to do.""" + + def __init__(self, active=True): + self.calls = [] + self.active = active + + def __call__(self, args, capture_output=False, text=False, check=False): + self.calls.append(args) + code = 0 if (args[2] != "is-active" or self.active) else 3 + return types.SimpleNamespace(returncode=code, stdout="", stderr="") + + @property + def verbs(self): + return [call[2] for call in self.calls] + + +class TestSystemd: + def _services(self, systemctl, tmp_path): + return native.SystemdServices(runner=systemctl, home=tmp_path) + + def test_install_writes_the_unit_and_reloads(self, tmp_path): + systemctl = FakeSystemctl() + services = self._services(systemctl, tmp_path) + services.install(native.units_for(tmp_path, "/venv/bin/docsgpt", 7091, tmp_path)[0]) + unit = tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service" + assert "ExecStart=/venv/bin/docsgpt api" in unit.read_text() + assert systemctl.verbs == ["daemon-reload"] + + def test_start_restarts_so_a_rewritten_unit_takes_effect(self, tmp_path): + """`enable --now` starts nothing when the unit is already active: it would keep the old ExecStart.""" + systemctl = FakeSystemctl() + self._services(systemctl, tmp_path).start("docsgpt-api") + assert "restart" in systemctl.verbs, systemctl.verbs + assert systemctl.verbs.index("enable") < systemctl.verbs.index("restart") + assert "--now" not in [argument for call in systemctl.calls for argument in call] + + def test_stop_and_remove(self, tmp_path): + systemctl = FakeSystemctl() + services = self._services(systemctl, tmp_path) + services.install(native.units_for(tmp_path, "/venv/bin/docsgpt", 7091, tmp_path)[0]) + unit = tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service" + services.stop("docsgpt-api") + assert unit.is_file(), "stopping keeps the unit" + services.remove("docsgpt-api") + assert not unit.exists() + assert systemctl.verbs[-1] == "daemon-reload", "systemd is told the unit is gone" + + def test_is_running_asks_systemd(self, tmp_path): + assert self._services(FakeSystemctl(active=True), tmp_path).is_running("docsgpt-api") is True + assert self._services(FakeSystemctl(active=False), tmp_path).is_running("docsgpt-api") is False + + def test_an_explicit_home_wins_over_xdg_config_home(self, tmp_path, monkeypatch): + """home is how a caller redirects the unit directory; the environment must not override it.""" + monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path / "xdg")) + services = native.SystemdServices(runner=FakeSystemctl(), home=tmp_path) + assert services.directory == tmp_path / ".config" / "systemd" / "user" + + def test_xdg_config_home_is_used_when_no_home_is_given(self, tmp_path, monkeypatch): + monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path / "xdg")) + services = native.SystemdServices(runner=FakeSystemctl()) + assert services.directory == tmp_path / "xdg" / "systemd" / "user" + + def test_a_failing_systemctl_is_reported(self, tmp_path): + class Failing(FakeSystemctl): + def __call__(self, args, **kwargs): + self.calls.append(args) + return types.SimpleNamespace(returncode=1, stdout="", stderr="Failed to enable unit") + + with pytest.raises(DeployError, match="Failed to enable unit"): + self._services(Failing(), tmp_path).start("docsgpt-api") + + +class TestPlatformChoice: + def test_each_platform_gets_its_service_manager(self): + assert isinstance(native.services_for_platform("darwin"), native.LaunchdServices) + assert isinstance(native.services_for_platform("linux"), native.SystemdServices) + + def test_windows_says_what_to_do_instead(self): + with pytest.raises(DeployError, match="Windows"): + native.services_for_platform("win32") + + class TestUnitFiles: def _unit(self, tmp_path): return native.Unit( From a96734da07f6792cffc0c7312b89b94a9e8e9c3f Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:45:54 +0100 Subject: [PATCH 036/130] fix: run the module when the docsgpt command is not on PATH Refusing to write the service units when `docsgpt` is not on PATH was wrong. A package installed in a virtualenv is runnable whether or not its console script is on PATH, and CI runs pytest as `python -m pytest`, where argv[0] is a module file: the refusal failed thirteen native tests there. The launcher now prefers the command on PATH, resolved to an absolute path since PATH can hold relative entries, then an argv[0] that can be executed, and otherwise this interpreter with `-m docsgpt`, which works wherever the package is importable. `python -m docsgpt` became an entrypoint of its own and has a test that runs it. --- docsgpt/__main__.py | 8 ++++++++ docsgpt/deploy/commands.py | 25 ++++++++++++++++++------- docsgpt/deploy/native.py | 11 +++++++---- tests/deploy/test_native.py | 19 ++++++++++--------- tests/test_cli.py | 10 ++++++++++ 5 files changed, 53 insertions(+), 20 deletions(-) create mode 100644 docsgpt/__main__.py diff --git a/docsgpt/__main__.py b/docsgpt/__main__.py new file mode 100644 index 00000000..397614de --- /dev/null +++ b/docsgpt/__main__.py @@ -0,0 +1,8 @@ +"""``python -m docsgpt`` runs what the ``docsgpt`` command runs.""" + +import sys + +from docsgpt.cli import main + +if __name__ == "__main__": + sys.exit(main()) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 24504a7d..0988044d 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -218,13 +218,24 @@ def _redis_urls(base: str) -> dict[str, str]: } -def _own_executable() -> str: - """The docsgpt command to put in the service files, or a DeployError naming what to install.""" +def _native_launcher() -> list[str]: + """How the service files start DocsGPT: the command on PATH, or this interpreter and the module. + + A unit has to name something that can be executed, and ``sys.argv[0]`` often cannot be: under + ``python -m docsgpt``, or pytest, it is a module file. Falling back to the running interpreter + works wherever the package is importable, which it must be to have got here. + """ + found = shutil.which("docsgpt") + if found: + # Absolute: PATH can hold relative entries, and a unit file needs a program it can exec. + return [str(Path(found).resolve())] candidate = Path(sys.argv[0]).resolve() if candidate.is_file() and os.access(candidate, os.X_OK): - return str(candidate) + return [str(candidate)] + if sys.executable: + return [sys.executable, "-m", "docsgpt"] raise DeployError( - "the docsgpt command is not on PATH, so the services would have nothing to run. " + "could not work out how to start docsgpt for the services. " "Install it with `uv tool install docsgpt` (or `pip install docsgpt`) and run this again." ) @@ -274,14 +285,14 @@ def _native_up(args, context: Context, directory: Path) -> int: (directory / "logs").mkdir(exist_ok=True) envfile.update(env_path, updates) - executable = shutil.which("docsgpt") or _own_executable() + launcher = _native_launcher() print("Applying database migrations ...") # The child reads the stack's settings, not a .env in whatever directory this was run from. stack_env = {"DOCSGPT_HOME": str(directory), "DOCSGPT_ENV_FILE": str(env_path)} - if context.run([executable, "migrate"], stack_env) != 0: + if context.run([*launcher, "migrate"], stack_env) != 0: raise DeployError("`docsgpt migrate` failed; check the database URL and that the server is reachable.") - for unit in native.units_for(directory, executable, port, directory): + for unit in native.units_for(directory, launcher, port, directory): services.install(unit) for name in native.SERVICES: print(f"Starting {name} ...") diff --git a/docsgpt/deploy/native.py b/docsgpt/deploy/native.py index 2c6376f2..8191a33f 100644 --- a/docsgpt/deploy/native.py +++ b/docsgpt/deploy/native.py @@ -204,20 +204,23 @@ def services_for_platform(platform: str = sys.platform): ) -def units_for(stack_directory: Path, executable: str, port: int, home: Path) -> list[Unit]: - """The API and worker services for a native install in ``stack_directory``.""" +def units_for(stack_directory: Path, launcher: list[str], port: int, home: Path) -> list[Unit]: + """The API and worker services for a native install in ``stack_directory``. + + ``launcher`` is how DocsGPT is started: the ``docsgpt`` command, or an interpreter and ``-m``. + """ environment = {"DOCSGPT_HOME": str(home)} return [ Unit( name=API_SERVICE, - arguments=[executable, "api", "--host", "127.0.0.1", "--port", str(port)], + arguments=[*launcher, "api", "--host", "127.0.0.1", "--port", str(port)], environment=dict(environment), working_directory=str(stack_directory), log_file=str(stack_directory / "logs" / "api.log"), ), Unit( name=WORKER_SERVICE, - arguments=[executable, "worker"], + arguments=[*launcher, "worker"], environment=dict(environment), working_directory=str(stack_directory), log_file=str(stack_directory / "logs" / "worker.log"), diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 630e0097..cc4b13d3 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -90,8 +90,8 @@ class TestNativeUp: "--postgres-uri", "postgresql://localhost/docsgpt"], _native_context(services)) == 0 api = services.units["docsgpt-api"] worker = services.units["docsgpt-worker"] - assert api.arguments[1:] == ["api", "--host", "127.0.0.1", "--port", "7091"] - assert worker.arguments[1:] == ["worker"] + assert api.arguments[-5:] == ["api", "--host", "127.0.0.1", "--port", "7091"] + assert worker.arguments[-1] == "worker" for unit in (api, worker): assert unit.environment["DOCSGPT_HOME"] == str(tmp_path) assert unit.working_directory == str(tmp_path) @@ -105,7 +105,7 @@ class TestNativeUp: assert _run(["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/docsgpt"], context) == 0 arguments, environment = calls[0] - assert arguments[1] == "migrate" + assert arguments[-1] == "migrate", "the launcher can be an interpreter and -m docsgpt" assert environment["DOCSGPT_HOME"] == str(tmp_path) assert environment["DOCSGPT_ENV_FILE"] == str(tmp_path / ".env") @@ -139,13 +139,14 @@ class TestNativeUp: with pytest.raises(DeployError, match="holds a Docker install"): _run(argv, _native_context()) - def test_services_are_not_written_without_a_docsgpt_command(self, tmp_path, monkeypatch): - """A unit pointing at something unexecutable would fail to exec with only a health check to show it.""" + def test_without_the_command_on_path_the_services_run_the_module(self, tmp_path, monkeypatch): + """A unit must name something executable; argv[0] is a module file under `python -m`.""" monkeypatch.setattr("docsgpt.deploy.commands.shutil.which", lambda name: None) monkeypatch.setattr(sys, "argv", [str(tmp_path / "not-a-program")]) + services = FakeServices() argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] - with pytest.raises(DeployError, match="not on PATH"): - _run(argv, _native_context()) + assert _run(argv, _native_context(services)) == 0 + assert services.units["docsgpt-api"].arguments[:3] == [sys.executable, "-m", "docsgpt"] def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): @@ -292,7 +293,7 @@ class TestSystemd: def test_install_writes_the_unit_and_reloads(self, tmp_path): systemctl = FakeSystemctl() services = self._services(systemctl, tmp_path) - services.install(native.units_for(tmp_path, "/venv/bin/docsgpt", 7091, tmp_path)[0]) + services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path)[0]) unit = tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service" assert "ExecStart=/venv/bin/docsgpt api" in unit.read_text() assert systemctl.verbs == ["daemon-reload"] @@ -308,7 +309,7 @@ class TestSystemd: def test_stop_and_remove(self, tmp_path): systemctl = FakeSystemctl() services = self._services(systemctl, tmp_path) - services.install(native.units_for(tmp_path, "/venv/bin/docsgpt", 7091, tmp_path)[0]) + services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path)[0]) unit = tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service" services.stop("docsgpt-api") assert unit.is_file(), "stopping keeps the unit" diff --git a/tests/test_cli.py b/tests/test_cli.py index 452f5598..dcad465e 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -34,6 +34,16 @@ class TestTopLevel: subprocess.run([sys.executable, "-c", code], cwd=Path(__file__).resolve().parents[1], check=True) +class TestModuleEntrypoint: + def test_python_m_docsgpt_runs_the_cli(self): + """`python -m docsgpt` is what a native service falls back to when the script is not on PATH.""" + result = subprocess.run( + [sys.executable, "-m", "docsgpt", "--version"], + cwd=Path(__file__).resolve().parents[1], capture_output=True, text=True, check=True, + ) + assert result.stdout.strip() == f"docsgpt {__version__}" + + class TestHome: @staticmethod def _installed(monkeypatch, tmp_path): From dbbed28f88e17d6eefdb77bcac1335aad099afe9 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:55:09 +0100 Subject: [PATCH 037/130] fix: upgrade re-execs a docsgpt it can actually find `upgrade` exec'd the bare name `docsgpt`, so after `python -m docsgpt upgrade` in a virtualenv without the console script on PATH, os.execv failed with a traceback. It now uses the same launcher the service units get, which is why that helper is no longer named for native mode. --- docsgpt/deploy/commands.py | 9 +++++---- tests/deploy/test_commands.py | 13 ++++++++++++- 2 files changed, 17 insertions(+), 5 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 0988044d..f718e2e0 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -218,8 +218,8 @@ def _redis_urls(base: str) -> dict[str, str]: } -def _native_launcher() -> list[str]: - """How the service files start DocsGPT: the command on PATH, or this interpreter and the module. +def _docsgpt_launcher() -> list[str]: + """How to start DocsGPT from another process: the command on PATH, or this interpreter and the module. A unit has to name something that can be executed, and ``sys.argv[0]`` often cannot be: under ``python -m docsgpt``, or pytest, it is a module file. Falling back to the running interpreter @@ -285,7 +285,7 @@ def _native_up(args, context: Context, directory: Path) -> int: (directory / "logs").mkdir(exist_ok=True) envfile.update(env_path, updates) - launcher = _native_launcher() + launcher = _docsgpt_launcher() print("Applying database migrations ...") # The child reads the stack's settings, not a .env in whatever directory this was run from. stack_env = {"DOCSGPT_HOME": str(directory), "DOCSGPT_ENV_FILE": str(env_path)} @@ -692,7 +692,8 @@ def upgrade(args, context: Optional[Context] = None) -> int: if installer == "uv": if context.run(["uv", "tool", "install", "--force", spec]) != 0: raise DeployError(f"uv could not install {spec}") - return context.exec_up(["docsgpt", "up", "--dir", str(stack.stack_dir(args.dir))]) + # The same launcher the service units get: a bare name is not always on PATH to exec. + return context.exec_up([*_docsgpt_launcher(), "up", "--dir", str(stack.stack_dir(args.dir))]) command = f"pipx install --force {spec}" if installer == "pipx" else f"pip install -U {spec}" print(f"Upgrade the package with `{command}`, then run `docsgpt up` to move the stack to it.", file=sys.stderr) return 1 diff --git a/tests/deploy/test_commands.py b/tests/deploy/test_commands.py index b73fedb2..9e9e684b 100644 --- a/tests/deploy/test_commands.py +++ b/tests/deploy/test_commands.py @@ -334,7 +334,8 @@ class TestUpgrade: ) assert _run(["upgrade", "--dir", str(tmp_path), "--version", "0.22.0"], context) == 0 assert self_calls[0] == ["uv", "tool", "install", "--force", "docsgpt==0.22.0"] - assert self_calls[1] == ["exec", "docsgpt", "up", "--dir", str(tmp_path)] + assert self_calls[1][0] == "exec" + assert self_calls[1][-3:] == ["up", "--dir", str(tmp_path)], "the launcher can be an interpreter and -m" def test_latest_when_no_version_is_given(self, tmp_path): self_calls = [] @@ -343,6 +344,16 @@ class TestUpgrade: assert _run(["upgrade", "--dir", str(tmp_path)], context) == 0 assert self_calls[0] == ["uv", "tool", "install", "--force", "docsgpt"] + def test_without_the_command_on_path_it_re_execs_the_module(self, tmp_path, monkeypatch): + """After `python -m docsgpt upgrade` there may be no docsgpt on PATH for execv to find.""" + monkeypatch.setattr("docsgpt.deploy.commands.shutil.which", lambda name: None) + monkeypatch.setattr(sys, "argv", [str(tmp_path / "not-a-program")]) + self_calls = [] + context = _context(installer=lambda: "uv", run=lambda args: 0, + exec_up=lambda argv: self_calls.append(argv) or 0) + assert _run(["upgrade", "--dir", str(tmp_path)], context) == 0 + assert self_calls[0][:3] == [sys.executable, "-m", "docsgpt"] + def test_a_pip_install_is_told_what_to_run(self, tmp_path, capsys): context = _context(installer=lambda: "pip", run=lambda args: pytest.fail("must not run")) assert _run(["upgrade", "--dir", str(tmp_path)], context) == 1 From 4d9f1d47a98ad9852a4dd23bbd203043ace3bbcc Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 22:58:52 +0100 Subject: [PATCH 038/130] fix: make a native install recoverable, honest and safe to quote From the outside-diff findings on #2800: - install.json is written before the services are installed and started. A service that fails to start used to leave units behind in a directory that status, down and uninstall no longer recognised as a native install, so nothing could clean them up. - systemd stop and removal propagate failures: `down` reporting success while the unit still runs, or `uninstall` dropping the unit file and the record while systemd still runs the service, is worse than an error. Removing a unit that is already gone stays harmless. - WorkingDirectory and each Environment value are quoted and escaped for systemd. `--dir` takes a free-form path, and one with a space in it is not hypothetical: this checkout lives in one. - Native status checks and prints http://localhost:, which is what the units bind. With a LAN DOCSGPT_BIND it used to poll an address nothing listened on and call a healthy install dead. --- docsgpt/deploy/commands.py | 20 +++++++---- docsgpt/deploy/native.py | 23 ++++++++---- tests/deploy/test_native.py | 70 +++++++++++++++++++++++++++++++++++++ 3 files changed, 100 insertions(+), 13 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index f718e2e0..45b9c461 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -292,17 +292,19 @@ def _native_up(args, context: Context, directory: Path) -> int: if context.run([*launcher, "migrate"], stack_env) != 0: raise DeployError("`docsgpt migrate` failed; check the database URL and that the server is reachable.") + # The record goes in before the services: if one fails to start, this is still a native install + # that status, down and uninstall can see and clean up, rather than orphaned units. + now = datetime.now(timezone.utc).isoformat(timespec="seconds") + record = record or {"installed_at": now} + record.update(version=context.version, mode="native", updated_at=now) + (directory / stack.RECORD_FILE).write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") + for unit in native.units_for(directory, launcher, port, directory): services.install(unit) for name in native.SERVICES: print(f"Starting {name} ...") services.start(name) - now = datetime.now(timezone.utc).isoformat(timespec="seconds") - record = record or {"installed_at": now} - record.update(version=context.version, mode="native", updated_at=now) - (directory / stack.RECORD_FILE).write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") - health = f"http://127.0.0.1:{port}/api/health" print("Waiting for DocsGPT to answer ...") if not context.wait(health, args.timeout): @@ -443,11 +445,15 @@ def status(args, context: Optional[Context] = None) -> int: return 1 if _mode(directory) == "native": services = context.service_manager() + # The units bind 127.0.0.1 whatever DOCSGPT_BIND says, so that is what gets printed and checked: + # against a LAN bind, stack.url and stack.health_url would advertise and poll an address + # nothing listens on, and status would call a healthy install dead. + port = env.get("DOCSGPT_PORT") or stack.DEFAULT_PORT print(f"DocsGPT {_record(directory).get('version', 'unknown')} in {directory} (native, {services.name})") - print(f"Address: {stack.url(env, context.lan_ip())}") + print(f"Address: http://localhost:{port}") for name in native.SERVICES: print(f" {name}: {'running' if services.is_running(name) else 'stopped'}") - healthy = context.wait(stack.health_url(env), 0) + healthy = context.wait(f"http://127.0.0.1:{port}/api/health", 0) print("API: answering" if healthy else "API: not answering (see `docsgpt logs`)") return 0 if healthy else 1 diff --git a/docsgpt/deploy/native.py b/docsgpt/deploy/native.py index 8191a33f..876e912a 100644 --- a/docsgpt/deploy/native.py +++ b/docsgpt/deploy/native.py @@ -61,9 +61,17 @@ def launchd_plist(unit: Unit, label: Optional[str] = None) -> str: return plistlib.dumps(body).decode("utf-8") +def _systemd_quote(value: str) -> str: + """A unit-file value, double-quoted with backslashes and quotes escaped, as systemd reads them.""" + escaped = value.replace("\\", "\\\\").replace('"', '\\"') + return f'"{escaped}"' + + def systemd_unit(unit: Unit) -> str: - """The systemd user unit for ``unit``; arguments are quoted, so a path with spaces survives.""" - environment = "\n".join(f'Environment="{key}={value}"' for key, value in sorted(unit.environment.items())) + """The systemd user unit for ``unit``; values are quoted, so a path with spaces survives.""" + environment = "\n".join( + f"Environment={_systemd_quote(f'{key}={value}')}" for key, value in sorted(unit.environment.items()) + ) command = " ".join(shlex.quote(argument) for argument in unit.arguments) return f"""[Unit] Description=DocsGPT ({unit.name}) @@ -72,7 +80,7 @@ After=network-online.target [Service] Type=simple ExecStart={command} -WorkingDirectory={unit.working_directory} +WorkingDirectory={_systemd_quote(unit.working_directory)} {environment} Restart=always RestartSec=5 @@ -181,12 +189,15 @@ class SystemdServices: self._systemctl("restart", f"{name}.service", check=True) def stop(self, name: str) -> None: - self._systemctl("stop", f"{name}.service") + self._systemctl("stop", f"{name}.service", check=True) def remove(self, name: str) -> None: - self._systemctl("disable", "--now", f"{name}.service") + # Only disable a unit systemd still knows about, so removing twice stays harmless; a real + # failure has to surface before the unit file and the install record are thrown away. + if self._unit_file(name).is_file(): + self._systemctl("disable", "--now", f"{name}.service", check=True) self._unit_file(name).unlink(missing_ok=True) - self._systemctl("daemon-reload") + self._systemctl("daemon-reload", check=True) def is_running(self, name: str) -> bool: return self._systemctl("is-active", "--quiet", f"{name}.service").returncode == 0 diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index cc4b13d3..54cb681d 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -148,6 +148,22 @@ class TestNativeUp: assert _run(argv, _native_context(services)) == 0 assert services.units["docsgpt-api"].arguments[:3] == [sys.executable, "-m", "docsgpt"] + def test_a_failed_start_leaves_an_install_the_other_commands_can_clean_up(self, tmp_path): + """Without the record, half-started units are orphaned: status, down and uninstall refuse the dir.""" + class FailingStart(FakeServices): + def start(self, name): + super().start(name) + if name == "docsgpt-worker": + raise DeployError("systemctl could not start docsgpt-worker") + + services = FailingStart() + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + with pytest.raises(DeployError, match="could not start"): + _run(argv, _native_context(services)) + assert json.loads((tmp_path / "install.json").read_text())["mode"] == "native" + assert _run(["uninstall", "--yes", "--dir", str(tmp_path)], _native_context(services)) == 0 + assert sorted(services.removed) == ["docsgpt-api", "docsgpt-worker"] + def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): _run(["up", "--native", "--dir", str(tmp_path), "--yes"], _native_context()) @@ -182,6 +198,17 @@ class TestNativeLifecycle: assert "native" in out.lower() assert "docsgpt-api" in out + def test_status_checks_the_address_the_services_actually_listen_on(self, tmp_path): + """A LAN DOCSGPT_BIND would have status poll an address the native units never bind.""" + services = FakeServices() + self._installed(tmp_path, services) + envfile.update(tmp_path / ".env", {"DOCSGPT_BIND": "192.168.1.50"}) + checked = [] + context = _native_context(services) + context.wait = lambda url, timeout: checked.append(url) or True + assert _run(["status", "--dir", str(tmp_path)], context) == 0 + assert checked == ["http://127.0.0.1:7091/api/health"] + def test_down_stops_them_and_leaves_the_settings(self, tmp_path): services = FakeServices() self._installed(tmp_path, services) @@ -321,6 +348,35 @@ class TestSystemd: assert self._services(FakeSystemctl(active=True), tmp_path).is_running("docsgpt-api") is True assert self._services(FakeSystemctl(active=False), tmp_path).is_running("docsgpt-api") is False + def test_a_stop_that_fails_is_not_reported_as_success(self, tmp_path): + """`down` saying it stopped while the service still runs is worse than an error.""" + class Failing(FakeSystemctl): + def __call__(self, args, **kwargs): + self.calls.append(args) + return types.SimpleNamespace(returncode=1, stdout="", stderr="Failed to stop docsgpt-api") + + with pytest.raises(DeployError, match="Failed to stop"): + self._services(Failing(), tmp_path).stop("docsgpt-api") + + def test_a_removal_that_fails_keeps_the_unit_file(self, tmp_path): + """uninstall must not drop the unit file and the record while systemd still runs the service.""" + class Failing(FakeSystemctl): + def __call__(self, args, **kwargs): + self.calls.append(args) + code = 1 if args[2] == "disable" else 0 + return types.SimpleNamespace(returncode=code, stdout="", stderr="Failed to disable") + + services = self._services(Failing(), tmp_path) + services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path)[0]) + with pytest.raises(DeployError, match="Failed to disable"): + services.remove("docsgpt-api") + assert (tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service").is_file() + + def test_removing_a_unit_that_is_already_gone_is_harmless(self, tmp_path): + systemctl = FakeSystemctl() + self._services(systemctl, tmp_path).remove("docsgpt-api") + assert "disable" not in systemctl.verbs, "nothing to disable when the unit file is gone" + def test_an_explicit_home_wins_over_xdg_config_home(self, tmp_path, monkeypatch): """home is how a caller redirects the unit directory; the environment must not override it.""" monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path / "xdg")) @@ -378,6 +434,20 @@ class TestUnitFiles: assert "ExecStart=/opt/venv/bin/docsgpt api --port 7091" in body assert "WantedBy=default.target" in body + def test_a_directory_with_spaces_and_quotes_survives_the_systemd_unit(self, tmp_path): + """--dir is free-form, and this repo itself lives under a path with a space in it.""" + odd = '/srv/my "odd" dir' + unit = native.Unit( + name="docsgpt-api", + arguments=["/venv/bin/docsgpt", "api"], + environment={"DOCSGPT_HOME": odd}, + working_directory=odd, + log_file="/srv/api.log", + ) + body = native.systemd_unit(unit) + assert 'WorkingDirectory="/srv/my \\"odd\\" dir"' in body + assert 'Environment="DOCSGPT_HOME=/srv/my \\"odd\\" dir"' in body + def test_an_argument_with_spaces_survives_the_systemd_unit(self, tmp_path): unit = self._unit(tmp_path) unit.arguments = ["/opt/my venv/bin/docsgpt", "api"] From 61f06fcce8732ad69874ce5c34bbbde4361cce1c Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 23:00:33 +0100 Subject: [PATCH 039/130] fix: docsgpt open points at the address a native install answers on `open` built its address from stack.url, which honours DOCSGPT_BIND, while the native units always bind 127.0.0.1: with a LAN bind it handed the browser an address nothing was listening on. status had the same mismatch and was fixed with it; both now go through one helper so they cannot drift apart again. --- docsgpt/deploy/commands.py | 22 +++++++++++++++++----- tests/deploy/test_native.py | 12 ++++++++++++ 2 files changed, 29 insertions(+), 5 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 45b9c461..aad99560 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -240,6 +240,16 @@ def _docsgpt_launcher() -> list[str]: ) +def _native_port(env: Mapping[str, str]) -> str: + """The port a native install listens on.""" + return str(env.get("DOCSGPT_PORT") or stack.DEFAULT_PORT) + + +def _native_address(env: Mapping[str, str]) -> str: + """Where a native install answers: its units bind 127.0.0.1, whatever DOCSGPT_BIND says.""" + return f"http://localhost:{_native_port(env)}" + + def _native_up(args, context: Context, directory: Path) -> int: """Run the API and the worker as services on this machine, against an existing Postgres and Redis.""" services = context.service_manager() @@ -445,12 +455,11 @@ def status(args, context: Optional[Context] = None) -> int: return 1 if _mode(directory) == "native": services = context.service_manager() - # The units bind 127.0.0.1 whatever DOCSGPT_BIND says, so that is what gets printed and checked: - # against a LAN bind, stack.url and stack.health_url would advertise and poll an address + # Against a LAN bind, stack.url and stack.health_url would advertise and poll an address # nothing listens on, and status would call a healthy install dead. - port = env.get("DOCSGPT_PORT") or stack.DEFAULT_PORT + port = _native_port(env) print(f"DocsGPT {_record(directory).get('version', 'unknown')} in {directory} (native, {services.name})") - print(f"Address: http://localhost:{port}") + print(f"Address: {_native_address(env)}") for name in native.SERVICES: print(f" {name}: {'running' if services.is_running(name) else 'stopped'}") healthy = context.wait(f"http://127.0.0.1:{port}/api/health", 0) @@ -510,7 +519,10 @@ def open_ui(args, context: Optional[Context] = None) -> int: env = _installed(directory) if env is None: return 1 - address = stack.url(env, context.lan_ip()) + # A native install answers on loopback only, so stack.url would hand the browser a LAN address + # or a domain that nothing behind this command is serving. + native_install = _mode(directory) == "native" + address = _native_address(env) if native_install else stack.url(env, context.lan_ip()) print(address) context.open_browser(address) return 0 diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 54cb681d..39a8938c 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -209,6 +209,18 @@ class TestNativeLifecycle: assert _run(["status", "--dir", str(tmp_path)], context) == 0 assert checked == ["http://127.0.0.1:7091/api/health"] + def test_open_hands_the_browser_the_address_that_answers(self, tmp_path, capsys): + """stack.url would offer a LAN address, but the native units listen on loopback only.""" + services = FakeServices() + self._installed(tmp_path, services) + envfile.update(tmp_path / ".env", {"DOCSGPT_BIND": "192.168.1.50"}) + opened = [] + context = _native_context(services) + context.open_browser = lambda url: opened.append(url) + assert _run(["open", "--dir", str(tmp_path)], context) == 0 + assert opened == ["http://localhost:7091"] + assert "192.168.1.50" not in capsys.readouterr().out + def test_down_stops_them_and_leaves_the_settings(self, tmp_path): services = FakeServices() self._installed(tmp_path, services) From 501baf8aae6988d369800a6a559f882a963db93c Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 23:01:30 +0100 Subject: [PATCH 040/130] fix: say why backup and restore do not apply to a native install Both accepted a native install and then drove `docker compose` in a directory with no compose file, so the user got "no configuration file provided" rather than an explanation. They now refuse with what to do instead, and restore refuses before it reads the archive or stops anything. --- docsgpt/deploy/commands.py | 10 +++++++++- tests/deploy/test_native.py | 11 +++++++++++ 2 files changed, 20 insertions(+), 1 deletion(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index aad99560..6353925a 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -568,6 +568,10 @@ def backup(args, context: Optional[Context] = None) -> int: env = _installed(directory) if env is None: return 1 + if _mode(directory) == "native": + raise DeployError( + f"{directory} is a native install: its database and its files are not in Docker volumes, so there is nothing here to archive. Back up the PostgreSQL that POSTGRES_URI points at with pg_dump, and copy the indexes, inputs and vectors folders from the data home." + ) out_dir = Path(args.out).expanduser() if args.out else directory / "backups" taken_at = datetime.now(timezone.utc) @@ -657,13 +661,17 @@ def _restore_data(context: Context, directory: Path, archive: Path, volumes: lis def restore(args, context: Optional[Context] = None) -> int: """Put a backup's database and data volumes back over this install.""" context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + if _mode(directory) == "native": + raise DeployError( + f"{directory} is a native install: its database and its files are not in Docker volumes, so there is nothing here to archive. Back up the PostgreSQL that POSTGRES_URI points at with pg_dump, and copy the indexes, inputs and vectors folders from the data home. `docsgpt restore` puts back what `docsgpt backup` wrote for a Docker install." + ) archive = Path(args.archive).expanduser() manifest = backup_format.read_manifest(archive) backup_format.check_version(manifest, context.version, args.force) # Everything the archive declares is checked here, while DocsGPT is still up. volumes = backup_format.validate(archive, manifest) - directory = stack.stack_dir(args.dir) env = _installed(directory) if env is None: return 1 diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 39a8938c..2b360803 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -238,6 +238,17 @@ class TestNativeLifecycle: self._installed(tmp_path, FakeServices()) assert _run(["logs", "--dir", str(tmp_path)], _native_context()) == 0, "logs work in both modes" + def test_backup_says_why_there_is_nothing_to_archive(self, tmp_path): + """Its data is not in Docker volumes, so compose would fail with no configuration file.""" + self._installed(tmp_path, FakeServices()) + with pytest.raises(DeployError, match="native install"): + _run(["backup", "--dir", str(tmp_path)], _native_context()) + + def test_restore_refuses_before_it_touches_anything(self, tmp_path): + self._installed(tmp_path, FakeServices()) + with pytest.raises(DeployError, match="native install"): + _run(["restore", str(tmp_path / "any.tar.gz"), "--dir", str(tmp_path), "--yes"], _native_context()) + class FakeLaunchctl: """launchctl, with a bootout that takes `unload_polls` polls to take effect.""" From 6cfd0948bc5b5604d3c76b9f188404e48e65605a Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 23:03:03 +0100 Subject: [PATCH 041/130] fix: refuse the Docker-only options in native mode instead of ignoring them `up --native` took --expose, --domain and --docling and did nothing with them. Asking for network exposure and silently getting a loopback-only install, or asking for docling and getting an install without it, is worse than being told. Each now says what to do instead: a reverse proxy or the Docker stack for exposure, and the docling extra for the parser engine. --expose local and --no-docling already describe native mode, so they stay silent. --- docsgpt/deploy/commands.py | 17 +++++++++++++++++ tests/deploy/test_native.py | 16 ++++++++++++++++ 2 files changed, 33 insertions(+) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 6353925a..1495bb5b 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -240,6 +240,22 @@ def _docsgpt_launcher() -> list[str]: ) +def _refuse_docker_only_options(args) -> None: + """Options that only mean something to the Docker stack, refused rather than quietly ignored.""" + if getattr(args, "domain", None) or getattr(args, "expose", None) in ("network", "domain"): + raise DeployError( + "a native install serves on 127.0.0.1 only, so --domain, --expose network and " + "--expose domain have nothing to act on here. Put a reverse proxy in front of it, or run " + "the Docker stack with `docsgpt up --expose ...`, which brings its own Caddy for a domain." + ) + if getattr(args, "docling", None): + raise DeployError( + "--docling selects a Docker image variant, which a native install does not use. Install the " + 'parser engine into this environment instead, with `uv tool install "docsgpt[docling]"` or ' + '`pip install "docsgpt[docling]"`, then run `docsgpt up --native` again.' + ) + + def _native_port(env: Mapping[str, str]) -> str: """The port a native install listens on.""" return str(env.get("DOCSGPT_PORT") or stack.DEFAULT_PORT) @@ -252,6 +268,7 @@ def _native_address(env: Mapping[str, str]) -> str: def _native_up(args, context: Context, directory: Path) -> int: """Run the API and the worker as services on this machine, against an existing Postgres and Redis.""" + _refuse_docker_only_options(args) services = context.service_manager() env_path = directory / ".env" existing = envfile.read(env_path) diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 2b360803..db3b1dfb 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -164,6 +164,22 @@ class TestNativeUp: assert _run(["uninstall", "--yes", "--dir", str(tmp_path)], _native_context(services)) == 0 assert sorted(services.removed) == ["docsgpt-api", "docsgpt-worker"] + def test_docker_only_options_are_refused_rather_than_ignored(self, tmp_path): + """Asking for network exposure and silently getting loopback is the worst of both.""" + base = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + with pytest.raises(DeployError, match="127.0.0.1 only"): + _run([*base, "--expose", "network"], _native_context()) + with pytest.raises(DeployError, match="127.0.0.1 only"): + _run([*base, "--domain", "docs.example.com"], _native_context()) + with pytest.raises(DeployError, match=r"docsgpt\[docling\]"): + _run([*base, "--docling"], _native_context()) + assert not (tmp_path / ".env").exists(), "it refuses before writing anything" + + def test_the_options_that_describe_what_native_already_does_are_kept(self, tmp_path): + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d", + "--expose", "local", "--no-docling"] + assert _run(argv, _native_context()) == 0 + def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): _run(["up", "--native", "--dir", str(tmp_path), "--yes"], _native_context()) From 76d81caaa468dd79de8bd20eda097379dc6bded4 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 23:10:53 +0100 Subject: [PATCH 042/130] fix: refuse control characters in the values of a service file MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Quoting cannot carry a newline into a unit file or a plist: the line ends and whatever follows becomes another directive. Every value bound for a service file — the working directory, the log path, environment names and values, and the command arguments — is checked before any of it is rendered, on both launchd and systemd. --- docsgpt/deploy/native.py | 23 +++++++++++++++++++++++ tests/deploy/test_native.py | 25 +++++++++++++++++++++++++ 2 files changed, 48 insertions(+) diff --git a/docsgpt/deploy/native.py b/docsgpt/deploy/native.py index 876e912a..ab78d0ed 100644 --- a/docsgpt/deploy/native.py +++ b/docsgpt/deploy/native.py @@ -13,6 +13,7 @@ from __future__ import annotations import os import plistlib +import re import shlex import subprocess import sys @@ -47,6 +48,7 @@ def label_for(name: str) -> str: def launchd_plist(unit: Unit, label: Optional[str] = None) -> str: """The launchd agent for ``unit``: kept alive, with its output in the stack's log file.""" + _check_unit_values(unit) body = { "Label": label or label_for(unit.name), "ProgramArguments": list(unit.arguments), @@ -61,6 +63,26 @@ def launchd_plist(unit: Unit, label: Optional[str] = None) -> str: return plistlib.dumps(body).decode("utf-8") +_CONTROL_CHARACTERS = re.compile(r"[\x00-\x1f\x7f]") + + +def _reject_control_characters(what: str, value: str) -> None: + """Service files are line-based, so a newline in a value adds a directive instead of text.""" + if _CONTROL_CHARACTERS.search(value): + raise DeployError(f"{what} contains a control character, which a service file cannot carry: {value!r}") + + +def _check_unit_values(unit: Unit) -> None: + """Everything bound for a service file, checked before any of it is rendered.""" + _reject_control_characters("the working directory", unit.working_directory) + _reject_control_characters("the log file path", unit.log_file) + for key, value in unit.environment.items(): + _reject_control_characters("an environment name", key) + _reject_control_characters(f"the environment value for {key}", value) + for argument in unit.arguments: + _reject_control_characters("a command argument", argument) + + def _systemd_quote(value: str) -> str: """A unit-file value, double-quoted with backslashes and quotes escaped, as systemd reads them.""" escaped = value.replace("\\", "\\\\").replace('"', '\\"') @@ -69,6 +91,7 @@ def _systemd_quote(value: str) -> str: def systemd_unit(unit: Unit) -> str: """The systemd user unit for ``unit``; values are quoted, so a path with spaces survives.""" + _check_unit_values(unit) environment = "\n".join( f"Environment={_systemd_quote(f'{key}={value}')}" for key, value in sorted(unit.environment.items()) ) diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index db3b1dfb..31dc6729 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -487,6 +487,31 @@ class TestUnitFiles: assert 'WorkingDirectory="/srv/my \\"odd\\" dir"' in body assert 'Environment="DOCSGPT_HOME=/srv/my \\"odd\\" dir"' in body + def test_a_newline_in_a_path_is_refused_rather_than_quoted(self, tmp_path): + """A service file is line-based: quoting cannot hold a newline, it would add a directive.""" + unit = native.Unit( + name="docsgpt-api", + arguments=["/venv/bin/docsgpt", "api"], + environment={"DOCSGPT_HOME": str(tmp_path)}, + working_directory="/srv/x\nExecStart=/bin/sh -c evil", + log_file="/srv/api.log", + ) + with pytest.raises(DeployError, match="control character"): + native.systemd_unit(unit) + with pytest.raises(DeployError, match="control character"): + native.launchd_plist(unit) + + def test_a_newline_in_the_environment_is_refused(self, tmp_path): + unit = native.Unit( + name="docsgpt-api", + arguments=["/venv/bin/docsgpt", "api"], + environment={"DOCSGPT_HOME": "/srv/x\nEnvironment=EVIL=1"}, + working_directory=str(tmp_path), + log_file="/srv/api.log", + ) + with pytest.raises(DeployError, match="control character"): + native.systemd_unit(unit) + def test_an_argument_with_spaces_survives_the_systemd_unit(self, tmp_path): unit = self._unit(tmp_path) unit.arguments = ["/opt/my venv/bin/docsgpt", "api"] From 65ee2f42019b7dd496d8cd6b9ac4f955e582525c Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 23:22:07 +0100 Subject: [PATCH 043/130] fix: keep the whole Redis URL when handing out its databases MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The three URLs were built by string surgery, so anything after the database number was mangled rather than kept: rediss://host:6380/0?ssl_cert_reqs=required came out as .../0?ssl_cert_reqs=required/0, and a URL carrying a query but no database had /0 appended after the query. TLS and managed Redis endpoints usually carry exactly those parameters. The URL is split properly now, the three databases go in the path, and scheme, credentials, host, query and fragment are preserved. A URL that cannot be numbered this way — one that is not redis:// or rediss://, or that has something other than a number where the database goes — is refused with a message instead of being turned into something that merely looks like a URL. --- docsgpt/deploy/commands.py | 27 +++++++++++++++++++-------- tests/deploy/test_native.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 48 insertions(+), 8 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 1495bb5b..8d00f7bb 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -17,6 +17,7 @@ from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path from typing import Any, Optional +from urllib.parse import urlsplit, urlunsplit from docsgpt.deploy import backup as backup_format from docsgpt.deploy import envfile, native, stack @@ -205,16 +206,26 @@ def _redis_urls(base: str) -> dict[str, str]: They start at database 0, or at the one the URL names: ``redis://host:6379/5`` puts them on 5, 6 and 7, which is how one Redis is shared with something that already uses the first databases. + Everything else in the URL is kept — TLS, credentials, and query parameters such as the + ``ssl_cert_reqs`` that a rediss:// endpoint usually needs. """ - trimmed = base.rstrip("/") - first = 0 - head, _, tail = trimmed.rpartition("/") - if tail.isdigit(): - trimmed, first = head, int(tail) + parts = urlsplit(base) + if parts.scheme not in ("redis", "rediss"): + raise DeployError( + f"the Redis URL {base!r} should start with redis:// or rediss://, with any options as " + "query parameters, so the broker, the result backend and the cache can be given a " + "database each." + ) + path = parts.path.rstrip("/").lstrip("/") + if path and not path.isdigit(): + raise DeployError( + f"the Redis URL {base!r} has {path!r} where a database number would go. Pass a URL like " + "redis://host:6379 or redis://host:6379/5." + ) + first = int(path) if path else 0 return { - "CELERY_BROKER_URL": f"{trimmed}/{first}", - "CELERY_RESULT_BACKEND": f"{trimmed}/{first + 1}", - "CACHE_REDIS_URL": f"{trimmed}/{first + 2}", + key: urlunsplit((parts.scheme, parts.netloc, f"/{first + offset}", parts.query, parts.fragment)) + for key, offset in (("CELERY_BROKER_URL", 0), ("CELERY_RESULT_BACKEND", 1), ("CACHE_REDIS_URL", 2)) } diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 31dc6729..cd7092ca 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -180,6 +180,35 @@ class TestNativeUp: "--expose", "local", "--no-docling"] assert _run(argv, _native_context()) == 0 + def test_a_redis_url_keeps_its_credentials_tls_and_query(self, tmp_path): + """A rediss:// endpoint usually needs ssl_cert_reqs, and losing it breaks every connection.""" + url = "rediss://user:pw@redis.example.com:6380/3?ssl_cert_reqs=required" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/d", "--redis-url", url] + assert _run(argv, _native_context()) == 0 + env = envfile.read(tmp_path / ".env") + host = "rediss://user:pw@redis.example.com:6380" + assert env["CELERY_BROKER_URL"] == f"{host}/3?ssl_cert_reqs=required" + assert env["CELERY_RESULT_BACKEND"] == f"{host}/4?ssl_cert_reqs=required" + assert env["CACHE_REDIS_URL"] == f"{host}/5?ssl_cert_reqs=required" + + def test_a_redis_url_with_a_query_and_no_database(self, tmp_path): + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d", + "--redis-url", "redis://localhost:6379?health_check_interval=30"] + assert _run(argv, _native_context()) == 0 + env = envfile.read(tmp_path / ".env") + assert env["CELERY_BROKER_URL"] == "redis://localhost:6379/0?health_check_interval=30" + assert env["CACHE_REDIS_URL"] == "redis://localhost:6379/2?health_check_interval=30" + + @pytest.mark.parametrize("url", ["localhost:6379", "redis+socket:///var/run/redis.sock", + "redis://localhost:6379/queue"]) + def test_a_redis_url_that_cannot_be_numbered_is_refused(self, tmp_path, url): + """Silently turning it into something that looks like a URL is the worse failure.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/d", "--redis-url", url] + with pytest.raises(DeployError, match="Redis URL"): + _run(argv, _native_context()) + def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): _run(["up", "--native", "--dir", str(tmp_path), "--yes"], _native_context()) From c8ab1cbab26f882d63a3a5a5d085e4cd3c5353d4 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 23:30:47 +0100 Subject: [PATCH 044/130] fix: report a malformed Redis URL instead of raising from urlsplit urlsplit raises ValueError on input such as redis://[::1 , which nothing converted, so a typo left native setup with a traceback rather than the message every other unusable URL gets. --- docsgpt/deploy/commands.py | 6 +++++- tests/deploy/test_native.py | 16 ++++++++++++++++ 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 8d00f7bb..5a853e8f 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -209,7 +209,11 @@ def _redis_urls(base: str) -> dict[str, str]: Everything else in the URL is kept — TLS, credentials, and query parameters such as the ``ssl_cert_reqs`` that a rediss:// endpoint usually needs. """ - parts = urlsplit(base) + try: + parts = urlsplit(base) + except ValueError as exc: + # urlsplit raises on things like redis://[::1 ; that is a typo, not a crash. + raise DeployError(f"the Redis URL {base!r} could not be read: {exc}") from exc if parts.scheme not in ("redis", "rediss"): raise DeployError( f"the Redis URL {base!r} should start with redis:// or rediss://, with any options as " diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index cd7092ca..0e09a382 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -200,6 +200,22 @@ class TestNativeUp: assert env["CELERY_BROKER_URL"] == "redis://localhost:6379/0?health_check_interval=30" assert env["CACHE_REDIS_URL"] == "redis://localhost:6379/2?health_check_interval=30" + def test_a_malformed_redis_url_is_reported_not_raised(self, tmp_path): + """urlsplit raises ValueError on an unterminated IPv6 bracket; a typo is not a traceback.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/d", "--redis-url", "redis://[::1"] + with pytest.raises(DeployError, match="could not be read"): + _run(argv, _native_context()) + + def test_an_ipv6_redis_url_still_works(self, tmp_path): + """Refusing the malformed form must not cost the valid one.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d", + "--redis-url", "redis://[::1]:6379/2"] + assert _run(argv, _native_context()) == 0 + env = envfile.read(tmp_path / ".env") + assert env["CELERY_BROKER_URL"] == "redis://[::1]:6379/2" + assert env["CACHE_REDIS_URL"] == "redis://[::1]:6379/4" + @pytest.mark.parametrize("url", ["localhost:6379", "redis+socket:///var/run/redis.sock", "redis://localhost:6379/queue"]) def test_a_redis_url_that_cannot_be_numbered_is_refused(self, tmp_path, url): From 4976b3104cbd9c2f832e6a909d076e697dcdc71a Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 16 Sep 2026 23:51:20 +0100 Subject: [PATCH 045/130] fix: give each native install its own services, and guard the port From the outside-diff findings on #2800: - Service names are derived from the install directory. A service manager has one namespace per user, so two installs in different --dir directories wrote over each other's units and down, status and uninstall acted on whichever was written last. The default install keeps the readable names; another directory gets a digest suffix. - `up --native` refuses when a Docker stack in another directory publishes the same port: its API would answer the health check while these services failed to bind. The check degrades quietly when Docker is absent, which is exactly the machine a native install targets. - A Redis database path is required to be ASCII digits: str.isdigit() is true for characters int() then refuses. - DOCSGPT_PORT from a hand-edited .env is validated before conversion, and the error names where the bad value came from. - Percent signs are doubled in systemd values, arguments and log paths, since systemd expands specifiers in all of them. --- docsgpt/deploy/commands.py | 58 ++++++++++++++++++---- docsgpt/deploy/native.py | 36 ++++++++++---- tests/deploy/test_native.py | 98 +++++++++++++++++++++++++++++++------ 3 files changed, 159 insertions(+), 33 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 5a853e8f..1f526e86 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -221,7 +221,7 @@ def _redis_urls(base: str) -> dict[str, str]: "database each." ) path = parts.path.rstrip("/").lstrip("/") - if path and not path.isdigit(): + if path and not (path.isascii() and path.isdigit()): raise DeployError( f"the Redis URL {base!r} has {path!r} where a database number would go. Pass a URL like " "redis://host:6379 or redis://host:6379/5." @@ -271,6 +271,36 @@ def _refuse_docker_only_options(args) -> None: ) +def _port_number(value: object, source: str) -> int: + """A port from the command line or a hand-edited .env, or a DeployError naming where it came from.""" + text = str(value) + if not (text.isascii() and text.isdigit() and 0 < int(text) < 65536): + raise DeployError(f"{source} is {value!r}, which is not a port number between 1 and 65535.") + return int(text) + + +def _service_names(directory: Path) -> tuple[str, str]: + """This install's service names; an install elsewhere gets its own, so the two cannot collide.""" + return native.service_names(directory, stack.stack_dir(None)) + + +def _refuse_other_docker_stack(context: Context, directory: Path, port: int) -> None: + """A Docker stack elsewhere on this port would answer the health check the native API failed.""" + try: + others = {path for path in context.docker.project_dirs(PROJECT) if path.resolve() != directory.resolve()} + except DeployError: + return # No Docker on this machine to ask, which is a fair reason to run natively. + for other in sorted(others): + other_port = envfile.read(other / ".env").get("DOCSGPT_PORT") or str(stack.DEFAULT_PORT) + if str(other_port) == str(port): + raise DeployError( + f"Docker already runs a DocsGPT stack from {other} on port {port}, which a native " + f"install would try to bind as well: the health check here could answer from that " + f"stack while these services failed to start. Stop it with `docsgpt down --dir " + f"{other}`, or give this install another port with --port." + ) + + def _native_port(env: Mapping[str, str]) -> str: """The port a native install listens on.""" return str(env.get("DOCSGPT_PORT") or stack.DEFAULT_PORT) @@ -304,7 +334,12 @@ def _native_up(args, context: Context, directory: Path) -> int: ) # An install keeps the port it was given: a later `up` with no --port must not move it back to the default. - port = int(args.port or existing.get("DOCSGPT_PORT") or stack.DEFAULT_PORT) + if args.port: + port = _port_number(args.port, "--port") + elif existing.get("DOCSGPT_PORT"): + port = _port_number(existing["DOCSGPT_PORT"], f"DOCSGPT_PORT in {env_path}") + else: + port = stack.DEFAULT_PORT updates: dict[str, Optional[str]] = { "POSTGRES_URI": postgres, "API_URL": f"http://127.0.0.1:{port}", @@ -334,6 +369,8 @@ def _native_up(args, context: Context, directory: Path) -> int: if context.run([*launcher, "migrate"], stack_env) != 0: raise DeployError("`docsgpt migrate` failed; check the database URL and that the server is reachable.") + _refuse_other_docker_stack(context, directory, port) + # The record goes in before the services: if one fails to start, this is still a native install # that status, down and uninstall can see and clean up, rather than orphaned units. now = datetime.now(timezone.utc).isoformat(timespec="seconds") @@ -341,9 +378,10 @@ def _native_up(args, context: Context, directory: Path) -> int: record.update(version=context.version, mode="native", updated_at=now) (directory / stack.RECORD_FILE).write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") - for unit in native.units_for(directory, launcher, port, directory): + names = _service_names(directory) + for unit in native.units_for(directory, launcher, port, directory, names): services.install(unit) - for name in native.SERVICES: + for name in names: print(f"Starting {name} ...") services.start(name) @@ -358,7 +396,7 @@ def _native_up(args, context: Context, directory: Path) -> int: return 1 print(f"\nDocsGPT is running at http://localhost:{port}") - print(f"Services: {', '.join(native.SERVICES)} under {services.name}") + print(f"Services: {', '.join(names)} under {services.name}") print(f"Settings: {env_path}") print("Manage it with: docsgpt status | logs | down | uninstall") return 0 @@ -471,7 +509,7 @@ def down(args, context: Optional[Context] = None) -> int: if _mode(directory) == "native": services = context.service_manager() # The worker goes first: it talks to the API, not the other way round. - for name in reversed(native.SERVICES): + for name in reversed(_service_names(directory)): services.stop(name) return 0 context.docker.compose(directory, *_EVERY_PROFILE, "down") @@ -492,7 +530,7 @@ def status(args, context: Optional[Context] = None) -> int: port = _native_port(env) print(f"DocsGPT {_record(directory).get('version', 'unknown')} in {directory} (native, {services.name})") print(f"Address: {_native_address(env)}") - for name in native.SERVICES: + for name in _service_names(directory): print(f" {name}: {'running' if services.is_running(name) else 'stopped'}") healthy = context.wait(f"http://127.0.0.1:{port}/api/health", 0) print("API: answering" if healthy else "API: not answering (see `docsgpt logs`)") @@ -515,7 +553,7 @@ def logs(args, context: Optional[Context] = None) -> int: return 1 if _mode(directory) == "native": logs_dir = directory / "logs" - wanted = args.services or [name.removeprefix("docsgpt-") for name in native.SERVICES] + wanted = args.services or ["api", "worker"] for service in wanted: path = logs_dir / f"{service}.log" print(f"=== {path}") @@ -775,9 +813,9 @@ def uninstall(args, context: Optional[Context] = None) -> int: return 1 if _mode(directory) == "native": services = context.service_manager() - for name in reversed(native.SERVICES): + for name in reversed(_service_names(directory)): services.remove(name) - print(f"Removed the {', '.join(native.SERVICES)} services.") + print(f"Removed the {', '.join(_service_names(directory))} services.") if args.purge: shutil.rmtree(directory) print(f"Removed {directory}. The database and Redis it used are untouched.") diff --git a/docsgpt/deploy/native.py b/docsgpt/deploy/native.py index ab78d0ed..85e02e01 100644 --- a/docsgpt/deploy/native.py +++ b/docsgpt/deploy/native.py @@ -11,6 +11,7 @@ already run (``--postgres-uri``, ``--redis-url``). from __future__ import annotations +import hashlib import os import plistlib import re @@ -26,7 +27,6 @@ from docsgpt.deploy.docker import DeployError API_SERVICE = "docsgpt-api" WORKER_SERVICE = "docsgpt-worker" -SERVICES = (API_SERVICE, WORKER_SERVICE) LABEL_PREFIX = "cloud.docsgpt" @@ -41,6 +41,20 @@ class Unit: log_file: str +def service_names(stack_directory: Path, default_directory: Path) -> tuple[str, str]: + """This install's two service names: the plain pair for the default install, suffixed for others. + + Service managers keep one namespace for the whole user, so two installs in different directories + would write over each other's units. These names show up in launchctl and systemctl output, so + the usual install keeps the readable ones and only a second install carries a digest. + """ + resolved = stack_directory.expanduser().resolve() + if resolved == default_directory.expanduser().resolve(): + return (API_SERVICE, WORKER_SERVICE) + digest = hashlib.sha256(str(resolved).encode("utf-8")).hexdigest()[:8] + return (f"{API_SERVICE}-{digest}", f"{WORKER_SERVICE}-{digest}") + + def label_for(name: str) -> str: """The launchd label for a service name (``docsgpt-api`` -> ``cloud.docsgpt.api``).""" return f"{LABEL_PREFIX}.{name.removeprefix('docsgpt-')}" @@ -84,8 +98,8 @@ def _check_unit_values(unit: Unit) -> None: def _systemd_quote(value: str) -> str: - """A unit-file value, double-quoted with backslashes and quotes escaped, as systemd reads them.""" - escaped = value.replace("\\", "\\\\").replace('"', '\\"') + """A unit-file value: quoted, with backslashes, quotes and percent signs escaped for systemd.""" + escaped = value.replace("\\", "\\\\").replace('"', '\\"').replace("%", "%%") return f'"{escaped}"' @@ -95,7 +109,9 @@ def systemd_unit(unit: Unit) -> str: environment = "\n".join( f"Environment={_systemd_quote(f'{key}={value}')}" for key, value in sorted(unit.environment.items()) ) - command = " ".join(shlex.quote(argument) for argument in unit.arguments) + # systemd expands % specifiers such as %h, so a literal percent has to be doubled everywhere. + command = " ".join(shlex.quote(argument).replace("%", "%%") for argument in unit.arguments) + log_file = unit.log_file.replace("%", "%%") return f"""[Unit] Description=DocsGPT ({unit.name}) After=network-online.target @@ -107,8 +123,8 @@ WorkingDirectory={_systemd_quote(unit.working_directory)} {environment} Restart=always RestartSec=5 -StandardOutput=append:{unit.log_file} -StandardError=append:{unit.log_file} +StandardOutput=append:{log_file} +StandardError=append:{log_file} [Install] WantedBy=default.target @@ -238,22 +254,24 @@ def services_for_platform(platform: str = sys.platform): ) -def units_for(stack_directory: Path, launcher: list[str], port: int, home: Path) -> list[Unit]: +def units_for(stack_directory: Path, launcher: list[str], port: int, home: Path, + names: tuple[str, str]) -> list[Unit]: """The API and worker services for a native install in ``stack_directory``. ``launcher`` is how DocsGPT is started: the ``docsgpt`` command, or an interpreter and ``-m``. """ environment = {"DOCSGPT_HOME": str(home)} + api_name, worker_name = names return [ Unit( - name=API_SERVICE, + name=api_name, arguments=[*launcher, "api", "--host", "127.0.0.1", "--port", str(port)], environment=dict(environment), working_directory=str(stack_directory), log_file=str(stack_directory / "logs" / "api.log"), ), Unit( - name=WORKER_SERVICE, + name=worker_name, arguments=[*launcher, "worker"], environment=dict(environment), working_directory=str(stack_directory), diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 0e09a382..55b6d3c3 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -4,13 +4,14 @@ import json import plistlib import sys import types +from pathlib import Path import pytest -from docsgpt.deploy import envfile, native +from docsgpt.deploy import envfile, native, stack from docsgpt.deploy.docker import DeployError -from .test_commands import FakePrompter, _context, _run +from .test_commands import FakeDocker, FakePrompter, _context, _run class FakeServices: @@ -44,6 +45,11 @@ class FakeServices: return name in self.running +def _names(directory): + """The two service names an install in ``directory`` gets; a second directory gets its own.""" + return native.service_names(Path(directory), stack.stack_dir(None)) + + def _native_context(services=None, migrated=None, **overrides): context = _context(**overrides) context.services = services or FakeServices() @@ -82,14 +88,15 @@ class TestNativeUp: assert _run(["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/docsgpt"], context) == 0 assert any("migrate" in " ".join(call) for call in migrated), migrated - assert services.started == ["docsgpt-api", "docsgpt-worker"] + assert services.started == list(_names(tmp_path)) def test_the_units_run_this_interpreter_with_the_stack_as_its_home(self, tmp_path): services = FakeServices() assert _run(["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/docsgpt"], _native_context(services)) == 0 - api = services.units["docsgpt-api"] - worker = services.units["docsgpt-worker"] + api_name, worker_name = _names(tmp_path) + api = services.units[api_name] + worker = services.units[worker_name] assert api.arguments[-5:] == ["api", "--host", "127.0.0.1", "--port", "7091"] assert worker.arguments[-1] == "worker" for unit in (api, worker): @@ -130,7 +137,7 @@ class TestNativeUp: env = envfile.read(tmp_path / ".env") assert env["DOCSGPT_PORT"] == "7099" assert env["API_URL"] == "http://127.0.0.1:7099" - assert services.units["docsgpt-api"].arguments[-1] == "7099" + assert services.units[_names(tmp_path)[0]].arguments[-1] == "7099" def test_it_refuses_to_run_beside_a_docker_install(self, tmp_path): """Native services on the same port would orphan the containers from down/status/uninstall.""" @@ -146,14 +153,14 @@ class TestNativeUp: services = FakeServices() argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] assert _run(argv, _native_context(services)) == 0 - assert services.units["docsgpt-api"].arguments[:3] == [sys.executable, "-m", "docsgpt"] + assert services.units[_names(tmp_path)[0]].arguments[:3] == [sys.executable, "-m", "docsgpt"] def test_a_failed_start_leaves_an_install_the_other_commands_can_clean_up(self, tmp_path): """Without the record, half-started units are orphaned: status, down and uninstall refuse the dir.""" class FailingStart(FakeServices): def start(self, name): super().start(name) - if name == "docsgpt-worker": + if name == _names(tmp_path)[1]: raise DeployError("systemctl could not start docsgpt-worker") services = FailingStart() @@ -162,7 +169,7 @@ class TestNativeUp: _run(argv, _native_context(services)) assert json.loads((tmp_path / "install.json").read_text())["mode"] == "native" assert _run(["uninstall", "--yes", "--dir", str(tmp_path)], _native_context(services)) == 0 - assert sorted(services.removed) == ["docsgpt-api", "docsgpt-worker"] + assert sorted(services.removed) == sorted(_names(tmp_path)) def test_docker_only_options_are_refused_rather_than_ignored(self, tmp_path): """Asking for network exposure and silently getting loopback is the worst of both.""" @@ -225,6 +232,51 @@ class TestNativeUp: with pytest.raises(DeployError, match="Redis URL"): _run(argv, _native_context()) + def test_a_port_that_is_not_a_number_is_reported(self, tmp_path): + """A hand-edited DOCSGPT_PORT reached int() unchecked and came out as a traceback.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + assert _run([*argv, "--port", "7099"], _native_context()) == 0 + envfile.update(tmp_path / ".env", {"DOCSGPT_PORT": "seven thousand"}) + with pytest.raises(DeployError, match="not a port number"): + _run(["up", "--dir", str(tmp_path), "--yes"], _native_context()) + + def test_a_redis_database_that_int_would_reject_is_reported(self, tmp_path): + """str.isdigit() is true for a superscript two, which int() then refuses.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d", + "--redis-url", "redis://localhost:6379/\u00b2"] + with pytest.raises(DeployError, match="Redis URL"): + _run(argv, _native_context()) + + def test_two_installs_do_not_share_service_names(self, tmp_path): + """A service manager has one namespace per user, so the second install would overwrite the first.""" + installed = {} + for name in ("one", "two"): + directory = tmp_path / name + services = FakeServices() + argv = ["up", "--native", "--dir", str(directory), "--yes", + "--postgres-uri", "postgresql://localhost/d"] + assert _run(argv, _native_context(services)) == 0 + installed[name] = set(services.units) + assert installed["one"].isdisjoint(installed["two"]), installed + + def test_the_default_install_keeps_the_readable_names(self): + """They are what shows up in launchctl and systemctl, so the usual install is not a digest.""" + default = Path("/opt/docsgpt") + assert native.service_names(default, default) == ("docsgpt-api", "docsgpt-worker") + + def test_a_docker_stack_elsewhere_on_the_same_port_is_refused(self, tmp_path): + """Its API would answer the health check while these services quietly failed to bind.""" + other = tmp_path / "docker-install" + other.mkdir() + (other / ".env").write_text("DOCSGPT_PORT=7099\n", encoding="utf-8") + context = _native_context(FakeServices(), docker=FakeDocker(project_dirs=[str(other)])) + directory = tmp_path / "native" + argv = ["up", "--native", "--dir", str(directory), "--yes", "--port", "7099", + "--postgres-uri", "postgresql://localhost/d"] + with pytest.raises(DeployError, match="already runs a DocsGPT stack"): + _run(argv, context) + assert not (directory / "install.json").exists(), "it refuses before recording the install" + def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): _run(["up", "--native", "--dir", str(tmp_path), "--yes"], _native_context()) @@ -286,14 +338,14 @@ class TestNativeLifecycle: services = FakeServices() self._installed(tmp_path, services) assert _run(["down", "--dir", str(tmp_path)], _native_context(services)) == 0 - assert services.stopped == ["docsgpt-worker", "docsgpt-api"], "the worker goes first" + assert services.stopped == list(reversed(_names(tmp_path))), "the worker goes first" assert (tmp_path / ".env").is_file() def test_uninstall_removes_the_services(self, tmp_path): services = FakeServices() self._installed(tmp_path, services) assert _run(["uninstall", "--yes", "--dir", str(tmp_path)], _native_context(services)) == 0 - assert sorted(services.removed) == ["docsgpt-api", "docsgpt-worker"] + assert sorted(services.removed) == sorted(_names(tmp_path)) def test_docker_commands_refuse_a_native_install(self, tmp_path, capsys): self._installed(tmp_path, FakeServices()) @@ -404,7 +456,8 @@ class TestSystemd: def test_install_writes_the_unit_and_reloads(self, tmp_path): systemctl = FakeSystemctl() services = self._services(systemctl, tmp_path) - services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path)[0]) + services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path, + ("docsgpt-api", "docsgpt-worker"))[0]) unit = tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service" assert "ExecStart=/venv/bin/docsgpt api" in unit.read_text() assert systemctl.verbs == ["daemon-reload"] @@ -420,7 +473,8 @@ class TestSystemd: def test_stop_and_remove(self, tmp_path): systemctl = FakeSystemctl() services = self._services(systemctl, tmp_path) - services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path)[0]) + services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path, + ("docsgpt-api", "docsgpt-worker"))[0]) unit = tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service" services.stop("docsgpt-api") assert unit.is_file(), "stopping keeps the unit" @@ -451,7 +505,8 @@ class TestSystemd: return types.SimpleNamespace(returncode=code, stdout="", stderr="Failed to disable") services = self._services(Failing(), tmp_path) - services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path)[0]) + services.install(native.units_for(tmp_path, ["/venv/bin/docsgpt"], 7091, tmp_path, + ("docsgpt-api", "docsgpt-worker"))[0]) with pytest.raises(DeployError, match="Failed to disable"): services.remove("docsgpt-api") assert (tmp_path / ".config" / "systemd" / "user" / "docsgpt-api.service").is_file() @@ -557,6 +612,21 @@ class TestUnitFiles: with pytest.raises(DeployError, match="control character"): native.systemd_unit(unit) + def test_a_percent_sign_is_doubled_for_systemd(self, tmp_path): + """systemd expands %h and friends, so a literal percent in a path has to be escaped.""" + unit = native.Unit( + name="docsgpt-api", + arguments=["/venv/bin/docsgpt", "api", "--flag", "100%"], + environment={"DOCSGPT_HOME": "/srv/100% full"}, + working_directory="/srv/100% full", + log_file="/srv/100% full/api.log", + ) + body = native.systemd_unit(unit) + assert 'WorkingDirectory="/srv/100%% full"' in body + assert 'Environment="DOCSGPT_HOME=/srv/100%% full"' in body + assert "append:/srv/100%% full/api.log" in body + assert "100%%" in body.split("ExecStart=")[1].split("\n")[0] + def test_an_argument_with_spaces_survives_the_systemd_unit(self, tmp_path): unit = self._unit(tmp_path) unit.arguments = ["/opt/my venv/bin/docsgpt", "api"] From 217310e201c71cfbbc0004bffa050277a8645eaa Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 00:03:06 +0100 Subject: [PATCH 046/130] fix: run the native preflights before anything is written MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Docker-stack check ran after logs/ was created, .env written and the migrations applied, so a conflict left a migrated database and partial files behind with no install.json — exactly the directory that down and uninstall then refuse. It runs immediately after the port is resolved now. Alongside it, `up --native` refuses a port it cannot bind. Neither service manager confirms that the API bound, and /api/health carries no installation identity, so a second install on the same port would have been answered by the first and reported success while its own API was dead. An install re-running on its own port is the exception, since its services are what hold it. --- docsgpt/deploy/commands.py | 36 +++++++++++++++++++++++++++++++++--- tests/deploy/test_native.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 62 insertions(+), 3 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 1f526e86..5ba2c4e3 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -8,6 +8,7 @@ import os import re import secrets import shutil +import socket import subprocess import sys import tempfile @@ -284,6 +285,33 @@ def _service_names(directory: Path) -> tuple[str, str]: return native.service_names(directory, stack.stack_dir(None)) +def _port_is_free(port: int) -> bool: + """Whether the loopback port can still be bound.""" + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: + try: + probe.bind(("127.0.0.1", port)) + except OSError: + return False + return True + + +def _refuse_busy_port(directory: Path, port: int, services, names: tuple[str, str]) -> None: + """Whatever holds the port would answer the health check while these services failed to bind. + + An install re-running on its own port is the exception: its services are holding it, and + starting them again replaces them. + """ + if _port_is_free(port): + return + if _record(directory).get("mode") == "native" and any(services.is_running(name) for name in names): + return + raise DeployError( + f"port {port} is already in use by something other than this install. The API would fail to " + f"bind it while the health check answered from whatever holds it, so the install would look " + f"healthy and be dead. Free the port, or give this install another one with --port." + ) + + def _refuse_other_docker_stack(context: Context, directory: Path, port: int) -> None: """A Docker stack elsewhere on this port would answer the health check the native API failed.""" try: @@ -340,6 +368,11 @@ def _native_up(args, context: Context, directory: Path) -> int: port = _port_number(existing["DOCSGPT_PORT"], f"DOCSGPT_PORT in {env_path}") else: port = stack.DEFAULT_PORT + + names = _service_names(directory) + _refuse_other_docker_stack(context, directory, port) + _refuse_busy_port(directory, port, services, names) + updates: dict[str, Optional[str]] = { "POSTGRES_URI": postgres, "API_URL": f"http://127.0.0.1:{port}", @@ -369,8 +402,6 @@ def _native_up(args, context: Context, directory: Path) -> int: if context.run([*launcher, "migrate"], stack_env) != 0: raise DeployError("`docsgpt migrate` failed; check the database URL and that the server is reachable.") - _refuse_other_docker_stack(context, directory, port) - # The record goes in before the services: if one fails to start, this is still a native install # that status, down and uninstall can see and clean up, rather than orphaned units. now = datetime.now(timezone.utc).isoformat(timespec="seconds") @@ -378,7 +409,6 @@ def _native_up(args, context: Context, directory: Path) -> int: record.update(version=context.version, mode="native", updated_at=now) (directory / stack.RECORD_FILE).write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") - names = _service_names(directory) for unit in native.units_for(directory, launcher, port, directory, names): services.install(unit) for name in names: diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 55b6d3c3..ce66744f 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -2,6 +2,7 @@ import json import plistlib +import socket import sys import types from pathlib import Path @@ -273,9 +274,37 @@ class TestNativeUp: directory = tmp_path / "native" argv = ["up", "--native", "--dir", str(directory), "--yes", "--port", "7099", "--postgres-uri", "postgresql://localhost/d"] + migrated = [] + context.run = lambda args, env=None: migrated.append(args) or 0 with pytest.raises(DeployError, match="already runs a DocsGPT stack"): _run(argv, context) assert not (directory / "install.json").exists(), "it refuses before recording the install" + assert not (directory / ".env").exists(), "and before writing any settings" + assert not (directory / "logs").exists() + assert migrated == [], "and before touching the database" + + def test_a_port_something_else_holds_is_refused(self, tmp_path): + """Nothing confirms the API bound its port, and a generic health check answers from the holder.""" + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as held: + held.bind(("127.0.0.1", 0)) + held.listen(1) + port = held.getsockname()[1] + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--port", str(port), + "--postgres-uri", "postgresql://localhost/d"] + with pytest.raises(DeployError, match="already in use"): + _run(argv, _native_context()) + assert not (tmp_path / ".env").exists(), "it refuses before writing anything" + + def test_an_install_may_keep_the_port_its_own_services_hold(self, tmp_path): + """Re-running an install must not trip over the services it is about to replace.""" + services = FakeServices() + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + assert _run(argv, _native_context(services)) == 0 + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as held: + held.bind(("127.0.0.1", 0)) + held.listen(1) + envfile.update(tmp_path / ".env", {"DOCSGPT_PORT": str(held.getsockname()[1])}) + assert _run(["up", "--dir", str(tmp_path), "--yes"], _native_context(services)) == 0 def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): From 0d153d3c0de23e06963f1f575cdcc2d54d83e8d5 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 00:12:40 +0100 Subject: [PATCH 047/130] fix: tie the busy-port exemption to the port the install is recorded on MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The exemption asked whether any service of the install was running, so moving an install onto a different port that something else held would pass the check and then fail to bind, with the health poll answered by whatever owned that port — the false success the check exists to prevent. install.json carries the API port now, and a busy port is allowed only when it is that port and the API service is running. --- docsgpt/deploy/commands.py | 23 ++++++++++++++++------- tests/deploy/test_native.py | 25 +++++++++++++++++++++---- 2 files changed, 37 insertions(+), 11 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 5ba2c4e3..c03ccecb 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -298,17 +298,26 @@ def _port_is_free(port: int) -> bool: def _refuse_busy_port(directory: Path, port: int, services, names: tuple[str, str]) -> None: """Whatever holds the port would answer the health check while these services failed to bind. - An install re-running on its own port is the exception: its services are holding it, and - starting them again replaces them. + The one exception is this install's own API holding the port it is recorded on. Ownership is not + inferred from the install having some service running: asking for a different port that something + else holds is how a move to a new port would look, and the API would fail to bind it. """ if _port_is_free(port): return - if _record(directory).get("mode") == "native" and any(services.is_running(name) for name in names): + record = _record(directory) + api_service = names[0] + owns_the_port = ( + record.get("mode") == "native" + and str(record.get("port") or "") == str(port) + and services.is_running(api_service) + ) + if owns_the_port: return raise DeployError( - f"port {port} is already in use by something other than this install. The API would fail to " - f"bind it while the health check answered from whatever holds it, so the install would look " - f"healthy and be dead. Free the port, or give this install another one with --port." + f"port {port} is already in use by something other than this install's API. It would fail to " + f"bind while the health check answered from whatever holds the port, so the install would " + f"look healthy and be dead. Free the port, stop this install first with `docsgpt down` if it " + f"is the one holding it on another port, or choose another port with --port." ) @@ -406,7 +415,7 @@ def _native_up(args, context: Context, directory: Path) -> int: # that status, down and uninstall can see and clean up, rather than orphaned units. now = datetime.now(timezone.utc).isoformat(timespec="seconds") record = record or {"installed_at": now} - record.update(version=context.version, mode="native", updated_at=now) + record.update(version=context.version, mode="native", port=port, updated_at=now) (directory / stack.RECORD_FILE).write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") for unit in native.units_for(directory, launcher, port, directory, names): diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index ce66744f..d5a41696 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -295,16 +295,33 @@ class TestNativeUp: _run(argv, _native_context()) assert not (tmp_path / ".env").exists(), "it refuses before writing anything" - def test_an_install_may_keep_the_port_its_own_services_hold(self, tmp_path): - """Re-running an install must not trip over the services it is about to replace.""" + def test_an_install_may_keep_the_port_its_own_api_is_recorded_on(self, tmp_path): + """Re-running an install must not trip over the API it is about to replace.""" + services = FakeServices() + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: + probe.bind(("127.0.0.1", 0)) + port = probe.getsockname()[1] + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--port", str(port), + "--postgres-uri", "postgresql://localhost/d"] + assert _run(argv, _native_context(services)) == 0 + assert json.loads((tmp_path / "install.json").read_text())["port"] == port + + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as held: + held.bind(("127.0.0.1", port)) + held.listen(1) + assert _run(["up", "--dir", str(tmp_path), "--yes"], _native_context(services)) == 0 + + def test_moving_an_install_onto_a_busy_port_is_refused(self, tmp_path): + """Ownership is of one port, not of any port while some service of the install runs.""" services = FakeServices() argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] assert _run(argv, _native_context(services)) == 0 with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as held: held.bind(("127.0.0.1", 0)) held.listen(1) - envfile.update(tmp_path / ".env", {"DOCSGPT_PORT": str(held.getsockname()[1])}) - assert _run(["up", "--dir", str(tmp_path), "--yes"], _native_context(services)) == 0 + elsewhere = held.getsockname()[1] + with pytest.raises(DeployError, match="already in use"): + _run(["up", "--dir", str(tmp_path), "--yes", "--port", str(elsewhere)], _native_context(services)) def test_a_database_url_is_required(self, tmp_path): with pytest.raises(DeployError, match="--postgres-uri"): From 147bb352eced358c07c8e5167a23faa13c7cf976 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 00:23:24 +0100 Subject: [PATCH 048/130] fix: report a Redis database number too long for int() to read The ASCII-digit check accepts any length, but since 3.11 Python refuses to convert a digit string past its conversion limit, so a long one raised ValueError straight through the CLI instead of the message every other unusable URL gets. --- docsgpt/deploy/commands.py | 10 +++++++++- tests/deploy/test_native.py | 7 +++++++ 2 files changed, 16 insertions(+), 1 deletion(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index c03ccecb..19babb5a 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -227,7 +227,15 @@ def _redis_urls(base: str) -> dict[str, str]: f"the Redis URL {base!r} has {path!r} where a database number would go. Pass a URL like " "redis://host:6379 or redis://host:6379/5." ) - first = int(path) if path else 0 + try: + first = int(path) if path else 0 + except ValueError as exc: + # Python refuses to convert a digit string past its conversion limit, and that is a typo + # rather than a crash. + raise DeployError( + f"the Redis URL {base!r} has a database number too long to read. Pass a URL like " + "redis://host:6379 or redis://host:6379/5." + ) from exc return { key: urlunsplit((parts.scheme, parts.netloc, f"/{first + offset}", parts.query, parts.fragment)) for key, offset in (("CELERY_BROKER_URL", 0), ("CELERY_RESULT_BACKEND", 1), ("CACHE_REDIS_URL", 2)) diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index d5a41696..33a50568 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -224,6 +224,13 @@ class TestNativeUp: assert env["CELERY_BROKER_URL"] == "redis://[::1]:6379/2" assert env["CACHE_REDIS_URL"] == "redis://[::1]:6379/4" + def test_a_redis_database_number_too_long_to_convert_is_reported(self, tmp_path): + """Since 3.11 int() refuses a digit string past its limit, which the digit check let through.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d", + "--redis-url", "redis://localhost:6379/" + "1" * 5000] + with pytest.raises(DeployError, match="too long to read"): + _run(argv, _native_context()) + @pytest.mark.parametrize("url", ["localhost:6379", "redis+socket:///var/run/redis.sock", "redis://localhost:6379/queue"]) def test_a_redis_url_that_cannot_be_numbered_is_refused(self, tmp_path, url): From 51b2fa38de42dba88cccb762dc8da60a0043a56f Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 00:34:42 +0100 Subject: [PATCH 049/130] fix: refuse a Redis URL whose port cannot be used urlsplit accepts an authority such as localhost:notaport or localhost:65536 and urlunsplit rebuilds it verbatim; only parts.port raises, and nothing read it. The unusable value reached .env, where the worker picked it up and failed to start its broker, while `up` reported success because it waits only on the API health endpoint. --- docsgpt/deploy/commands.py | 7 +++++++ tests/deploy/test_native.py | 9 +++++++++ 2 files changed, 16 insertions(+) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 19babb5a..db114639 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -215,6 +215,13 @@ def _redis_urls(base: str) -> dict[str, str]: except ValueError as exc: # urlsplit raises on things like redis://[::1 ; that is a typo, not a crash. raise DeployError(f"the Redis URL {base!r} could not be read: {exc}") from exc + try: + _ = parts.port # a non-numeric or out-of-range port raises here, not when the URL is split + except ValueError as exc: + raise DeployError( + f"the Redis URL {base!r} has an unusable port: {exc}. Pass a URL like " + "redis://host:6379 or redis://host:6379/5." + ) from exc if parts.scheme not in ("redis", "rediss"): raise DeployError( f"the Redis URL {base!r} should start with redis:// or rediss://, with any options as " diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index 33a50568..a9335068 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -231,6 +231,15 @@ class TestNativeUp: with pytest.raises(DeployError, match="too long to read"): _run(argv, _native_context()) + @pytest.mark.parametrize("url", ["redis://localhost:notaport/0", "redis://localhost:65536/0"]) + def test_a_redis_url_with_an_unusable_port_is_refused(self, tmp_path, url): + """urlsplit accepts it and rebuilds it verbatim; only .port notices, and the worker dies later.""" + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", + "--postgres-uri", "postgresql://localhost/d", "--redis-url", url] + with pytest.raises(DeployError, match="unusable port"): + _run(argv, _native_context()) + assert not (tmp_path / ".env").exists(), "it refuses before writing settings the worker would read" + @pytest.mark.parametrize("url", ["localhost:6379", "redis+socket:///var/run/redis.sock", "redis://localhost:6379/queue"]) def test_a_redis_url_that_cannot_be_numbered_is_refused(self, tmp_path, url): From 06f233925b3e8dd078d212d5238c71269042d0eb Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 00:44:19 +0100 Subject: [PATCH 050/130] fix: refuse a Redis URL on port zero parts.port returns 0 rather than raising, since 0 is inside the range it checks, so the URL reached .env and the worker and cache had nothing to connect to. --- docsgpt/deploy/commands.py | 8 +++++++- tests/deploy/test_native.py | 3 ++- 2 files changed, 9 insertions(+), 2 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index db114639..0d95655f 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -216,12 +216,18 @@ def _redis_urls(base: str) -> dict[str, str]: # urlsplit raises on things like redis://[::1 ; that is a typo, not a crash. raise DeployError(f"the Redis URL {base!r} could not be read: {exc}") from exc try: - _ = parts.port # a non-numeric or out-of-range port raises here, not when the URL is split + port = parts.port # a non-numeric or out-of-range port raises here, not when the URL is split except ValueError as exc: raise DeployError( f"the Redis URL {base!r} has an unusable port: {exc}. Pass a URL like " "redis://host:6379 or redis://host:6379/5." ) from exc + if port == 0: + # urlsplit is happy with it, since 0 is inside the range, but nothing can connect to it. + raise DeployError( + f"the Redis URL {base!r} has an unusable port: 0. Pass a URL like " + "redis://host:6379 or redis://host:6379/5." + ) if parts.scheme not in ("redis", "rediss"): raise DeployError( f"the Redis URL {base!r} should start with redis:// or rediss://, with any options as " diff --git a/tests/deploy/test_native.py b/tests/deploy/test_native.py index a9335068..3c6341f7 100644 --- a/tests/deploy/test_native.py +++ b/tests/deploy/test_native.py @@ -231,7 +231,8 @@ class TestNativeUp: with pytest.raises(DeployError, match="too long to read"): _run(argv, _native_context()) - @pytest.mark.parametrize("url", ["redis://localhost:notaport/0", "redis://localhost:65536/0"]) + @pytest.mark.parametrize("url", ["redis://localhost:notaport/0", "redis://localhost:65536/0", + "redis://localhost:0/0"]) def test_a_redis_url_with_an_unusable_port_is_refused(self, tmp_path, url): """urlsplit accepts it and rebuilds it verbatim; only .port notices, and the worker dies later.""" argv = ["up", "--native", "--dir", str(tmp_path), "--yes", From c17b23378e2ec7c1292b611871f0f3de69f8f6fb Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:04:01 +0100 Subject: [PATCH 051/130] refactor(settings): split Settings into per-domain modules docsgpt/core/settings.py had grown to 258 fields in one 600-line class, touched by about two commits a week, with related settings scattered (GitHub ingest caps inside the embeddings block, API keys in four places, the OpenAI Responses knobs 100 lines from the other OpenAI fields). It is now a package: one module per domain (auth, llm, embeddings, retrieval, vectorstores, database, workers, ingestion, ocr, storage, connectors, server, events, agents, guardrails, scheduler, sandbox, speech), each a SettingsGroup owning its fields and validators, composed by multiple inheritance into the same flat Settings class. Every attribute name, type, default, alias and constraint is unchanged, so settings.NAME reads, .env files and test monkeypatches all keep working; the import path docsgpt.core.settings is the package. Settings.normalize_api_key is kept as a classmethod for callers that reuse it. The comment above or beside each field became its Field(description=...), so the definitions are visible to tooling; the next commit generates the docs reference from them. Pitfall recorded for future groups: pydantic collects validators by method name across the MRO, so two groups naming a validator the same would silently keep only one. Each group's validator has a unique name. --- AGENTS.md | 2 +- docs/content/Guides/Architecture.mdx | 2 +- docs/content/Guides/compression.md | 2 +- docs/content/quickstart.mdx | 2 +- docs/runbooks/sse-notifications.md | 2 +- docsgpt/core/db_uri.py | 2 +- docsgpt/core/settings.py | 596 -------------------------- docsgpt/core/settings/__init__.py | 77 ++++ docsgpt/core/settings/_shared.py | 37 ++ docsgpt/core/settings/agents.py | 91 ++++ docsgpt/core/settings/auth.py | 87 ++++ docsgpt/core/settings/connectors.py | 50 +++ docsgpt/core/settings/database.py | 36 ++ docsgpt/core/settings/embeddings.py | 87 ++++ docsgpt/core/settings/events.py | 86 ++++ docsgpt/core/settings/guardrails.py | 40 ++ docsgpt/core/settings/ingestion.py | 144 +++++++ docsgpt/core/settings/llm.py | 121 ++++++ docsgpt/core/settings/ocr.py | 91 ++++ docsgpt/core/settings/retrieval.py | 40 ++ docsgpt/core/settings/sandbox.py | 88 ++++ docsgpt/core/settings/scheduler.py | 28 ++ docsgpt/core/settings/server.py | 39 ++ docsgpt/core/settings/speech.py | 31 ++ docsgpt/core/settings/storage.py | 49 +++ docsgpt/core/settings/vectorstores.py | 89 ++++ docsgpt/core/settings/workers.py | 37 ++ tests/test_remaining_coverage.py | 2 +- 28 files changed, 1355 insertions(+), 603 deletions(-) delete mode 100644 docsgpt/core/settings.py create mode 100644 docsgpt/core/settings/__init__.py create mode 100644 docsgpt/core/settings/_shared.py create mode 100644 docsgpt/core/settings/agents.py create mode 100644 docsgpt/core/settings/auth.py create mode 100644 docsgpt/core/settings/connectors.py create mode 100644 docsgpt/core/settings/database.py create mode 100644 docsgpt/core/settings/embeddings.py create mode 100644 docsgpt/core/settings/events.py create mode 100644 docsgpt/core/settings/guardrails.py create mode 100644 docsgpt/core/settings/ingestion.py create mode 100644 docsgpt/core/settings/llm.py create mode 100644 docsgpt/core/settings/ocr.py create mode 100644 docsgpt/core/settings/retrieval.py create mode 100644 docsgpt/core/settings/sandbox.py create mode 100644 docsgpt/core/settings/scheduler.py create mode 100644 docsgpt/core/settings/server.py create mode 100644 docsgpt/core/settings/speech.py create mode 100644 docsgpt/core/settings/storage.py create mode 100644 docsgpt/core/settings/vectorstores.py create mode 100644 docsgpt/core/settings/workers.py diff --git a/AGENTS.md b/AGENTS.md index 90356fc4..6aea7be9 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -198,7 +198,7 @@ vale . - Parsers live in `docsgpt/parser/` and handle different document formats in the ingestion stage. - Agents and tools are in `docsgpt/agents/` and `docsgpt/agents/tools/`. - Celery setup/config lives in `docsgpt/celery_init.py` and `docsgpt/celeryconfig.py`. -- Settings and env vars are managed via Pydantic in `docsgpt/core/settings.py`. +- Settings and env vars are managed via Pydantic in `docsgpt/core/settings/` (one module per domain, composed into `Settings`). Every field needs a `description`; regenerate the docs reference with `python -m docsgpt.core.settings.reference --write`. ### Frontend diff --git a/docs/content/Guides/Architecture.mdx b/docs/content/Guides/Architecture.mdx index a65f3dcb..a25d34e6 100644 --- a/docs/content/Guides/Architecture.mdx +++ b/docs/content/Guides/Architecture.mdx @@ -251,4 +251,4 @@ The main extension points in [arc53/DocsGPT](https://github.com/arc53/DocsGPT) a | Parsing and workers | [`docsgpt/parser/`](https://github.com/arc53/DocsGPT/tree/main/docsgpt/parser), [`docsgpt/worker.py`](https://github.com/arc53/DocsGPT/blob/main/docsgpt/worker.py), [`docsgpt/api/user/tasks.py`](https://github.com/arc53/DocsGPT/blob/main/docsgpt/api/user/tasks.py) | | Models and vector stores | [`docsgpt/llm/`](https://github.com/arc53/DocsGPT/tree/main/docsgpt/llm), [`docsgpt/vectorstore/`](https://github.com/arc53/DocsGPT/tree/main/docsgpt/vectorstore) | | Storage and events | [`docsgpt/storage/`](https://github.com/arc53/DocsGPT/tree/main/docsgpt/storage), [`docsgpt/streaming/`](https://github.com/arc53/DocsGPT/tree/main/docsgpt/streaming), [`docsgpt/events/`](https://github.com/arc53/DocsGPT/tree/main/docsgpt/events) | -| Configuration, UI, and deployment | [`docsgpt/core/settings.py`](https://github.com/arc53/DocsGPT/blob/main/docsgpt/core/settings.py), [`frontend/`](https://github.com/arc53/DocsGPT/tree/main/frontend), [`deployment/`](https://github.com/arc53/DocsGPT/tree/main/deployment) | +| Configuration, UI, and deployment | [`docsgpt/core/settings/`](https://github.com/arc53/DocsGPT/tree/main/docsgpt/core/settings), [`frontend/`](https://github.com/arc53/DocsGPT/tree/main/frontend), [`deployment/`](https://github.com/arc53/DocsGPT/tree/main/deployment) | diff --git a/docs/content/Guides/compression.md b/docs/content/Guides/compression.md index 14b90c62..53bcea57 100644 --- a/docs/content/Guides/compression.md +++ b/docs/content/Guides/compression.md @@ -19,7 +19,7 @@ The compression system operates on a "summarize and truncate" principle: ## Configuration -You can configure the compression behavior in your `.env` file or `docsgpt/core/settings.py`: +You can configure the compression behavior in your `.env` file or `docsgpt/core/settings/agents.py`: | Setting | Default | Description | | :--- | :--- | :--- | diff --git a/docs/content/quickstart.mdx b/docs/content/quickstart.mdx index 3b134fe6..a1702b3e 100644 --- a/docs/content/quickstart.mdx +++ b/docs/content/quickstart.mdx @@ -124,6 +124,6 @@ To work from the source tree, for example to build the images yourself, use `set ## Advanced Configuration -For more advanced customization of DocsGPT settings, such as configuring vector stores, embedding models, and other parameters, please refer to the [DocsGPT Settings documentation](/Deploying/DocsGPT-Settings). This guide explains how to modify the `.env` file or `settings.py` for deeper configuration. +For more advanced customization of DocsGPT settings, such as configuring vector stores, embedding models, and other parameters, please refer to the [DocsGPT Settings documentation](/Deploying/DocsGPT-Settings). This guide explains how to configure DocsGPT through the `.env` file, and links to the full settings reference. Enjoy using DocsGPT! diff --git a/docs/runbooks/sse-notifications.md b/docs/runbooks/sse-notifications.md index 1d100c43..6afaf225 100644 --- a/docs/runbooks/sse-notifications.md +++ b/docs/runbooks/sse-notifications.md @@ -320,7 +320,7 @@ redis-cli -n 2 DEL user::stream ## Settings reference -Everything in `docsgpt/core/settings.py`: +Everything in `docsgpt/core/settings/events.py`: | Setting | Default | Purpose | | --------------------------------------------- | ------- | --------------------------------------------- | diff --git a/docsgpt/core/db_uri.py b/docsgpt/core/db_uri.py index 99e93bc3..cb875b4e 100644 --- a/docsgpt/core/db_uri.py +++ b/docsgpt/core/db_uri.py @@ -15,7 +15,7 @@ have to know which driver a given field feeds. Each normalizer also silently upgrades the legacy ``postgresql+psycopg2://`` prefix since psycopg2 is no longer in the project. -This module is deliberately separate from ``docsgpt/core/settings.py`` +This module is deliberately separate from ``docsgpt/core/settings`` so the Settings class stays focused on field declarations, and the URI-rewriting logic can be unit-tested without triggering ``.env`` file loading from importing Settings. diff --git a/docsgpt/core/settings.py b/docsgpt/core/settings.py deleted file mode 100644 index 8bf864d0..00000000 --- a/docsgpt/core/settings.py +++ /dev/null @@ -1,596 +0,0 @@ -import os -from typing import Optional - -from pydantic import AliasChoices, Field, field_validator -from pydantic_settings import BaseSettings, SettingsConfigDict - -from docsgpt.core.db_uri import ( - normalize_pgvector_connection_string, - normalize_postgres_uri, -) -from docsgpt.core.paths import env_file, home_dir - -# Runtime data home (DOCSGPT_HOME, the checkout, or cwd); see docsgpt.core.paths. -current_dir = str(home_dir()) - - -class Settings(BaseSettings): - model_config = SettingsConfigDict(extra="ignore") - - AUTH_TYPE: Optional[str] = None # simple_jwt, session_jwt, oidc, or None - - # OIDC SSO (AUTH_TYPE=oidc) — any OpenID Connect IdP with discovery (Authentik, Keycloak, ...) - OIDC_ISSUER: Optional[str] = None # e.g. https://auth.example.com/application/o/docsgpt/ - OIDC_CLIENT_ID: Optional[str] = None - OIDC_CLIENT_SECRET: Optional[str] = None # optional; PKCE is always used - OIDC_SCOPES: str = "openid profile email" - OIDC_USER_ID_CLAIM: str = "sub" # ID-token claim mapped to the DocsGPT user id - OIDC_FRONTEND_URL: Optional[str] = None # browser-facing app origin, e.g. http://localhost:5173 - OIDC_REDIRECT_URI: Optional[str] = None # override; default /api/auth/oidc/callback - OIDC_SESSION_LIFETIME_SECONDS: int = 28800 # minted session JWT lifetime (8h) - OIDC_PROVIDER_NAME: Optional[str] = None # sign-in button label, e.g. "Acme SSO" - OIDC_ALLOWED_GROUPS: Optional[str] = None # comma-separated allowlist; unset = any authenticated user - OIDC_GROUPS_CLAIM: str = "groups" # ID-token/userinfo claim carrying group membership - OIDC_ADMIN_GROUPS: Optional[str] = None # comma-separated groups granted admin; unset = no OIDC admin mapping - - # RBAC: persisted admin grants live in user_roles (AUTH_TYPE=oidc only). This is the - # only non-DB admin path, for AUTH_TYPE=None self-host. MUST stay False if networked. - LOCAL_MODE_ADMIN: bool = False - - # SCIM 2.0 provisioning (IdP-driven user create/deactivate at /scim/v2) - SCIM_ENABLED: bool = False - SCIM_TOKEN: Optional[str] = None # bearer token for IdP SCIM clients (required when enabled) - - LLM_PROVIDER: str = "docsgpt" - LLM_NAME: Optional[str] = None # if LLM_PROVIDER is openai, LLM_NAME can be gpt-4 or gpt-3.5-turbo - # Legacy model on purpose: an install that never pinned this has vectors from it, and - # granite is the same width so a swap would fail silently. New installs get granite from - # .env-template; existing ones switch by setting this and running docsgpt.scripts.reembed. - EMBEDDINGS_NAME: str = "huggingface_sentence-transformers/all-mpnet-base-v2" - EMBEDDINGS_BASE_URL: Optional[str] = None # Remote embeddings API URL (OpenAI-compatible) - EMBEDDINGS_KEY: Optional[str] = None # api key for embeddings (if using openai, just copy API_KEY) - EMBEDDINGS_MAX_INPUT_TOKENS: Optional[int] = None # truncate each remote embed input to N tokens (overflow lost) - EMBEDDINGS_BATCH_SIZE: int = 32 # chunks per store transaction / remote embed request - # Documents per local ONNX forward pass. Each pass pads to its longest input, and that - # waste grows with the square of chunk length: at 1250 tokens, 32 peaked at 6.6 GB, 1 at 2.9 GB. - EMBEDDINGS_MODEL_BATCH_SIZE: int = 1 - # Intra-op threads for the local ONNX runner; None = every core. It scales sub-linearly, - # so several single-threaded workers beat one many-threaded process on the same cores. - EMBEDDINGS_THREADS: Optional[int] = None - # Embedding models and their tokenizers. Persistent by default: FastEmbed's own default is the temp dir. - EMBEDDINGS_CACHE_DIR: Optional[str] = Field(default_factory=lambda: str(home_dir() / "models")) - # Pooling ("cls"/"mean") and L2 normalisation. Read from the model's own repository; - # set these only for a repository that declares neither, or to override what it declares. - EMBEDDINGS_POOLING: Optional[str] = None - EMBEDDINGS_NORMALIZE: Optional[bool] = None - # Embed on the worker so the API holds no model (~890 MB), at one broker round trip per - # query. Ignored when EMBEDDINGS_BASE_URL is set, which is the better answer for production. - EMBEDDINGS_DELEGATE_TO_WORKER: bool = True - EMBEDDINGS_QUEUE: str = "embeddings" # queue the embed task is routed to - EMBEDDINGS_DELEGATE_TIMEOUT: int = 60 # seconds to wait for the worker - GITHUB_INGEST_MAX_FILE_BYTES: int = 1048576 # skip repo blobs larger than this (0 = no cap) - GITHUB_INGEST_MAX_WORKERS: int = 8 # parallel file fetches per GitHub repo ingest - # Operator-supplied model YAMLs, loaded after the built-in catalog; later wins on - # duplicate model id. See docsgpt/core/models/README.md. - MODELS_CONFIG_DIR: Optional[str] = None - - CELERY_BROKER_URL: str = "redis://localhost:6379/0" - CELERY_RESULT_BACKEND: str = "redis://localhost:6379/1" - # Prefetch=1 caps SIGKILL loss to one task. Visibility timeout must exceed the longest - # legitimate task runtime but stay short enough that SIGKILLed tasks redeliver promptly. - CELERY_WORKER_PREFETCH_MULTIPLIER: int = 1 - CELERY_VISIBILITY_TIMEOUT: int = 3600 - # Recycle a prefork child past this resident size in KB; backstops docling/torch heap growth. - # Checked between tasks, so it does not bound the peak within one. 0 disables. - CELERY_WORKER_MAX_MEMORY_PER_CHILD: int = 4194304 - CELERY_WORKER_MAX_TASKS_PER_CHILD: int = 0 # recycle after N tasks; 0 disables - # Only consulted when VECTOR_STORE=mongodb or when running scripts/db/backfill.py; user data lives in Postgres. - MONGO_URI: Optional[str] = None - # User-data Postgres DB. - POSTGRES_URI: Optional[str] = None - # On startup, apply pending Alembic migrations. Disable if you manage schema out-of-band. - AUTO_MIGRATE: bool = True - # On startup, create the target Postgres database if missing (needs CREATEDB privilege). - AUTO_CREATE_DB: bool = True - # On startup, create the pgvector/graph tables and verify the embedding dimension. No Alembic - # migration covers the vector DB (it may be a separate cluster); set False to manage it yourself. - AUTO_VECTOR_SCHEMA: bool = True - LLM_PATH: str = os.path.join(current_dir, "models/docsgpt-7b-f16.gguf") - DEFAULT_MAX_HISTORY: int = 150 - DEFAULT_LLM_TOKEN_LIMIT: int = 128000 # Fallback when model not found in registry - RESERVED_TOKENS: dict = { - "system_prompt": 500, - "current_query": 500, - "safety_buffer": 1000, - } - DEFAULT_AGENT_LIMITS: dict = { - "token_limit": 50000, - "request_limit": 500, - } - UPLOAD_FOLDER: str = "inputs" - # Serve the web UI shipped in the package (docsgpt/static) from the API process. - SERVE_UI: bool = True - # Request cap is applied by Flask before multipart parsing; the per-file cap also while copying. - UPLOAD_MAX_REQUEST_BYTES: int = Field(default=256 * 1024 * 1024, gt=0) - UPLOAD_MAX_FILE_BYTES: int = Field(default=100 * 1024 * 1024, gt=0) - PARSE_SPEC_MAX_BYTES: int = Field(default=10 * 1024 * 1024, gt=0) - # ZIP limits apply cumulatively across nested archives in one extraction. - UPLOAD_MAX_ARCHIVE_BYTES: int = Field(default=250 * 1024 * 1024, gt=0) - UPLOAD_MAX_ARCHIVE_FILES: int = Field(default=10_000, gt=0) - UPLOAD_MAX_ARCHIVE_RATIO: int = Field(default=1000, gt=0) - UPLOAD_MAX_ARCHIVE_DEPTH: int = Field(default=3, ge=0) - PARSE_PDF_AS_IMAGE: bool = False - PARSE_IMAGE_REMOTE: bool = False - # Document parser for source ingestion, chat attachments and the - # read_document tool. "anydoc" (default): firecrawl-anydoc, a Rust - # converter with no ML models — milliseconds per file, ~100 MB peak RSS. - # "docling": the layout/table-model pipeline (optional install; needed - # for read_document's structured output and the docling OCR backend). - # Files anydoc cannot convert (scanned PDFs, malformed input) fall back to - # docling when it is installed, otherwise to the native OCR parsers (OCR - # on) or the legacy parsers. Rollback to the previous behaviour is this - # one variable. - DOC_PARSER_ENGINE: str = "anydoc" - # OCR for scanned PDFs and images. OCR_ENABLED covers source ingestion, - # OCR_ATTACHMENTS_ENABLED chat attachments. Which stack performs it is - # OCR_BACKEND; which engine, OCR_ENGINE. The DOCLING_OCR_* names are the - # pre-2026-09 spellings and stay accepted as aliases. - OCR_ENABLED: bool = Field( - default=False, validation_alias=AliasChoices("OCR_ENABLED", "DOCLING_OCR_ENABLED") - ) - OCR_ATTACHMENTS_ENABLED: bool = Field( - default=False, - validation_alias=AliasChoices("OCR_ATTACHMENTS_ENABLED", "DOCLING_OCR_ATTACHMENTS_ENABLED"), - ) - # Which stack runs OCR when it is on: - # auto — docling when installed, otherwise native. - # docling — the layout-model pipeline (hybrid region OCR, reading order, - # table structure); needs the optional docling extra. - # native — pypdfium2/Pillow page rendering straight into tesseract or a - # DeepSeek-OCR endpoint (docsgpt/parser/file/ocr_parser.py). - # No ML models in the worker; tables come out as text lines - # under tesseract. - OCR_BACKEND: str = "auto" - # Pages docling's threaded pipeline buffers in flight; the library - # default (100) drives worker RSS to ~3 GB on a mid-size PDF. - DOCLING_PIPELINE_QUEUE_MAX_SIZE: int = 2 - DOCLING_COMPILE_TORCH_MODELS: bool = False - DOCLING_TABULAR_MAX_BYTES: int = 2_000_000 - DOCLING_MARKUP_MAX_BYTES: int = 8_000_000 - # HTML/XHTML larger than this (bytes) are head-truncated before the - # markdownify parser runs (the anydoc engine's HTML path). The tree that - # path builds costs ~50x the input — 30 MB of HTML measured at 1.6 GB RSS — - # and the upload cap is 100 MB, so the gate is what keeps one upload from - # taking the ingest worker down. 0 disables it. - MARKUP_MAX_BYTES: int = 8_000_000 - # Trust-check anydoc's PDF output (docsgpt/parser/file/pdf_trust.py): - # flag composite (Type0) fonts without a ToUnicode map, and CJK-declaring - # PDFs whose extracted text has almost no CJK — the two classes where - # anydoc drops text silently. A flagged file re-parses on the docling - # fallback when docling is installed; otherwise the anydoc output is kept - # and the document gets extra_info["parse_warnings"]. ~30 ms per scanned MB. - PDF_TRUST_CHECK: bool = True - # Rewrite dot-leader / whitespace-aligned table runs in anydoc's PDF - # markdown into GFM tables (docsgpt/parser/file/tableize.py). Off by - # default: it rewrites content on a heuristic (>=3 uniform label+numbers - # lines) validated only on a small corpus so far. - ANYDOC_TABLEIZE: bool = False - # OCR engine used when OCR is on (OCR_ENABLED / OCR_ATTACHMENTS_ENABLED). - # Benched 2026-08 on EN/ZH/table/degraded scans (docs/Guides/ocr has the - # menu): - # tesseract — recommended: best classic-engine accuracy (perfect EN word - # recall, 0.000 bilingual CER, 100% table cells), ~35 MB, CPU-only. - # Needs the system binary + language packs: an optional install like - # every OCR dependency (build with INSTALL_TESSERACT=true, or apt/brew - # install tesseract-ocr for a local run). Both backends. - # deepseek — DeepSeek-OCR against an Ollama/vLLM endpoint - # (OCR_DEEPSEEK_*). Best table/CJK quality; the worker stays light - # (no layout models) but each page costs seconds on the model server. - # Both backends. - # auto — docling's pick: ocrmac on macOS (excellent), rapidocr on Linux - # (silently shreds some long text lines — avoid as a server default). - # ocrmac | rapidocr — force one of those. - # auto/ocrmac/rapidocr exist only inside docling; the native backend runs - # tesseract for them. An engine that is not installed degrades (docling: - # to "auto") with a warning instead of failing the parse. - OCR_ENGINE: str = "tesseract" - # Tesseract language packs, "+"-separated (e.g. "eng+chi_sim+deu"). Other - # engines keep their own defaults — their language codes differ. - OCR_LANGS: str = "eng" - OCR_DEEPSEEK_URL: str = "http://localhost:11434/v1/chat/completions" - OCR_DEEPSEEK_MODEL: str = "deepseek-ocr:3b" - # Seconds allowed per page request to the DeepSeek endpoint, on both - # backends (native sends pages one at a time; docling's VLM pipeline - # keeps its own concurrency). A 3B model on a laptop needs minutes; a - # vLLM GPU deployment, seconds. - OCR_DEEPSEEK_TIMEOUT: float = 300.0 - # Native backend only: resolution at which pages without a text layer are - # rendered before OCR. 200 suits tesseract; clamped to 72-600. - OCR_RENDER_DPI: int = 200 - # Chars-per-page floor below which an OCR'd PDF/image parse is treated as an OCR - # dropout (long-running docling workers were observed returning zero characters for - # every scanned page after a long scanned PDF, with no error) rather than as content. - # docling retries once on a fresh full-page-OCR converter; both backends then fail - # loudly instead of indexing an empty document. 0 disables the guard. - OCR_MIN_CHARS_PER_PAGE: int = Field( - default=20, validation_alias=AliasChoices("OCR_MIN_CHARS_PER_PAGE", "DOCLING_OCR_MIN_CHARS_PER_PAGE") - ) - # Read PDF *attachments* via their embedded text layer (pypdfium2) instead - # of docling, falling back to docling when there is no text layer to read. - # Attachments go into a prompt, so docling's structural markdown earns far - # less than the tens of seconds per file it costs; source ingestion is - # unaffected and always uses docling, because chunking and retrieval do - # depend on that structure. - ATTACHMENT_PDF_TEXT_FAST_PATH: bool = True - # Median chars per sampled page below which a PDF is treated as a scan and handed to docling. - # Measured on real uploads: scans at 0-17 chars/page, text-layer documents at 433-6834. - ATTACHMENT_PDF_TEXT_MIN_MEDIAN_CHARS: int = 32 - ATTACHMENT_TEXT_MAX_BYTES: int = 5_000_000 - AGENT_IMAGE_MAX_BYTES: int = 5_000_000 - AGENT_IMAGE_MAX_PIXELS: int = 16_777_216 - VECTOR_STORE: str = "faiss" # "faiss" or "elasticsearch" or "qdrant" or "milvus" or "lancedb" or "pgvector" - # Retriever keys an agent may use; must match RetrieverCreator.retrievers registry keys, - # NOT the legacy ``classic_rag`` label which never matched the registry. - RETRIEVERS_ENABLED: list = ["classic", "default"] - # Concurrent per-source searches in one retrieval; the query is embedded once and shared. - RETRIEVAL_MAX_PARALLEL_SOURCES: int = 4 - # Kill-switch for per-source retrieval dispatch; False collapses to a single retriever. - PER_SOURCE_RETRIEVAL_ENABLED: bool = True - GRAPHRAG_ENABLED: bool = False # gates graph-aware ingestion/retrieval - # Model for ingest-time graph extraction; None reuses LLM_PROVIDER/LLM_NAME. - GRAPHRAG_EXTRACTION_MODEL: Optional[str] = None - # Hard cap on chunks extracted per source (cost control). - GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = 2000 - AGENT_NAME: str = "classic" - FALLBACK_LLM_PROVIDER: Optional[str] = None # provider for fallback llm - FALLBACK_LLM_NAME: Optional[str] = None # model name for fallback llm - FALLBACK_LLM_API_KEY: Optional[str] = None # api key for fallback llm - - # Google Drive integration - GOOGLE_CLIENT_ID: Optional[str] = None # Replace with your actual Google OAuth client ID - GOOGLE_CLIENT_SECRET: Optional[str] = None # Replace with your actual Google OAuth client secret - CONNECTOR_REDIRECT_BASE_URI: Optional[str] = ( - "http://127.0.0.1:7091/api/connectors/callback" ##add redirect url as it is to your provider's console(gcp) - ) - # Comma-separated frontend origins allowed to receive connector OAuth results, e.g. https://docsgpt.example.com. - # The callback origin and OIDC_FRONTEND_URL are always allowed; a loopback callback also allows localhost:5173. - CONNECTOR_ALLOWED_ORIGINS: Optional[str] = None - - # Microsoft Entra ID (Azure AD) integration - MICROSOFT_CLIENT_ID: Optional[str] = None # Azure AD Application (client) ID - MICROSOFT_CLIENT_SECRET: Optional[str] = None # Azure AD Application client secret - MICROSOFT_TENANT_ID: Optional[str] = "common" # Azure AD Tenant ID (or 'common' for multi-tenant) - MICROSOFT_AUTHORITY: Optional[str] = None # e.g., "https://login.microsoftonline.com/{tenant_id}" - - # Confluence Cloud integration - CONFLUENCE_CLIENT_ID: Optional[str] = None - CONFLUENCE_CLIENT_SECRET: Optional[str] = None - - # GitHub source - GITHUB_ACCESS_TOKEN: Optional[str] = None # PAT token with read repo access - - # LLM Cache - CACHE_REDIS_URL: str = "redis://localhost:6379/2" - - API_URL: str = "http://localhost:7091" # backend url for celery worker - - # Public base URL for user-facing endpoint references in prompts - PUBLIC_API_BASE_URL: Optional[str] = None - MCP_OAUTH_REDIRECT_URI: Optional[str] = None # public callback URL for MCP OAuth - INTERNAL_KEY: Optional[str] = None # internal api key for worker-to-backend auth - - API_KEY: Optional[str] = None # LLM api key (used by LLM_PROVIDER) - - # Provider-specific API keys (for multi-model support) - OPENAI_API_KEY: Optional[str] = None - ANTHROPIC_API_KEY: Optional[str] = None - GOOGLE_API_KEY: Optional[str] = None - GROQ_API_KEY: Optional[str] = None - HUGGINGFACE_API_KEY: Optional[str] = None - OPEN_ROUTER_API_KEY: Optional[str] = None - NOVITA_API_KEY: Optional[str] = None - - OPENAI_API_BASE: Optional[str] = None # azure openai api base url - OPENAI_API_VERSION: Optional[str] = None # azure openai api version - AZURE_DEPLOYMENT_NAME: Optional[str] = None # azure deployment name for answering - AZURE_EMBEDDINGS_DEPLOYMENT_NAME: Optional[str] = None # azure deployment name for embeddings - OPENAI_BASE_URL: Optional[str] = None # openai base url for open ai compatable models - - # elasticsearch - ELASTIC_CLOUD_ID: Optional[str] = None # cloud id for elasticsearch - ELASTIC_USERNAME: Optional[str] = None # username for elasticsearch - ELASTIC_PASSWORD: Optional[str] = None # password for elasticsearch - ELASTIC_URL: Optional[str] = None # url for elasticsearch - ELASTIC_INDEX: Optional[str] = "docsgpt" # index name for elasticsearch - - # Legacy AWS credentials from the retired SageMaker provider. Still read as a deprecated - # fallback by S3 storage; do not use for new deployments. - SAGEMAKER_REGION: Optional[str] = None - SAGEMAKER_ACCESS_KEY: Optional[str] = None - SAGEMAKER_SECRET_KEY: Optional[str] = None - - # Qdrant vectorstore config - QDRANT_COLLECTION_NAME: Optional[str] = "docsgpt" - QDRANT_LOCATION: Optional[str] = None - QDRANT_URL: Optional[str] = None - QDRANT_PORT: Optional[int] = 6333 - QDRANT_GRPC_PORT: int = 6334 - QDRANT_PREFER_GRPC: bool = False - QDRANT_HTTPS: Optional[bool] = None - QDRANT_API_KEY: Optional[str] = None - QDRANT_PREFIX: Optional[str] = None - QDRANT_TIMEOUT: Optional[float] = None - QDRANT_HOST: Optional[str] = None - QDRANT_PATH: Optional[str] = None - QDRANT_DISTANCE_FUNC: str = "Cosine" - - # PGVector config. postgres://, postgresql:// and postgresql+psycopg:// are all accepted - # and normalized internally for psycopg.connect(). - PGVECTOR_CONNECTION_STRING: Optional[str] = None - PGVECTOR_POOL_MAX_SIZE: int = 8 # per-process pool; 0 = one direct connection per store - # IVFFlat probes; None derives sqrt(lists) from the index. Higher = better recall, more scan. - PGVECTOR_IVFFLAT_PROBES: Optional[int] = None - # Milvus vectorstore config - MILVUS_COLLECTION_NAME: Optional[str] = "docsgpt" - # milvus-lite (embedded) database file, under the data home like the other local stores - MILVUS_URI: Optional[str] = Field(default_factory=lambda: str(home_dir() / "milvus_local.db")) - MILVUS_TOKEN: Optional[str] = "" - - # LanceDB vectorstore config - LANCEDB_PATH: str = Field(default_factory=lambda: str(home_dir() / "data" / "lancedb")) # LanceDB local data - LANCEDB_TABLE_NAME: Optional[str] = "docsgpts" # Name of the table to use for storing vectors - - FLASK_DEBUG_MODE: bool = False - STORAGE_TYPE: str = "local" # local or s3 - - # S3-compatible object storage (STORAGE_TYPE=s3): AWS S3, MinIO, R2, B2, Spaces, ... - # For non-AWS, set S3_ENDPOINT_URL and usually S3_PATH_STYLE=true. - S3_BUCKET_NAME: str = "docsgpt-test-bucket" - S3_ENDPOINT_URL: Optional[str] = None # custom endpoint for S3-compatible services; omit for AWS - S3_ACCESS_KEY_ID: Optional[str] = None - S3_SECRET_ACCESS_KEY: Optional[str] = None - S3_REGION: Optional[str] = None # AWS region; use "auto" for Cloudflare R2 - S3_PATH_STYLE: bool = False # path-style addressing (required by most non-AWS services) - - # Anonymous startup version check for security issues. - VERSION_CHECK: bool = True - URL_STRATEGY: str = "backend" # backend or s3 - - JWT_SECRET_KEY: str = "" - - # Encryption settings - ENCRYPTION_SECRET_KEY: str = "default-docsgpt-encryption-key" - - TTS_PROVIDER: str = "google_tts" # google_tts, elevenlabs, or none to switch text-to-speech off - ELEVENLABS_API_KEY: Optional[str] = None - STT_PROVIDER: str = "openai" # openai, faster_whisper, or none to switch speech-to-text off - OPENAI_STT_MODEL: str = "gpt-4o-mini-transcribe" - STT_LANGUAGE: Optional[str] = None - STT_MAX_FILE_SIZE_MB: int = 50 - STT_ENABLE_TIMESTAMPS: bool = False - STT_ENABLE_DIARIZATION: bool = False - - # Tool pre-fetch settings - ENABLE_TOOL_PREFETCH: bool = True - - # True persists Responses API calls server-side so previous_response_id can chain turns. - # False keeps them stateless, carrying reasoning across the tool loop as encrypted items. - OPENAI_RESPONSES_STORE: bool = False - # Cross-turn ``previous_response_id`` chaining (store mode only). The - # chained transcript lives on the provider and is invisible to every - # local guard, so it is bounded: a turn starts from the local history - # when the previous turn's reported prompt already reached the budget - # (default: the model's context window) or when the conversation was - # compressed after that turn was produced. - OPENAI_RESPONSES_CHAIN_ACROSS_TURNS: bool = True - OPENAI_RESPONSES_CHAIN_BUDGET_TOKENS: Optional[int] = None - # ``truncation: "auto"`` lets the provider drop the oldest input items - # instead of failing every request once a chain exceeds the model's window. - OPENAI_RESPONSES_TRUNCATION_AUTO: bool = False - # Prompt-cache hints on the Responses API: route a user's calls to the - # same cache shard (opaque per-user key), and request extended retention - # where offered. - OPENAI_PROMPT_CACHE_KEY: bool = True - OPENAI_PROMPT_CACHE_RETENTION: Optional[str] = None - OPENAI_REASONING_SUMMARY: str = "auto" - - # Lets OpenAI-compatible clients identify a logical chat by session header, which - # chat-completions itself has no field for. - V1_SESSION_TTL_SECONDS: int = 24 * 60 * 60 - # Optional cheaper model for conversation titles; unset reuses the answer model. - TITLE_MODEL_ID: Optional[str] = None - - # Config-free tools on by default in agentless chats. ``scheduler`` is dual-registered in - # BUILTIN_AGENT_TOOLS so one synthetic id resolves via defaults or the agent picker. - # Add "code_executor" and "artifact_generator" once a sandbox runner is configured — both - # execute through it and would fail on every call without one. - DEFAULT_CHAT_TOOLS: list = [ - "memory", - "read_webpage", - "scheduler", - ] - - # Conversation Compression Settings - ENABLE_CONVERSATION_COMPRESSION: bool = True - COMPRESSION_THRESHOLD_PERCENTAGE: float = 0.8 # Trigger at 80% of context - COMPRESSION_MODEL_OVERRIDE: Optional[str] = None # Use different model for compression - COMPRESSION_PROMPT_VERSION: str = "v1.0" # Track prompt iterations - COMPRESSION_MAX_HISTORY_POINTS: int = 3 # Keep only last N compression points to prevent DB bloat - # Per-field cap on the verbatim tail kept after a compression point (0 disables). - COMPRESSION_RECENT_FIELD_MAX_TOKENS: int = 8000 - # Cap on one tool result entering the LLM context (0 disables); journal/DB keep it whole. - TOOL_RESULT_MAX_TOKENS: int = 20000 - - # Agent Guardrails - GUARDRAILS_ENABLED: bool = True # master switch; False disables every stage - # Allowlist of GuardrailCreator.checks keys; empty means every registered check. - GUARDRAILS_CHECKS_ENABLED: list = [] - # A GuardrailsConfig fragment every agent inherits and cannot weaken; agents may add - # controls or make an action stricter, never looser. "enabled" is required — without it - # the floor parses but applies to nothing. Example: - # {"enabled": true, "mode": "scan_all", - # "controls": [{"check": "secrets", "stage": "output", "action": "redact"}]} - GUARDRAILS_FLOOR: dict = {} - # Judge model for the topic/policy checks; None reuses the request's model. - GUARDRAILS_JUDGE_MODEL: Optional[str] = None - # Persist scanned text alongside guardrail_events. Off by default: pre-redaction text is - # exactly the material a PII control exists to keep out of storage. - GUARDRAILS_STORE_SCANNED_TEXT: bool = False - GUARDRAILS_EVENTS_RETENTION_DAYS: int = Field(default=30, ge=1) - - # Internal SSE push channel (notifications + durable replay journal). - # False makes /api/events emit "push_disabled" and return; clients fall back to polling. - ENABLE_SSE_PUSH: bool = True - # Per-user durable backlog cap in entries; ~24h of replay at typical rates. - EVENTS_STREAM_MAXLEN: int = 1000 - # Bounds uvicorn's shutdown drain (uvicorn_worker doesn't forward --graceful-timeout). - # Keep below the gunicorn --timeout (180) watchdog. Used by BoundedDrainUvicornWorker. - GRACEFUL_SHUTDOWN_TIMEOUT_SECONDS: int = 30 - WSGI_THREADPOOL_WORKERS: int = 96 - SSE_KEEPALIVE_SECONDS: int = Field(default=15, ge=1) - # Simultaneous SSE connections per user; each holds a pooled async Redis connection for - # its lifetime. 8 covers multi-tab use without one user starving the pool. 0 disables. - SSE_MAX_CONCURRENT_PER_USER: int = 8 - # Pool size of the async Redis client behind the event-loop routes, per process. Every - # open notification tab, chat reconnect and device session holds one connection, so this - # caps concurrent streams per worker (redis-py's own default is 100). Keep the total - # across workers below the Redis server's maxclients (10000 by default). - ASYNC_REDIS_MAX_CONNECTIONS: int = Field(default=2000, ge=1) - # Backlog entries XRANGE returns per /api/events snapshot. Bounds what one replay moves - # from Redis to the wire: a client looping Last-Event-ID reconnects enumerates at most - # this many per round-trip, and the budget below bounds total throughput. - EVENTS_REPLAY_MAX_PER_REQUEST: int = 200 - EVENTS_REPLAY_MAX_AGE_HOURS: int = 48 - # Sliding-window cap on snapshot replays per user; exhausting it returns 429 with the - # cursor pinned so the client backs off until the window rolls over. - EVENTS_REPLAY_BUDGET_REQUESTS_PER_WINDOW: int = 30 - EVENTS_REPLAY_BUDGET_WINDOW_SECONDS: int = 60 - - # Retention for the message_events journal, enforced by the cleanup_message_events beat - # task. Replay only needs streams a client could still be tailing. - MESSAGE_EVENTS_RETENTION_DAYS: int = 14 - - # Remote Device feature. - REMOTE_DEVICE_SESSION_IDLE_SECONDS: int = 60 - REMOTE_DEVICE_REQUIRE_SIGNATURE: bool = False - REMOTE_DEVICE_PAIRING_TTL_SECONDS: int = 600 - # Redis broker tunables, routing invocations cross-process so a scheduled run reaches the - # web-held device session. The queue TTL must exceed the max drain deadline (605s) so a - # command for a briefly-offline device isn't evicted before its own drain gives up. - REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS: int = 900 - REMOTE_DEVICE_INVOCATION_TTL_SECONDS: int = 900 - REMOTE_DEVICE_OUTPUT_STREAM_MAXLEN: int = 10_000 - - # Scheduler (see scheduler.md). - SCHEDULE_DISPATCHER_INTERVAL: int = 30 - SCHEDULE_MIN_INTERVAL: int = 900 - SCHEDULE_MAX_PER_USER: int = 50 - SCHEDULE_RUN_TIMEOUT: int = 600 - SCHEDULE_MISFIRE_GRACE: int = 60 - SCHEDULE_AUTOPAUSE_FAILURES: int = 3 - SCHEDULE_ONCE_MAX_HORIZON: int = 31_536_000 - SCHEDULE_RUN_OUTPUT_RETENTION_DAYS: int = 90 - - # Code-execution sandbox. The app is a CLIENT of an always-on runner; defaults are safe so - # app import never fails when the sandbox is unconfigured. - SANDBOX_BACKEND: str = "jupyter" # "jupyter" (self-host) | "daytona" (Daytona Cloud) - # URL of the Jupyter Kernel Gateway runner (the docsgpt-sandbox service). - SANDBOX_GATEWAY_URL: str = "http://localhost:8888" - SANDBOX_GATEWAY_AUTH_TOKEN: Optional[str] = None # gateway auth token, if set - # Kernelspec per session. The env-scrubbing "docsgpt-python" spec keeps kernel code from - # reading the gateway token or operator secrets from os.environ; the stock "python3" spec - # inherits the gateway env verbatim and must not be used with untrusted code. - SANDBOX_KERNEL_NAME: str = "docsgpt-python" - SANDBOX_MAX_TTL: int = 1200 # hard cap (s) on agent-selectable keep-alive TTL - # Concurrent live sessions per process, backend-agnostic; at the cap an LRU-idle session is - # evicted. 0 or negative disables the cap. - SANDBOX_MAX_SESSIONS: int = 32 - SANDBOX_EXEC_TIMEOUT: int = 60 # default wall-clock cap (s) per exec call - SANDBOX_HTTP_TIMEOUT: int = 10 # fixed cap (s) for REST control calls (create/delete/alive/interrupt) - SANDBOX_MAX_OUTPUT_BYTES: int = 8 * 1024 * 1024 # cap on buffered stdout+stderr per exec - SANDBOX_MAX_FILE_BYTES: int = 10 * 1024 * 1024 # cap on get_file size routed through stdout - SANDBOX_MAX_INPUT_BYTES: int = 25 * 1024 * 1024 # cap on an input document staged into a sandbox session - # ``read_document`` parsing on a dedicated Celery ``parsing`` queue (backend parser). - DOCUMENT_PARSE_QUEUE: str = "parsing" # queue the parse_document task is routed to - DOCUMENT_PARSE_TIMEOUT: int = 120 # seconds the tool awaits the enqueued parse before degrading - # The base timeout is a FLOOR: the window grows with document size, because OCR cost scales - # with pages. Without this a large scan is silently dropped at the base window. - DOCUMENT_PARSE_TIMEOUT_PER_MB: int = 60 # extra seconds of parse window per MiB of input - DOCUMENT_PARSE_TIMEOUT_MAX: int = 900 # absolute ceiling on the size-scaled parse window - DOCUMENT_PARSE_MAX_BYTES: int = 0 # cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES) - DOCUMENT_MAX_DECOMPRESSED_BYTES: int = 300 * 1024 * 1024 - DOCUMENT_MAX_ARCHIVE_ENTRIES: int = 10000 - # Files per node passed natively to the LLM; past the cap they are extracted to text or - # dropped, to bound context and cost. Re-uses SANDBOX_MAX_INPUT_BYTES per file. - WORKFLOW_NODE_NATIVE_MAX_FILES: int = 5 - # Documents per node extracted via the parsing worker. Each issues a separate blocking - # parse; past the cap they are skipped with a truncation note. - WORKFLOW_NODE_EXTRACT_MAX_FILES: int = 5 - # Wall clock one node may spend on blocking parses, shared across all of them. Without it a - # node could serialize WORKFLOW_NODE_EXTRACT_MAX_FILES full windows on a web threadpool slot. - WORKFLOW_NODE_EXTRACT_BUDGET_SECONDS: int = 900 - # A run row is pre-created as ``running``; a disconnect or crash can strand it there. The - # beat reaper fails runs still ``running`` past this. Generous so a long run is never cut off. - WORKFLOW_RUN_STALE_SECONDS: int = 3600 - # Runner container caps, consumed by the docsgpt-sandbox compose service, not the app. - # These cgroup limits are part of the untrusted-code security boundary. - SANDBOX_MEMORY: str = "1g" # docker mem_limit for the runner container - SANDBOX_CPUS: str = "1.0" # docker cpu quota for the runner container - # Daytona Cloud backend (SANDBOX_BACKEND="daytona"). All knobs are optional so app import - # never fails when the backend is unused. - DAYTONA_API_KEY: Optional[str] = None # Daytona Cloud API key (secret) - DAYTONA_API_URL: Optional[str] = None # override Daytona API base URL, if self-targeting - DAYTONA_TARGET: Optional[str] = None # Daytona region/target, e.g. "us" - DAYTONA_SNAPSHOT: Optional[str] = None # image for new sandboxes; render libs via scripts/build_daytona_snapshot.py - DAYTONA_LANGUAGE: str = "python" # default runtime language for created sandboxes - DAYTONA_AUTO_STOP_INTERVAL: int = 15 # minutes idle before Daytona auto-stops a sandbox (0 disables) - DAYTONA_AUTO_DELETE_INTERVAL: int = 60 # minutes after stop before Daytona auto-deletes (-1 disables) - DAYTONA_MAX_SANDBOXES: int = 50 # cap on concurrent live Daytona sandboxes (cost-DoS guard) - # Per-user artifact quotas, enforced at persistence time. 0 or negative disables a quota. - ARTIFACT_MAX_BYTES: int = 50 * 1024 * 1024 # cap on a single stored artifact version's bytes - ARTIFACT_MAX_COUNT_PER_USER: int = 5000 # cap on artifacts a user may own - ARTIFACT_MAX_TOTAL_BYTES_PER_USER: int = 5 * 1024 * 1024 * 1024 # cap on a user's total stored bytes - - @field_validator("POSTGRES_URI", mode="before") - @classmethod - def _normalize_postgres_uri_validator(cls, v): - return normalize_postgres_uri(v) - - @field_validator("PGVECTOR_CONNECTION_STRING", mode="before") - @classmethod - def _normalize_pgvector_connection_string_validator(cls, v): - return normalize_pgvector_connection_string(v) - - @field_validator( - "API_KEY", - "OPENAI_API_KEY", - "ANTHROPIC_API_KEY", - "GOOGLE_API_KEY", - "GROQ_API_KEY", - "HUGGINGFACE_API_KEY", - "NOVITA_API_KEY", - "EMBEDDINGS_KEY", - "FALLBACK_LLM_API_KEY", - "QDRANT_API_KEY", - "ELEVENLABS_API_KEY", - "INTERNAL_KEY", - mode="before", - ) - @classmethod - def normalize_api_key(cls, v: Optional[str]) -> Optional[str]: - """ - Normalize API keys: convert 'None', 'none', empty strings, - and whitespace-only strings to actual None. - Handles Pydantic loading 'None' from .env as string "None". - """ - if v is None: - return None - if not isinstance(v, str): - return v - stripped = v.strip() - if stripped == "" or stripped.lower() == "none": - return None - return stripped - - -settings = Settings(_env_file=env_file(), _env_file_encoding="utf-8") diff --git a/docsgpt/core/settings/__init__.py b/docsgpt/core/settings/__init__.py new file mode 100644 index 00000000..2af13cea --- /dev/null +++ b/docsgpt/core/settings/__init__.py @@ -0,0 +1,77 @@ +"""Application settings. + +``settings`` is the process-wide instance, loaded from the environment and the +``.env`` file in the data home (see ``docsgpt.core.paths``). Every setting is a +flat attribute, ``settings.NAME``, matching the environment variable of the +same name. + +The definitions are split by domain into the modules of this package; each +module owns one ``SettingsGroup`` and ``Settings`` composes them all. Add a new +setting to the group it belongs to (or add a group and list it in +``SETTINGS_GROUPS``), with a ``description`` -- the settings reference in the +docs is generated from these definitions. +""" + +from __future__ import annotations + +from typing import Optional + +from docsgpt.core.paths import env_file, home_dir +from docsgpt.core.settings._shared import SettingsGroup, normalize_secret +from docsgpt.core.settings.agents import AgentSettings +from docsgpt.core.settings.auth import AuthSettings +from docsgpt.core.settings.connectors import ConnectorSettings +from docsgpt.core.settings.database import DatabaseSettings +from docsgpt.core.settings.embeddings import EmbeddingsSettings +from docsgpt.core.settings.events import EventsSettings +from docsgpt.core.settings.guardrails import GuardrailSettings +from docsgpt.core.settings.ingestion import IngestionSettings +from docsgpt.core.settings.llm import LLMSettings +from docsgpt.core.settings.ocr import OCRSettings +from docsgpt.core.settings.retrieval import RetrievalSettings +from docsgpt.core.settings.sandbox import SandboxSettings +from docsgpt.core.settings.scheduler import SchedulerSettings +from docsgpt.core.settings.server import ServerSettings +from docsgpt.core.settings.speech import SpeechSettings +from docsgpt.core.settings.storage import StorageSettings +from docsgpt.core.settings.vectorstores import VectorStoreSettings +from docsgpt.core.settings.workers import WorkerSettings + +#: Every settings group, in the order the generated reference lists them. +SETTINGS_GROUPS: tuple[tuple[str, type[SettingsGroup]], ...] = ( + ("Authentication", AuthSettings), + ("LLM providers", LLMSettings), + ("Embeddings", EmbeddingsSettings), + ("Retrieval", RetrievalSettings), + ("Vector stores", VectorStoreSettings), + ("User-data database", DatabaseSettings), + ("Workers", WorkerSettings), + ("Ingestion and parsing", IngestionSettings), + ("OCR", OCRSettings), + ("File storage", StorageSettings), + ("Connectors", ConnectorSettings), + ("Server", ServerSettings), + ("Events and devices", EventsSettings), + ("Agents", AgentSettings), + ("Guardrails", GuardrailSettings), + ("Scheduler", SchedulerSettings), + ("Sandbox", SandboxSettings), + ("Speech", SpeechSettings), +) + +# Runtime data home (DOCSGPT_HOME, the checkout, or cwd); see docsgpt.core.paths. +current_dir = str(home_dir()) + + +class Settings(*(group for _, group in SETTINGS_GROUPS)): + """All settings, composed from the per-domain groups in this package.""" + + @classmethod + def normalize_api_key(cls, v: Optional[str]) -> Optional[str]: + """Normalize a secret the way the per-field validators do; kept for callers that reuse it.""" + return normalize_secret(v) + + +settings = Settings(_env_file=env_file(), _env_file_encoding="utf-8") + +__all__ = ["SETTINGS_GROUPS", "Settings", "SettingsGroup", "current_dir", "settings"] diff --git a/docsgpt/core/settings/_shared.py b/docsgpt/core/settings/_shared.py new file mode 100644 index 00000000..4eb5c48d --- /dev/null +++ b/docsgpt/core/settings/_shared.py @@ -0,0 +1,37 @@ +"""Building blocks shared by the settings groups. + +Every group in this package is a :class:`SettingsGroup`: a ``BaseSettings`` +subclass that owns one domain's fields. ``docsgpt.core.settings.Settings`` +inherits from all of them, so the composed class keeps the flat +``settings.NAME`` attributes the rest of the codebase reads while each +domain's definitions live in their own module. +""" + +from __future__ import annotations + +from typing import Optional + +from pydantic_settings import BaseSettings, SettingsConfigDict + + +class SettingsGroup(BaseSettings): + """Base for one domain's settings; groups are composed into ``Settings``.""" + + model_config = SettingsConfigDict(extra="ignore") + + +def normalize_secret(value: Optional[str]) -> Optional[str]: + """Map the ways an unset secret reaches us from ``.env`` to ``None``. + + ``.env`` files carry ``KEY=None`` and ``KEY=`` for "not set", and pydantic + would otherwise keep those as the strings ``"None"`` and ``""``. Whitespace + around a real value is stripped. + """ + if value is None: + return None + if not isinstance(value, str): + return value + stripped = value.strip() + if stripped == "" or stripped.lower() == "none": + return None + return stripped diff --git a/docsgpt/core/settings/agents.py b/docsgpt/core/settings/agents.py new file mode 100644 index 00000000..fdee1cc3 --- /dev/null +++ b/docsgpt/core/settings/agents.py @@ -0,0 +1,91 @@ +"""Agent runtime: default tools, context management, workflows and artifacts.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class AgentSettings(SettingsGroup): + """What an agent may do per turn and how its context is kept within budget.""" + + AGENT_NAME: str = Field(default="classic", description="Default agent type for agentless chats.") + DEFAULT_MAX_HISTORY: int = Field(default=150, description="Default number of history messages kept.") + DEFAULT_AGENT_LIMITS: dict = Field( + default={"token_limit": 50000, "request_limit": 500}, + description="Per-agent default quotas: tokens and requests.", + ) + DEFAULT_CHAT_TOOLS: list = Field( + default=["memory", "read_webpage", "scheduler"], + description=( + "Config-free tools on by default in agentless chats. scheduler is dual-registered in " + "BUILTIN_AGENT_TOOLS so one synthetic id resolves via defaults or the agent picker. Add " + "code_executor and artifact_generator once a sandbox runner is configured; both execute through " + "it and would fail on every call without one." + ), + ) + ENABLE_TOOL_PREFETCH: bool = Field(default=True, description="Pre-fetch retrieval before the agent's first turn.") + TOOL_RESULT_MAX_TOKENS: int = Field( + default=20000, + description="Cap on one tool result entering the LLM context (0 disables); journal and DB keep it whole.", + ) + + # Conversation compression. + ENABLE_CONVERSATION_COMPRESSION: bool = Field( + default=True, description="Compress long conversations once they approach the context window." + ) + COMPRESSION_THRESHOLD_PERCENTAGE: float = Field( + default=0.8, description="Fraction of the context window at which compression triggers." + ) + COMPRESSION_MODEL_OVERRIDE: Optional[str] = Field( + default=None, description="Use a different model for compression; unset reuses the answer model." + ) + COMPRESSION_PROMPT_VERSION: str = Field(default="v1.0", description="Tracks compression prompt iterations.") + COMPRESSION_MAX_HISTORY_POINTS: int = Field( + default=3, description="Keep only the last N compression points to prevent DB bloat." + ) + COMPRESSION_RECENT_FIELD_MAX_TOKENS: int = Field( + default=8000, description="Per-field cap on the verbatim tail kept after a compression point (0 disables)." + ) + + # Workflows. + WORKFLOW_NODE_NATIVE_MAX_FILES: int = Field( + default=5, + description=( + "Files per node passed natively to the LLM; past the cap they are extracted to text or dropped, to " + "bound context and cost. Re-uses SANDBOX_MAX_INPUT_BYTES per file." + ), + ) + WORKFLOW_NODE_EXTRACT_MAX_FILES: int = Field( + default=5, + description=( + "Documents per node extracted via the parsing worker. Each issues a separate blocking parse; past " + "the cap they are skipped with a truncation note." + ), + ) + WORKFLOW_NODE_EXTRACT_BUDGET_SECONDS: int = Field( + default=900, + description=( + "Wall clock one node may spend on blocking parses, shared across all of them. Without it a node " + "could serialize WORKFLOW_NODE_EXTRACT_MAX_FILES full windows on a web threadpool slot." + ), + ) + WORKFLOW_RUN_STALE_SECONDS: int = Field( + default=3600, + description=( + "A run row is pre-created as running; a disconnect or crash can strand it there. The beat reaper " + "fails runs still running past this. Generous so a long run is never cut off." + ), + ) + + # Per-user artifact quotas, enforced at persistence time. 0 or negative disables a quota. + ARTIFACT_MAX_BYTES: int = Field( + default=50 * 1024 * 1024, description="Cap on a single stored artifact version's bytes (0 disables)." + ) + ARTIFACT_MAX_COUNT_PER_USER: int = Field(default=5000, description="Cap on artifacts a user may own (0 disables).") + ARTIFACT_MAX_TOTAL_BYTES_PER_USER: int = Field( + default=5 * 1024 * 1024 * 1024, description="Cap on a user's total stored artifact bytes (0 disables)." + ) diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py new file mode 100644 index 00000000..d0763934 --- /dev/null +++ b/docsgpt/core/settings/auth.py @@ -0,0 +1,87 @@ +"""Authentication, SSO and provisioning.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field, field_validator + +from docsgpt.core.settings._shared import SettingsGroup, normalize_secret + + +class AuthSettings(SettingsGroup): + """How users authenticate: none, a shared token, per-session JWTs, or OIDC SSO.""" + + AUTH_TYPE: Optional[str] = Field( + default=None, + description="Authentication mode: simple_jwt, session_jwt, oidc, or unset for no authentication.", + ) + JWT_SECRET_KEY: str = Field( + default="", + description=( + "Signing key for session tokens and other signed capabilities. Required on every replica in " + "production; local development may fall back to a key generated on disk." + ), + ) + ENCRYPTION_SECRET_KEY: str = Field( + default="default-docsgpt-encryption-key", + description="Key used to encrypt stored credentials such as tool and connector secrets.", + ) + INTERNAL_KEY: Optional[str] = Field( + default=None, description="Internal API key for worker-to-backend authentication." + ) + + # OIDC SSO (AUTH_TYPE=oidc): any OpenID Connect IdP with discovery (Authentik, Keycloak, ...). + OIDC_ISSUER: Optional[str] = Field( + default=None, + description="OIDC issuer URL with discovery, e.g. https://auth.example.com/application/o/docsgpt/.", + ) + OIDC_CLIENT_ID: Optional[str] = Field(default=None, description="OIDC client id.") + OIDC_CLIENT_SECRET: Optional[str] = Field( + default=None, description="OIDC client secret. Optional; PKCE is always used." + ) + OIDC_SCOPES: str = Field(default="openid profile email", description="Scopes requested from the IdP.") + OIDC_USER_ID_CLAIM: str = Field( + default="sub", description="ID-token claim mapped to the DocsGPT user id." + ) + OIDC_FRONTEND_URL: Optional[str] = Field( + default=None, description="Browser-facing app origin, e.g. http://localhost:5173." + ) + OIDC_REDIRECT_URI: Optional[str] = Field( + default=None, description="Override for the callback URL; default is /api/auth/oidc/callback." + ) + OIDC_SESSION_LIFETIME_SECONDS: int = Field( + default=28800, description="Lifetime of the minted session JWT in seconds (8h)." + ) + OIDC_PROVIDER_NAME: Optional[str] = Field( + default=None, description='Sign-in button label, e.g. "Acme SSO".' + ) + OIDC_ALLOWED_GROUPS: Optional[str] = Field( + default=None, description="Comma-separated group allowlist; unset admits any authenticated user." + ) + OIDC_GROUPS_CLAIM: str = Field( + default="groups", description="ID-token/userinfo claim carrying group membership." + ) + OIDC_ADMIN_GROUPS: Optional[str] = Field( + default=None, description="Comma-separated groups granted admin; unset means no OIDC admin mapping." + ) + + LOCAL_MODE_ADMIN: bool = Field( + default=False, + description=( + "Grant admin without a database role. Persisted admin grants live in user_roles (AUTH_TYPE=oidc " + "only); this is the only non-DB admin path, for AUTH_TYPE=None self-host. MUST stay False if " + "networked." + ), + ) + + # SCIM 2.0 provisioning (IdP-driven user create/deactivate at /scim/v2). + SCIM_ENABLED: bool = Field(default=False, description="Enable SCIM 2.0 provisioning at /scim/v2.") + SCIM_TOKEN: Optional[str] = Field( + default=None, description="Bearer token for IdP SCIM clients (required when SCIM is enabled)." + ) + + @field_validator("INTERNAL_KEY", mode="before") + @classmethod + def _normalize_auth_secrets(cls, v): + return normalize_secret(v) diff --git a/docsgpt/core/settings/connectors.py b/docsgpt/core/settings/connectors.py new file mode 100644 index 00000000..303c00b6 --- /dev/null +++ b/docsgpt/core/settings/connectors.py @@ -0,0 +1,50 @@ +"""OAuth credentials for external source connectors.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class ConnectorSettings(SettingsGroup): + """Client credentials and callback URLs for Google Drive, Microsoft, Confluence, GitHub and MCP.""" + + # Google Drive integration. + GOOGLE_CLIENT_ID: Optional[str] = Field(default=None, description="Google OAuth client id.") + GOOGLE_CLIENT_SECRET: Optional[str] = Field(default=None, description="Google OAuth client secret.") + CONNECTOR_REDIRECT_BASE_URI: Optional[str] = Field( + default="http://127.0.0.1:7091/api/connectors/callback", + description="OAuth callback URL; register it as-is in your provider's console (e.g. GCP).", + ) + CONNECTOR_ALLOWED_ORIGINS: Optional[str] = Field( + default=None, + description=( + "Comma-separated frontend origins allowed to receive connector OAuth results, e.g. " + "https://docsgpt.example.com. The callback origin and OIDC_FRONTEND_URL are always allowed; a " + "loopback callback also allows localhost:5173." + ), + ) + + # Microsoft Entra ID (Azure AD) integration. + MICROSOFT_CLIENT_ID: Optional[str] = Field(default=None, description="Azure AD application (client) id.") + MICROSOFT_CLIENT_SECRET: Optional[str] = Field(default=None, description="Azure AD application client secret.") + MICROSOFT_TENANT_ID: Optional[str] = Field( + default="common", description="Azure AD tenant id, or 'common' for multi-tenant." + ) + MICROSOFT_AUTHORITY: Optional[str] = Field( + default=None, description='Authority URL override, e.g. "https://login.microsoftonline.com/{tenant_id}".' + ) + + # Confluence Cloud integration. + CONFLUENCE_CLIENT_ID: Optional[str] = Field(default=None, description="Confluence Cloud OAuth client id.") + CONFLUENCE_CLIENT_SECRET: Optional[str] = Field(default=None, description="Confluence Cloud OAuth client secret.") + + # GitHub source. + GITHUB_ACCESS_TOKEN: Optional[str] = Field(default=None, description="GitHub PAT with read access to repositories.") + + MCP_OAUTH_REDIRECT_URI: Optional[str] = Field( + default=None, description="Public callback URL for MCP OAuth; unset derives it from CONNECTOR_REDIRECT_BASE_URI." + ) diff --git a/docsgpt/core/settings/database.py b/docsgpt/core/settings/database.py new file mode 100644 index 00000000..e88e03a9 --- /dev/null +++ b/docsgpt/core/settings/database.py @@ -0,0 +1,36 @@ +"""User-data Postgres and schema management at startup.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field, field_validator + +from docsgpt.core.db_uri import normalize_postgres_uri +from docsgpt.core.settings._shared import SettingsGroup + + +class DatabaseSettings(SettingsGroup): + """The Postgres database holding users, conversations and sources, and what startup may do to it.""" + + POSTGRES_URI: Optional[str] = Field(default=None, description="User-data Postgres connection URI.") + AUTO_MIGRATE: bool = Field( + default=True, + description="On startup, apply pending Alembic migrations. Disable if you manage schema out-of-band.", + ) + AUTO_CREATE_DB: bool = Field( + default=True, + description="On startup, create the target Postgres database if missing (needs CREATEDB privilege).", + ) + AUTO_VECTOR_SCHEMA: bool = Field( + default=True, + description=( + "On startup, create the pgvector/graph tables and verify the embedding dimension. No Alembic " + "migration covers the vector DB (it may be a separate cluster); set False to manage it yourself." + ), + ) + + @field_validator("POSTGRES_URI", mode="before") + @classmethod + def _normalize_postgres_uri(cls, v): + return normalize_postgres_uri(v) diff --git a/docsgpt/core/settings/embeddings.py b/docsgpt/core/settings/embeddings.py new file mode 100644 index 00000000..85de4d15 --- /dev/null +++ b/docsgpt/core/settings/embeddings.py @@ -0,0 +1,87 @@ +"""Embedding model selection and where it runs.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field, field_validator + +from docsgpt.core.paths import home_dir +from docsgpt.core.settings._shared import SettingsGroup, normalize_secret + + +class EmbeddingsSettings(SettingsGroup): + """The embedding model, remote or local, and the batching around it.""" + + EMBEDDINGS_NAME: str = Field( + default="huggingface_sentence-transformers/all-mpnet-base-v2", + description=( + "Embedding model. The legacy model is the default on purpose: an install that never pinned this " + "has vectors from it, and granite is the same width so a swap would fail silently. New installs " + "get granite from .env-template; existing ones switch by setting this and running " + "docsgpt.scripts.reembed." + ), + ) + EMBEDDINGS_BASE_URL: Optional[str] = Field( + default=None, description="Remote embeddings API URL (OpenAI-compatible)." + ) + EMBEDDINGS_KEY: Optional[str] = Field( + default=None, description="API key for embeddings (with OpenAI, the same value as API_KEY)." + ) + EMBEDDINGS_MAX_INPUT_TOKENS: Optional[int] = Field( + default=None, description="Truncate each remote embed input to N tokens (overflow is lost)." + ) + EMBEDDINGS_BATCH_SIZE: int = Field( + default=32, description="Chunks per store transaction and per remote embed request." + ) + EMBEDDINGS_MODEL_BATCH_SIZE: int = Field( + default=1, + description=( + "Documents per local ONNX forward pass. Each pass pads to its longest input, and that waste grows " + "with the square of chunk length: at 1250 tokens, 32 peaked at 6.6 GB, 1 at 2.9 GB." + ), + ) + EMBEDDINGS_THREADS: Optional[int] = Field( + default=None, + description=( + "Intra-op threads for the local ONNX runner; unset uses every core. It scales sub-linearly, so " + "several single-threaded workers beat one many-threaded process on the same cores." + ), + ) + EMBEDDINGS_CACHE_DIR: Optional[str] = Field( + default_factory=lambda: str(home_dir() / "models"), + description=( + "Where embedding models and their tokenizers are cached. Persistent by default: FastEmbed's own " + "default is the temp dir." + ), + ) + EMBEDDINGS_POOLING: Optional[str] = Field( + default=None, + description=( + 'Pooling strategy ("cls" or "mean"). Read from the model\'s own repository; set only for a ' + "repository that declares none, or to override what it declares." + ), + ) + EMBEDDINGS_NORMALIZE: Optional[bool] = Field( + default=None, + description=( + "L2-normalise embeddings. Read from the model's own repository; set only for a repository that " + "declares nothing, or to override what it declares." + ), + ) + EMBEDDINGS_DELEGATE_TO_WORKER: bool = Field( + default=True, + description=( + "Embed on the worker so the API holds no model (~890 MB), at one broker round trip per query. " + "Ignored when EMBEDDINGS_BASE_URL is set, which is the better answer for production." + ), + ) + EMBEDDINGS_QUEUE: str = Field(default="embeddings", description="Celery queue the embed task is routed to.") + EMBEDDINGS_DELEGATE_TIMEOUT: int = Field( + default=60, description="Seconds the API waits for the worker to return an embedding." + ) + + @field_validator("EMBEDDINGS_KEY", mode="before") + @classmethod + def _normalize_embeddings_secrets(cls, v): + return normalize_secret(v) diff --git a/docsgpt/core/settings/events.py b/docsgpt/core/settings/events.py new file mode 100644 index 00000000..0ab01d4c --- /dev/null +++ b/docsgpt/core/settings/events.py @@ -0,0 +1,86 @@ +"""Server-sent events, replay journal and remote-device sessions.""" + +from __future__ import annotations + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class EventsSettings(SettingsGroup): + """The internal push channel (notifications and durable replay) and the Redis pool behind it.""" + + ENABLE_SSE_PUSH: bool = Field( + default=True, + description=( + "Internal SSE push channel (notifications and durable replay journal). False makes /api/events emit " + '"push_disabled" and return; clients fall back to polling.' + ), + ) + EVENTS_STREAM_MAXLEN: int = Field( + default=1000, description="Per-user durable backlog cap in entries; ~24h of replay at typical rates." + ) + SSE_KEEPALIVE_SECONDS: int = Field(default=15, ge=1, description="Interval between SSE keepalive comments.") + SSE_MAX_CONCURRENT_PER_USER: int = Field( + default=8, + description=( + "Simultaneous SSE connections per user; each holds a pooled async Redis connection for its lifetime. " + "8 covers multi-tab use without one user starving the pool. 0 disables." + ), + ) + ASYNC_REDIS_MAX_CONNECTIONS: int = Field( + default=2000, + ge=1, + description=( + "Pool size of the async Redis client behind the event-loop routes, per process. Every open " + "notification tab, chat reconnect and device session holds one connection, so this caps concurrent " + "streams per worker (redis-py's own default is 100). Keep the total across workers below the Redis " + "server's maxclients (10000 by default)." + ), + ) + EVENTS_REPLAY_MAX_PER_REQUEST: int = Field( + default=200, + description=( + "Backlog entries XRANGE returns per /api/events snapshot. Bounds what one replay moves from Redis to " + "the wire: a client looping Last-Event-ID reconnects enumerates at most this many per round-trip." + ), + ) + EVENTS_REPLAY_MAX_AGE_HOURS: int = Field(default=48, description="Oldest backlog entry a replay will return.") + EVENTS_REPLAY_BUDGET_REQUESTS_PER_WINDOW: int = Field( + default=30, + description=( + "Sliding-window cap on snapshot replays per user; exhausting it returns 429 with the cursor pinned " + "so the client backs off until the window rolls over." + ), + ) + EVENTS_REPLAY_BUDGET_WINDOW_SECONDS: int = Field(default=60, description="Length of the replay budget window.") + MESSAGE_EVENTS_RETENTION_DAYS: int = Field( + default=14, + description=( + "Retention for the message_events journal, enforced by the cleanup_message_events beat task. Replay " + "only needs streams a client could still be tailing." + ), + ) + + # Remote Device feature. + REMOTE_DEVICE_SESSION_IDLE_SECONDS: int = Field( + default=60, description="Seconds without a heartbeat before a remote-device session is considered idle." + ) + REMOTE_DEVICE_REQUIRE_SIGNATURE: bool = Field( + default=False, description="Require signed commands from remote devices." + ) + REMOTE_DEVICE_PAIRING_TTL_SECONDS: int = Field(default=600, description="Lifetime of a pairing code.") + REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS: int = Field( + default=900, + description=( + "Redis TTL of the per-device command queue, routing invocations cross-process so a scheduled run " + "reaches the web-held device session. Must exceed the max drain deadline (605s) so a command for a " + "briefly-offline device isn't evicted before its own drain gives up." + ), + ) + REMOTE_DEVICE_INVOCATION_TTL_SECONDS: int = Field( + default=900, description="Redis TTL of a pending remote-device invocation." + ) + REMOTE_DEVICE_OUTPUT_STREAM_MAXLEN: int = Field( + default=10_000, description="Cap on buffered output entries per remote-device invocation stream." + ) diff --git a/docsgpt/core/settings/guardrails.py b/docsgpt/core/settings/guardrails.py new file mode 100644 index 00000000..81ac60f2 --- /dev/null +++ b/docsgpt/core/settings/guardrails.py @@ -0,0 +1,40 @@ +"""Agent guardrails.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class GuardrailSettings(SettingsGroup): + """Input/output checks every agent runs, and the floor no agent may weaken.""" + + GUARDRAILS_ENABLED: bool = Field(default=True, description="Master switch; False disables every stage.") + GUARDRAILS_CHECKS_ENABLED: list = Field( + default=[], description="Allowlist of GuardrailCreator.checks keys; empty means every registered check." + ) + GUARDRAILS_FLOOR: dict = Field( + default={}, + description=( + "A GuardrailsConfig fragment every agent inherits and cannot weaken; agents may add controls or " + 'make an action stricter, never looser. "enabled" is required; without it the floor parses but ' + 'applies to nothing. Example: {"enabled": true, "mode": "scan_all", "controls": [{"check": ' + '"secrets", "stage": "output", "action": "redact"}]}' + ), + ) + GUARDRAILS_JUDGE_MODEL: Optional[str] = Field( + default=None, description="Judge model for the topic/policy checks; unset reuses the request's model." + ) + GUARDRAILS_STORE_SCANNED_TEXT: bool = Field( + default=False, + description=( + "Persist scanned text alongside guardrail_events. Off by default: pre-redaction text is exactly the " + "material a PII control exists to keep out of storage." + ), + ) + GUARDRAILS_EVENTS_RETENTION_DAYS: int = Field( + default=30, ge=1, description="Days guardrail events are kept before the cleanup task removes them." + ) diff --git a/docsgpt/core/settings/ingestion.py b/docsgpt/core/settings/ingestion.py new file mode 100644 index 00000000..38da376d --- /dev/null +++ b/docsgpt/core/settings/ingestion.py @@ -0,0 +1,144 @@ +"""Uploads, document parsing and the size caps that keep one file from taking a worker down.""" + +from __future__ import annotations + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class IngestionSettings(SettingsGroup): + """Upload limits, the parser engine, and per-format byte caps for ingestion and attachments.""" + + UPLOAD_FOLDER: str = Field(default="inputs", description="Directory under the data home for uploaded sources.") + UPLOAD_MAX_REQUEST_BYTES: int = Field( + default=256 * 1024 * 1024, + gt=0, + description="Cap on an upload request body; applied by Flask before multipart parsing.", + ) + UPLOAD_MAX_FILE_BYTES: int = Field( + default=100 * 1024 * 1024, gt=0, description="Cap on a single uploaded file; also enforced while copying." + ) + PARSE_SPEC_MAX_BYTES: int = Field( + default=10 * 1024 * 1024, gt=0, description="Cap on an OpenAPI/tool spec file accepted for parsing." + ) + # ZIP limits apply cumulatively across nested archives in one extraction. + UPLOAD_MAX_ARCHIVE_BYTES: int = Field( + default=250 * 1024 * 1024, gt=0, description="Cap on total bytes extracted from one uploaded archive." + ) + UPLOAD_MAX_ARCHIVE_FILES: int = Field( + default=10_000, gt=0, description="Cap on files extracted from one uploaded archive." + ) + UPLOAD_MAX_ARCHIVE_RATIO: int = Field( + default=1000, gt=0, description="Maximum decompressed-to-compressed ratio before an archive is rejected." + ) + UPLOAD_MAX_ARCHIVE_DEPTH: int = Field( + default=3, ge=0, description="Maximum nesting depth of archives inside archives." + ) + PARSE_PDF_AS_IMAGE: bool = Field(default=False, description="Render PDF pages to images before parsing.") + PARSE_IMAGE_REMOTE: bool = Field(default=False, description="Send images to a remote parser.") + DOC_PARSER_ENGINE: str = Field( + default="anydoc", + description=( + 'Document parser for source ingestion, chat attachments and the read_document tool. "anydoc" ' + "(default): firecrawl-anydoc, a Rust converter with no ML models; milliseconds per file, ~100 MB " + 'peak RSS. "docling": the layout/table-model pipeline (optional install; needed for ' + "read_document's structured output and the docling OCR backend). Files anydoc cannot convert " + "(scanned PDFs, malformed input) fall back to docling when it is installed, otherwise to the native " + "OCR parsers (OCR on) or the legacy parsers. Rollback to the previous behaviour is this one variable." + ), + ) + DOCLING_PIPELINE_QUEUE_MAX_SIZE: int = Field( + default=2, + description=( + "Pages docling's threaded pipeline buffers in flight; the library default (100) drives worker RSS " + "to ~3 GB on a mid-size PDF." + ), + ) + DOCLING_COMPILE_TORCH_MODELS: bool = Field( + default=False, description="Let docling torch.compile its models (slower start, faster pages)." + ) + DOCLING_TABULAR_MAX_BYTES: int = Field( + default=2_000_000, description="Largest CSV/XLSX docling will parse, in bytes." + ) + DOCLING_MARKUP_MAX_BYTES: int = Field( + default=8_000_000, description="Largest HTML/XML docling will parse, in bytes." + ) + MARKUP_MAX_BYTES: int = Field( + default=8_000_000, + description=( + "HTML/XHTML larger than this (bytes) are head-truncated before the markdownify parser runs (the " + "anydoc engine's HTML path). The tree that path builds costs ~50x the input (30 MB of HTML measured " + "at 1.6 GB RSS) and the upload cap is 100 MB, so the gate is what keeps one upload from taking the " + "ingest worker down. 0 disables it." + ), + ) + PDF_TRUST_CHECK: bool = Field( + default=True, + description=( + "Trust-check anydoc's PDF output (docsgpt/parser/file/pdf_trust.py): flag composite (Type0) fonts " + "without a ToUnicode map, and CJK-declaring PDFs whose extracted text has almost no CJK, the two " + "classes where anydoc drops text silently. A flagged file re-parses on the docling fallback when " + "docling is installed; otherwise the anydoc output is kept and the document gets " + 'extra_info["parse_warnings"]. ~30 ms per scanned MB.' + ), + ) + ANYDOC_TABLEIZE: bool = Field( + default=False, + description=( + "Rewrite dot-leader / whitespace-aligned table runs in anydoc's PDF markdown into GFM tables " + "(docsgpt/parser/file/tableize.py). Off by default: it rewrites content on a heuristic (>=3 uniform " + "label+numbers lines) validated only on a small corpus so far." + ), + ) + ATTACHMENT_PDF_TEXT_FAST_PATH: bool = Field( + default=True, + description=( + "Read PDF attachments via their embedded text layer (pypdfium2) instead of docling, falling back to " + "docling when there is no text layer. Attachments go into a prompt, so docling's structural " + "markdown earns far less than the tens of seconds per file it costs; source ingestion is " + "unaffected because chunking and retrieval do depend on that structure." + ), + ) + ATTACHMENT_PDF_TEXT_MIN_MEDIAN_CHARS: int = Field( + default=32, + description=( + "Median chars per sampled page below which a PDF attachment is treated as a scan and handed to " + "docling. Measured on real uploads: scans at 0-17 chars/page, text-layer documents at 433-6834." + ), + ) + ATTACHMENT_TEXT_MAX_BYTES: int = Field(default=5_000_000, description="Cap on extracted attachment text.") + AGENT_IMAGE_MAX_BYTES: int = Field(default=5_000_000, description="Cap on an image passed to an agent.") + AGENT_IMAGE_MAX_PIXELS: int = Field( + default=16_777_216, description="Cap on the pixel count of an image passed to an agent." + ) + GITHUB_INGEST_MAX_FILE_BYTES: int = Field( + default=1048576, description="Skip GitHub repo blobs larger than this (0 = no cap)." + ) + GITHUB_INGEST_MAX_WORKERS: int = Field(default=8, description="Parallel file fetches per GitHub repo ingest.") + + # read_document parsing on a dedicated Celery queue (backend parser). + DOCUMENT_PARSE_QUEUE: str = Field(default="parsing", description="Celery queue the parse_document task is routed to.") + DOCUMENT_PARSE_TIMEOUT: int = Field( + default=120, description="Seconds the read_document tool awaits the enqueued parse before degrading." + ) + DOCUMENT_PARSE_TIMEOUT_PER_MB: int = Field( + default=60, + description=( + "Extra seconds of parse window per MiB of input. The base timeout is a FLOOR: the window grows with " + "document size because OCR cost scales with pages. Without this a large scan is silently dropped at " + "the base window." + ), + ) + DOCUMENT_PARSE_TIMEOUT_MAX: int = Field( + default=900, description="Absolute ceiling on the size-scaled parse window, in seconds." + ) + DOCUMENT_PARSE_MAX_BYTES: int = Field( + default=0, description="Cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES)." + ) + DOCUMENT_MAX_DECOMPRESSED_BYTES: int = Field( + default=300 * 1024 * 1024, description="Cap on bytes decompressed from an archive handed to read_document." + ) + DOCUMENT_MAX_ARCHIVE_ENTRIES: int = Field( + default=10000, description="Cap on entries in an archive handed to read_document." + ) diff --git a/docsgpt/core/settings/llm.py b/docsgpt/core/settings/llm.py new file mode 100644 index 00000000..232ad3be --- /dev/null +++ b/docsgpt/core/settings/llm.py @@ -0,0 +1,121 @@ +"""LLM providers, API keys and per-provider tunables.""" + +from __future__ import annotations + +import os +from typing import Optional + +from pydantic import Field, field_validator + +from docsgpt.core.paths import home_dir +from docsgpt.core.settings._shared import SettingsGroup, normalize_secret + + +class LLMSettings(SettingsGroup): + """Which model answers, how it is reached, and provider-specific behaviour.""" + + LLM_PROVIDER: str = Field(default="docsgpt", description="LLM provider key, e.g. openai, anthropic, docsgpt.") + LLM_NAME: Optional[str] = Field( + default=None, description="Model name for the provider; with openai, e.g. gpt-4 or gpt-3.5-turbo." + ) + API_KEY: Optional[str] = Field(default=None, description="LLM API key used by LLM_PROVIDER.") + + # Provider-specific API keys (for multi-model support). + OPENAI_API_KEY: Optional[str] = Field(default=None, description="OpenAI API key.") + ANTHROPIC_API_KEY: Optional[str] = Field(default=None, description="Anthropic API key.") + GOOGLE_API_KEY: Optional[str] = Field(default=None, description="Google AI API key.") + GROQ_API_KEY: Optional[str] = Field(default=None, description="Groq API key.") + HUGGINGFACE_API_KEY: Optional[str] = Field(default=None, description="Hugging Face API key.") + OPEN_ROUTER_API_KEY: Optional[str] = Field(default=None, description="OpenRouter API key.") + NOVITA_API_KEY: Optional[str] = Field(default=None, description="Novita API key.") + + OPENAI_API_BASE: Optional[str] = Field(default=None, description="Azure OpenAI API base URL.") + OPENAI_API_VERSION: Optional[str] = Field(default=None, description="Azure OpenAI API version.") + AZURE_DEPLOYMENT_NAME: Optional[str] = Field(default=None, description="Azure deployment name for answering.") + AZURE_EMBEDDINGS_DEPLOYMENT_NAME: Optional[str] = Field( + default=None, description="Azure deployment name for embeddings." + ) + OPENAI_BASE_URL: Optional[str] = Field( + default=None, description="Base URL for OpenAI-compatible model servers." + ) + LLM_PATH: str = Field( + default=os.path.join(str(home_dir()), "models/docsgpt-7b-f16.gguf"), + description="Path to the local GGUF model used by the llama.cpp provider.", + ) + + FALLBACK_LLM_PROVIDER: Optional[str] = Field(default=None, description="Provider for the fallback LLM.") + FALLBACK_LLM_NAME: Optional[str] = Field(default=None, description="Model name for the fallback LLM.") + FALLBACK_LLM_API_KEY: Optional[str] = Field(default=None, description="API key for the fallback LLM.") + TITLE_MODEL_ID: Optional[str] = Field( + default=None, description="Optional cheaper model for conversation titles; unset reuses the answer model." + ) + MODELS_CONFIG_DIR: Optional[str] = Field( + default=None, + description=( + "Directory of operator-supplied model YAMLs, loaded after the built-in catalog; later wins on " + "duplicate model id. See docsgpt/core/models/README.md." + ), + ) + DEFAULT_LLM_TOKEN_LIMIT: int = Field( + default=128000, description="Context window assumed when the model is not found in the registry." + ) + RESERVED_TOKENS: dict = Field( + default={"system_prompt": 500, "current_query": 500, "safety_buffer": 1000}, + description="Tokens held back from the context window for the system prompt, the query and a safety buffer.", + ) + CACHE_REDIS_URL: str = Field(default="redis://localhost:6379/2", description="Redis URL for the LLM cache.") + + # OpenAI Responses API. + OPENAI_RESPONSES_STORE: bool = Field( + default=False, + description=( + "True persists Responses API calls server-side so previous_response_id can chain turns. False keeps " + "them stateless, carrying reasoning across the tool loop as encrypted items." + ), + ) + OPENAI_RESPONSES_CHAIN_ACROSS_TURNS: bool = Field( + default=True, + description=( + "Cross-turn previous_response_id chaining (store mode only). The chained transcript lives on the " + "provider and is invisible to every local guard, so it is bounded: a turn starts from the local " + "history when the previous turn's reported prompt already reached the budget (default: the model's " + "context window) or when the conversation was compressed after that turn was produced." + ), + ) + OPENAI_RESPONSES_CHAIN_BUDGET_TOKENS: Optional[int] = Field( + default=None, description="Prompt-token budget for cross-turn chaining; unset uses the model's context window." + ) + OPENAI_RESPONSES_TRUNCATION_AUTO: bool = Field( + default=False, + description=( + 'Send truncation: "auto" so the provider drops the oldest input items instead of failing every ' + "request once a chain exceeds the model's window." + ), + ) + OPENAI_PROMPT_CACHE_KEY: bool = Field( + default=True, + description=( + "Route a user's Responses API calls to the same prompt-cache shard with an opaque per-user key." + ), + ) + OPENAI_PROMPT_CACHE_RETENTION: Optional[str] = Field( + default=None, description="Request extended prompt-cache retention where the provider offers it." + ) + OPENAI_REASONING_SUMMARY: str = Field( + default="auto", description="Reasoning summary mode requested from the Responses API." + ) + + @field_validator( + "API_KEY", + "OPENAI_API_KEY", + "ANTHROPIC_API_KEY", + "GOOGLE_API_KEY", + "GROQ_API_KEY", + "HUGGINGFACE_API_KEY", + "NOVITA_API_KEY", + "FALLBACK_LLM_API_KEY", + mode="before", + ) + @classmethod + def _normalize_llm_secrets(cls, v): + return normalize_secret(v) diff --git a/docsgpt/core/settings/ocr.py b/docsgpt/core/settings/ocr.py new file mode 100644 index 00000000..264b21c5 --- /dev/null +++ b/docsgpt/core/settings/ocr.py @@ -0,0 +1,91 @@ +"""OCR for scanned PDFs and images.""" + +from __future__ import annotations + +from pydantic import AliasChoices, Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class OCRSettings(SettingsGroup): + """Whether OCR runs, which stack performs it, and which engine it uses. + + OCR_ENABLED covers source ingestion, OCR_ATTACHMENTS_ENABLED chat attachments. Which stack performs + it is OCR_BACKEND; which engine, OCR_ENGINE. The DOCLING_OCR_* names are the pre-2026-09 spellings + and stay accepted as aliases. + """ + + OCR_ENABLED: bool = Field( + default=False, + validation_alias=AliasChoices("OCR_ENABLED", "DOCLING_OCR_ENABLED"), + description="OCR scanned PDFs and images during source ingestion.", + ) + OCR_ATTACHMENTS_ENABLED: bool = Field( + default=False, + validation_alias=AliasChoices("OCR_ATTACHMENTS_ENABLED", "DOCLING_OCR_ATTACHMENTS_ENABLED"), + description="OCR scanned PDFs and images attached to a chat.", + ) + OCR_BACKEND: str = Field( + default="auto", + description=( + "Which stack runs OCR when it is on. auto: docling when installed, otherwise native. docling: the " + "layout-model pipeline (hybrid region OCR, reading order, table structure); needs the optional " + "docling extra. native: pypdfium2/Pillow page rendering straight into tesseract or a DeepSeek-OCR " + "endpoint (docsgpt/parser/file/ocr_parser.py); no ML models in the worker, tables come out as text " + "lines under tesseract." + ), + ) + OCR_ENGINE: str = Field( + default="tesseract", + description=( + "OCR engine used when OCR is on. Benched 2026-08 on EN/ZH/table/degraded scans (docs/Guides/ocr has " + "the menu). tesseract (recommended): best classic-engine accuracy (perfect EN word recall, 0.000 " + "bilingual CER, 100% table cells), ~35 MB, CPU-only; needs the system binary and language packs, an " + "optional install like every OCR dependency (build with INSTALL_TESSERACT=true, or apt/brew install " + "tesseract-ocr for a local run); both backends. deepseek: DeepSeek-OCR against an Ollama/vLLM " + "endpoint (OCR_DEEPSEEK_*); best table/CJK quality, the worker stays light (no layout models) but " + "each page costs seconds on the model server; both backends. auto: docling's pick, ocrmac on macOS " + "(excellent), rapidocr on Linux (silently shreds some long text lines; avoid as a server default). " + "ocrmac | rapidocr: force one of those. auto/ocrmac/rapidocr exist only inside docling; the native " + "backend runs tesseract for them. An engine that is not installed degrades (docling: to auto) with " + "a warning instead of failing the parse." + ), + ) + OCR_LANGS: str = Field( + default="eng", + description=( + 'Tesseract language packs, "+"-separated (e.g. "eng+chi_sim+deu"). Other engines keep their own ' + "defaults; their language codes differ." + ), + ) + OCR_DEEPSEEK_URL: str = Field( + default="http://localhost:11434/v1/chat/completions", + description="Chat-completions URL of the DeepSeek-OCR endpoint (Ollama or vLLM).", + ) + OCR_DEEPSEEK_MODEL: str = Field(default="deepseek-ocr:3b", description="Model name at the DeepSeek-OCR endpoint.") + OCR_DEEPSEEK_TIMEOUT: float = Field( + default=300.0, + description=( + "Seconds allowed per page request to the DeepSeek endpoint, on both backends (native sends pages one " + "at a time; docling's VLM pipeline keeps its own concurrency). A 3B model on a laptop needs minutes; " + "a vLLM GPU deployment, seconds." + ), + ) + OCR_RENDER_DPI: int = Field( + default=200, + description=( + "Native backend only: resolution at which pages without a text layer are rendered before OCR. 200 " + "suits tesseract; clamped to 72-600." + ), + ) + OCR_MIN_CHARS_PER_PAGE: int = Field( + default=20, + validation_alias=AliasChoices("OCR_MIN_CHARS_PER_PAGE", "DOCLING_OCR_MIN_CHARS_PER_PAGE"), + description=( + "Chars-per-page floor below which an OCR'd PDF/image parse is treated as an OCR dropout rather than " + "as content (long-running docling workers were observed returning zero characters for every " + "scanned page after a long scanned PDF, with no error). docling retries once on a fresh full-page-OCR " + "converter; both backends then fail loudly instead of indexing an empty document. 0 disables the " + "guard." + ), + ) diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py new file mode 100644 index 00000000..39672d29 --- /dev/null +++ b/docsgpt/core/settings/retrieval.py @@ -0,0 +1,40 @@ +"""Retrieval strategy and GraphRAG.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class RetrievalSettings(SettingsGroup): + """Which vector store answers searches and how retrieval fans out across sources.""" + + VECTOR_STORE: str = Field( + default="faiss", + description="Vector store backend: faiss, elasticsearch, mongodb, qdrant, milvus or pgvector.", + ) + RETRIEVERS_ENABLED: list = Field( + default=["classic", "default"], + description=( + "Retriever keys an agent may use; must match RetrieverCreator.retrievers registry keys, NOT the " + "legacy classic_rag label which never matched the registry." + ), + ) + RETRIEVAL_MAX_PARALLEL_SOURCES: int = Field( + default=4, + description="Concurrent per-source searches in one retrieval; the query is embedded once and shared.", + ) + PER_SOURCE_RETRIEVAL_ENABLED: bool = Field( + default=True, + description="Kill-switch for per-source retrieval dispatch; False collapses to a single retriever.", + ) + GRAPHRAG_ENABLED: bool = Field(default=False, description="Gates graph-aware ingestion and retrieval.") + GRAPHRAG_EXTRACTION_MODEL: Optional[str] = Field( + default=None, description="Model for ingest-time graph extraction; unset reuses LLM_PROVIDER/LLM_NAME." + ) + GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field( + default=2000, description="Hard cap on chunks extracted per source (cost control)." + ) diff --git a/docsgpt/core/settings/sandbox.py b/docsgpt/core/settings/sandbox.py new file mode 100644 index 00000000..958734bc --- /dev/null +++ b/docsgpt/core/settings/sandbox.py @@ -0,0 +1,88 @@ +"""Code-execution sandbox: the Jupyter gateway runner or Daytona Cloud.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class SandboxSettings(SettingsGroup): + """The app is a CLIENT of an always-on runner; defaults are safe so app import never fails unconfigured.""" + + SANDBOX_BACKEND: str = Field( + default="jupyter", description="Sandbox backend: jupyter (self-host) or daytona (Daytona Cloud)." + ) + SANDBOX_GATEWAY_URL: str = Field( + default="http://localhost:8888", + description="URL of the Jupyter Kernel Gateway runner (the docsgpt-sandbox service).", + ) + SANDBOX_GATEWAY_AUTH_TOKEN: Optional[str] = Field(default=None, description="Gateway auth token, if set.") + SANDBOX_KERNEL_NAME: str = Field( + default="docsgpt-python", + description=( + "Kernelspec per session. The env-scrubbing docsgpt-python spec keeps kernel code from reading the " + "gateway token or operator secrets from os.environ; the stock python3 spec inherits the gateway env " + "verbatim and must not be used with untrusted code." + ), + ) + SANDBOX_MAX_TTL: int = Field(default=1200, description="Hard cap (s) on agent-selectable keep-alive TTL.") + SANDBOX_MAX_SESSIONS: int = Field( + default=32, + description=( + "Concurrent live sessions per process, backend-agnostic; at the cap an LRU-idle session is evicted. " + "0 or negative disables the cap." + ), + ) + SANDBOX_EXEC_TIMEOUT: int = Field(default=60, description="Default wall-clock cap (s) per exec call.") + SANDBOX_HTTP_TIMEOUT: int = Field( + default=10, description="Fixed cap (s) for REST control calls (create/delete/alive/interrupt)." + ) + SANDBOX_MAX_OUTPUT_BYTES: int = Field( + default=8 * 1024 * 1024, description="Cap on buffered stdout+stderr per exec." + ) + SANDBOX_MAX_FILE_BYTES: int = Field( + default=10 * 1024 * 1024, description="Cap on get_file size routed through stdout." + ) + SANDBOX_MAX_INPUT_BYTES: int = Field( + default=25 * 1024 * 1024, description="Cap on an input document staged into a sandbox session." + ) + # Runner container caps, consumed by the docsgpt-sandbox compose service, not the app. + SANDBOX_MEMORY: str = Field( + default="1g", + description=( + "Docker mem_limit for the runner container. Consumed by the docsgpt-sandbox compose service, not " + "the app; part of the untrusted-code security boundary." + ), + ) + SANDBOX_CPUS: str = Field( + default="1.0", + description=( + "Docker CPU quota for the runner container. Consumed by the docsgpt-sandbox compose service, not " + "the app; part of the untrusted-code security boundary." + ), + ) + + # Daytona Cloud backend (SANDBOX_BACKEND=daytona). All knobs are optional so app import never fails + # when the backend is unused. + DAYTONA_API_KEY: Optional[str] = Field(default=None, description="Daytona Cloud API key (secret).") + DAYTONA_API_URL: Optional[str] = Field( + default=None, description="Override for the Daytona API base URL, if self-targeting." + ) + DAYTONA_TARGET: Optional[str] = Field(default=None, description='Daytona region/target, e.g. "us".') + DAYTONA_SNAPSHOT: Optional[str] = Field( + default=None, + description="Image for new sandboxes; render libs via scripts/build_daytona_snapshot.py.", + ) + DAYTONA_LANGUAGE: str = Field(default="python", description="Default runtime language for created sandboxes.") + DAYTONA_AUTO_STOP_INTERVAL: int = Field( + default=15, description="Minutes idle before Daytona auto-stops a sandbox (0 disables)." + ) + DAYTONA_AUTO_DELETE_INTERVAL: int = Field( + default=60, description="Minutes after stop before Daytona auto-deletes a sandbox (-1 disables)." + ) + DAYTONA_MAX_SANDBOXES: int = Field( + default=50, description="Cap on concurrent live Daytona sandboxes (cost-DoS guard)." + ) diff --git a/docsgpt/core/settings/scheduler.py b/docsgpt/core/settings/scheduler.py new file mode 100644 index 00000000..00a9f455 --- /dev/null +++ b/docsgpt/core/settings/scheduler.py @@ -0,0 +1,28 @@ +"""Scheduled agent runs (see scheduler.md).""" + +from __future__ import annotations + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class SchedulerSettings(SettingsGroup): + """Cadence, quotas and timeouts of scheduled runs.""" + + SCHEDULE_DISPATCHER_INTERVAL: int = Field( + default=30, description="Seconds between dispatcher passes that enqueue due schedules." + ) + SCHEDULE_MIN_INTERVAL: int = Field(default=900, description="Smallest allowed recurrence interval in seconds.") + SCHEDULE_MAX_PER_USER: int = Field(default=50, description="Cap on schedules a user may own.") + SCHEDULE_RUN_TIMEOUT: int = Field(default=600, description="Wall-clock cap on one scheduled run, in seconds.") + SCHEDULE_MISFIRE_GRACE: int = Field( + default=60, description="Seconds past the due time within which a missed run still fires." + ) + SCHEDULE_AUTOPAUSE_FAILURES: int = Field( + default=3, description="Consecutive failures after which a schedule is paused automatically." + ) + SCHEDULE_ONCE_MAX_HORIZON: int = Field( + default=31_536_000, description="How far ahead a one-off run may be scheduled, in seconds (one year)." + ) + SCHEDULE_RUN_OUTPUT_RETENTION_DAYS: int = Field(default=90, description="Days scheduled-run output is kept.") diff --git a/docsgpt/core/settings/server.py b/docsgpt/core/settings/server.py new file mode 100644 index 00000000..84bb2510 --- /dev/null +++ b/docsgpt/core/settings/server.py @@ -0,0 +1,39 @@ +"""The API process itself.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class ServerSettings(SettingsGroup): + """Serving the UI, public URLs, and process-level knobs of the API server.""" + + SERVE_UI: bool = Field( + default=True, description="Serve the web UI shipped in the package (docsgpt/static) from the API process." + ) + FLASK_DEBUG_MODE: bool = Field(default=False, description="Run Flask in debug mode.") + VERSION_CHECK: bool = Field(default=True, description="Anonymous startup version check for security issues.") + PUBLIC_API_BASE_URL: Optional[str] = Field( + default=None, description="Public base URL for user-facing endpoint references in prompts." + ) + GRACEFUL_SHUTDOWN_TIMEOUT_SECONDS: int = Field( + default=30, + description=( + "Bounds uvicorn's shutdown drain (uvicorn_worker doesn't forward --graceful-timeout). Keep below the " + "gunicorn --timeout (180) watchdog. Used by BoundedDrainUvicornWorker." + ), + ) + WSGI_THREADPOOL_WORKERS: int = Field( + default=96, description="Threads serving the WSGI (Flask) part of the app under the ASGI server." + ) + V1_SESSION_TTL_SECONDS: int = Field( + default=24 * 60 * 60, + description=( + "Lets OpenAI-compatible clients identify a logical chat by session header, which chat-completions " + "itself has no field for; TTL of that session mapping." + ), + ) diff --git a/docsgpt/core/settings/speech.py b/docsgpt/core/settings/speech.py new file mode 100644 index 00000000..465c8fc5 --- /dev/null +++ b/docsgpt/core/settings/speech.py @@ -0,0 +1,31 @@ +"""Text-to-speech and speech-to-text.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field, field_validator + +from docsgpt.core.settings._shared import SettingsGroup, normalize_secret + + +class SpeechSettings(SettingsGroup): + """Voice providers and transcription options.""" + + TTS_PROVIDER: str = Field( + default="google_tts", description="Text-to-speech provider: google_tts, elevenlabs, or none to switch it off." + ) + ELEVENLABS_API_KEY: Optional[str] = Field(default=None, description="ElevenLabs API key.") + STT_PROVIDER: str = Field( + default="openai", description="Speech-to-text provider: openai, faster_whisper, or none to switch it off." + ) + OPENAI_STT_MODEL: str = Field(default="gpt-4o-mini-transcribe", description="OpenAI transcription model.") + STT_LANGUAGE: Optional[str] = Field(default=None, description="Language hint for transcription; unset auto-detects.") + STT_MAX_FILE_SIZE_MB: int = Field(default=50, description="Cap on an audio file accepted for transcription.") + STT_ENABLE_TIMESTAMPS: bool = Field(default=False, description="Return word/segment timestamps.") + STT_ENABLE_DIARIZATION: bool = Field(default=False, description="Label speakers in the transcript.") + + @field_validator("ELEVENLABS_API_KEY", mode="before") + @classmethod + def _normalize_speech_secrets(cls, v): + return normalize_secret(v) diff --git a/docsgpt/core/settings/storage.py b/docsgpt/core/settings/storage.py new file mode 100644 index 00000000..ac970a54 --- /dev/null +++ b/docsgpt/core/settings/storage.py @@ -0,0 +1,49 @@ +"""Where uploaded files and generated artifacts are stored.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class StorageSettings(SettingsGroup): + """Local disk or an S3-compatible bucket, and how download URLs are produced.""" + + STORAGE_TYPE: str = Field(default="local", description="File storage backend: local or s3.") + URL_STRATEGY: str = Field( + default="backend", + description="How download links are produced: backend (streamed through the API) or s3 (presigned URLs).", + ) + + # S3-compatible object storage (STORAGE_TYPE=s3): AWS S3, MinIO, R2, B2, Spaces, ... + # For non-AWS, set S3_ENDPOINT_URL and usually S3_PATH_STYLE=true. + S3_BUCKET_NAME: str = Field(default="docsgpt-test-bucket", description="Bucket name.") + S3_ENDPOINT_URL: Optional[str] = Field( + default=None, description="Custom endpoint for S3-compatible services (MinIO, R2, B2, Spaces); omit for AWS." + ) + S3_ACCESS_KEY_ID: Optional[str] = Field(default=None, description="Access key id.") + S3_SECRET_ACCESS_KEY: Optional[str] = Field(default=None, description="Secret access key.") + S3_REGION: Optional[str] = Field(default=None, description='AWS region; use "auto" for Cloudflare R2.') + S3_PATH_STYLE: bool = Field( + default=False, description="Path-style addressing (required by most non-AWS services)." + ) + + # Legacy AWS credentials from the retired SageMaker provider. + SAGEMAKER_REGION: Optional[str] = Field( + default=None, + description="Legacy AWS region from the retired SageMaker provider; deprecated fallback for S3_REGION.", + ) + SAGEMAKER_ACCESS_KEY: Optional[str] = Field( + default=None, + description="Legacy AWS access key from the retired SageMaker provider; deprecated fallback for S3_ACCESS_KEY_ID.", + ) + SAGEMAKER_SECRET_KEY: Optional[str] = Field( + default=None, + description=( + "Legacy AWS secret key from the retired SageMaker provider; deprecated fallback for " + "S3_SECRET_ACCESS_KEY." + ), + ) diff --git a/docsgpt/core/settings/vectorstores.py b/docsgpt/core/settings/vectorstores.py new file mode 100644 index 00000000..13c504df --- /dev/null +++ b/docsgpt/core/settings/vectorstores.py @@ -0,0 +1,89 @@ +"""Connection settings for each vector store backend.""" + +from __future__ import annotations + +from typing import Optional + +from pydantic import Field, field_validator + +from docsgpt.core.db_uri import normalize_pgvector_connection_string +from docsgpt.core.paths import home_dir +from docsgpt.core.settings._shared import SettingsGroup, normalize_secret + + +class VectorStoreSettings(SettingsGroup): + """Per-backend connection details; only the backend named by VECTOR_STORE is read.""" + + MONGO_URI: Optional[str] = Field( + default=None, + description=( + "Only consulted when VECTOR_STORE=mongodb or when running scripts/db/backfill.py; user data lives " + "in Postgres." + ), + ) + + # Elasticsearch. + ELASTIC_CLOUD_ID: Optional[str] = Field(default=None, description="Elastic Cloud id.") + ELASTIC_USERNAME: Optional[str] = Field(default=None, description="Elasticsearch username.") + ELASTIC_PASSWORD: Optional[str] = Field(default=None, description="Elasticsearch password.") + ELASTIC_URL: Optional[str] = Field(default=None, description="Elasticsearch URL.") + ELASTIC_INDEX: Optional[str] = Field(default="docsgpt", description="Elasticsearch index name.") + + # Qdrant. + QDRANT_COLLECTION_NAME: Optional[str] = Field(default="docsgpt", description="Qdrant collection name.") + QDRANT_LOCATION: Optional[str] = Field(default=None, description="Qdrant location (':memory:' or a URL).") + QDRANT_URL: Optional[str] = Field(default=None, description="Qdrant server URL.") + QDRANT_PORT: Optional[int] = Field(default=6333, description="Qdrant REST port.") + QDRANT_GRPC_PORT: int = Field(default=6334, description="Qdrant gRPC port.") + QDRANT_PREFER_GRPC: bool = Field(default=False, description="Use gRPC instead of REST where possible.") + QDRANT_HTTPS: Optional[bool] = Field(default=None, description="Use HTTPS for the Qdrant connection.") + QDRANT_API_KEY: Optional[str] = Field(default=None, description="Qdrant API key.") + QDRANT_PREFIX: Optional[str] = Field(default=None, description="URL prefix for a Qdrant behind a proxy.") + QDRANT_TIMEOUT: Optional[float] = Field(default=None, description="Qdrant request timeout in seconds.") + QDRANT_HOST: Optional[str] = Field(default=None, description="Qdrant host (alternative to QDRANT_URL).") + QDRANT_PATH: Optional[str] = Field(default=None, description="Path for an embedded on-disk Qdrant.") + QDRANT_DISTANCE_FUNC: str = Field(default="Cosine", description="Qdrant distance function.") + + # PGVector. + PGVECTOR_CONNECTION_STRING: Optional[str] = Field( + default=None, + description=( + "pgvector connection string. postgres://, postgresql:// and postgresql+psycopg:// are all accepted " + "and normalized internally for psycopg.connect(). Unset falls back to POSTGRES_URI." + ), + ) + PGVECTOR_POOL_MAX_SIZE: int = Field( + default=8, description="Per-process connection pool size; 0 uses one direct connection per store." + ) + PGVECTOR_IVFFLAT_PROBES: Optional[int] = Field( + default=None, + description="IVFFlat probes; unset derives sqrt(lists) from the index. Higher means better recall, more scan.", + ) + + # Milvus. + MILVUS_COLLECTION_NAME: Optional[str] = Field(default="docsgpt", description="Milvus collection name.") + MILVUS_URI: Optional[str] = Field( + default_factory=lambda: str(home_dir() / "milvus_local.db"), + description=( + "Milvus server URI. The default is a milvus-lite (embedded) database file under the data home, " + "like the other local stores." + ), + ) + MILVUS_TOKEN: Optional[str] = Field(default="", description="Milvus auth token.") + + # LanceDB. + LANCEDB_PATH: str = Field( + default_factory=lambda: str(home_dir() / "data" / "lancedb"), + description="LanceDB local data directory.", + ) + LANCEDB_TABLE_NAME: Optional[str] = Field(default="docsgpts", description="LanceDB table for stored vectors.") + + @field_validator("PGVECTOR_CONNECTION_STRING", mode="before") + @classmethod + def _normalize_pgvector_connection_string(cls, v): + return normalize_pgvector_connection_string(v) + + @field_validator("QDRANT_API_KEY", mode="before") + @classmethod + def _normalize_vectorstore_secrets(cls, v): + return normalize_secret(v) diff --git a/docsgpt/core/settings/workers.py b/docsgpt/core/settings/workers.py new file mode 100644 index 00000000..fbc7daae --- /dev/null +++ b/docsgpt/core/settings/workers.py @@ -0,0 +1,37 @@ +"""Celery broker, result backend and worker process limits.""" + +from __future__ import annotations + +from pydantic import Field + +from docsgpt.core.settings._shared import SettingsGroup + + +class WorkerSettings(SettingsGroup): + """How background tasks are queued and how worker processes are recycled.""" + + CELERY_BROKER_URL: str = Field(default="redis://localhost:6379/0", description="Celery broker URL.") + CELERY_RESULT_BACKEND: str = Field(default="redis://localhost:6379/1", description="Celery result backend URL.") + CELERY_WORKER_PREFETCH_MULTIPLIER: int = Field( + default=1, description="Tasks prefetched per worker process; 1 caps SIGKILL loss to one task." + ) + CELERY_VISIBILITY_TIMEOUT: int = Field( + default=3600, + description=( + "Broker visibility timeout in seconds. Must exceed the longest legitimate task runtime but stay " + "short enough that SIGKILLed tasks redeliver promptly." + ), + ) + CELERY_WORKER_MAX_MEMORY_PER_CHILD: int = Field( + default=4194304, + description=( + "Recycle a prefork child past this resident size in KB; backstops docling/torch heap growth. " + "Checked between tasks, so it does not bound the peak within one. 0 disables." + ), + ) + CELERY_WORKER_MAX_TASKS_PER_CHILD: int = Field( + default=0, description="Recycle a worker child after N tasks; 0 disables." + ) + API_URL: str = Field( + default="http://localhost:7091", description="Backend URL the Celery worker calls back into." + ) diff --git a/tests/test_remaining_coverage.py b/tests/test_remaining_coverage.py index b251a6b1..89fbded2 100644 --- a/tests/test_remaining_coverage.py +++ b/tests/test_remaining_coverage.py @@ -421,7 +421,7 @@ class TestBaseLLMAbstractRawGen: # --------------------------------------------------------------------------- -# docsgpt/core/settings.py (line 184 - clean_none_string) +# docsgpt/core/settings (normalize_api_key) # --------------------------------------------------------------------------- @pytest.mark.unit class TestSettingsNormalizeApiKey: From 5a5226ebe08262cc760a046dd8d2987ca88b4be6 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:04:11 +0100 Subject: [PATCH 052/130] docs(settings): generate the settings reference from the definitions The hand-maintained settings page documented 95 of 258 settings and .env-template 42, and both drifted as fields were added. The field descriptions now live on the model, so the reference is rendered from it: python -m docsgpt.core.settings.reference --write writes docs/content/Deploying/Settings-Reference.mdx, one section per settings group with each field's type, default, constraints, aliases and description. --check reports a stale page, and tests/core/test_settings.py fails when the checked-in page no longer matches the definitions, so a new setting cannot land undocumented. test_settings.py also pins the composition contract: every group field is a flat Settings attribute, no field is defined twice, every field has a description, and the secret-normalising validator of every group is applied (the case that a shared method name would silently drop). The App Configuration page points at the reference instead of at settings.py, and the reference is listed in the Deploying navigation. --- docs/content/Deploying/DocsGPT-Settings.mdx | 19 +- docs/content/Deploying/Settings-Reference.mdx | 1652 +++++++++++++++++ docs/content/Deploying/_meta.js | 4 + docsgpt/core/settings/reference.py | 169 ++ tests/core/test_settings.py | 109 ++ 5 files changed, 1943 insertions(+), 10 deletions(-) create mode 100644 docs/content/Deploying/Settings-Reference.mdx create mode 100644 docsgpt/core/settings/reference.py create mode 100644 tests/core/test_settings.py diff --git a/docs/content/Deploying/DocsGPT-Settings.mdx b/docs/content/Deploying/DocsGPT-Settings.mdx index e21e3be3..534413d6 100644 --- a/docs/content/Deploying/DocsGPT-Settings.mdx +++ b/docs/content/Deploying/DocsGPT-Settings.mdx @@ -27,13 +27,13 @@ API_KEY=YOUR_OPENAI_API_KEY LLM_NAME=gpt-4o ``` -### 2. Configuration via `settings.py` file (Advanced) +### 2. Configuration in code (Advanced) -For more advanced configurations or if you prefer to manage settings directly in code, you can modify the `settings.py` file. This file is located in the `docsgpt/core` directory of your DocsGPT project. +Every setting is defined in the `docsgpt/core/settings/` package, one module per domain (`auth.py`, `llm.py`, `embeddings.py`, ...). If you prefer to manage defaults directly in code, change them there; the `.env` file and the process environment still override whatever the code says. -While modifying `settings.py` offers more flexibility, it's generally recommended to use the `.env` file for basic settings and reserve `settings.py` for more complex adjustments or when you need to configure settings programmatically. +Using the `.env` file is recommended for day-to-day configuration. Reserve code changes for new settings or for defaults you want every deployment of your fork to share. -**Location of `settings.py`:** `docsgpt/core/settings.py` +The [Settings Reference](/Deploying/Settings-Reference) lists every setting with its type, default and description, generated from those definitions. ## Basic Settings Explained @@ -289,7 +289,7 @@ DocsGPT includes a JWT (JSON Web Token) based authentication feature for managin ### `AUTH_TYPE` Overview -The `AUTH_TYPE` setting in your `.env` file or `settings.py` determines the authentication method used by DocsGPT. This allows you to control how users authenticate with your DocsGPT instance. +The `AUTH_TYPE` setting in your `.env` file determines the authentication method used by DocsGPT. This allows you to control how users authenticate with your DocsGPT instance. | Value | Description | | ------------- | ------------------------------------------------------------------------------------------- | @@ -300,7 +300,7 @@ The `AUTH_TYPE` setting in your `.env` file or `settings.py` determines the auth #### How to Configure -Add the following to your `.env` file (or set in `settings.py`): +Add the following to your `.env` file: ```env # Shared signing key (required in production for every authentication mode) @@ -547,11 +547,10 @@ recovers. ## Exploring More Settings -These are just the basic settings to get you started. The `settings.py` file contains many more advanced options that you can explore to further customize DocsGPT, such as: +These are just the basic settings to get you started. DocsGPT has many more advanced options, such as: - Vector store configuration (`VECTOR_STORE`, Qdrant, Milvus, LanceDB settings) If you're looking for an easy way to set up a vector store with pgvector, try [Neon](https://get.neon.com/docsgpt). -- Retriever settings (`RETRIEVERS_ENABLED`) - Cache settings (`CACHE_REDIS_URL`) -- And many more! +- Sandbox, scheduler, guardrails and event-stream tuning -For a complete list of available settings and their descriptions, refer to the `settings.py` file in `docsgpt/core`. Remember to restart your Docker containers after making changes to your `.env` file or `settings.py` for the changes to take effect. +The [Settings Reference](/Deploying/Settings-Reference) lists every setting with its type, default and description. Remember to restart your Docker containers after making changes to your `.env` file for the changes to take effect. diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx new file mode 100644 index 00000000..6b935de9 --- /dev/null +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -0,0 +1,1652 @@ +--- +title: Settings Reference +description: Every DocsGPT setting, grouped by domain, with its type, default and purpose. +--- + +{/* GENERATED FILE. Do not edit by hand: run `python -m docsgpt.core.settings.reference --write`. */} + +# Settings Reference + +Every setting DocsGPT reads, generated from `docsgpt/core/settings/`. Each one is +an environment variable of the same name, set in `.env` or the process +environment; see [App Configuration](/Deploying/DocsGPT-Settings) for how the +file is found and for worked examples. `` below is the data home +described there. + + +## Authentication + +How users authenticate: none, a shared token, per-session JWTs, or OIDC SSO. + +### `AUTH_TYPE` + +Type `str`, default unset. + +Authentication mode: simple_jwt, session_jwt, oidc, or unset for no authentication. + +### `JWT_SECRET_KEY` + +Type `str`, default `""`. + +Signing key for session tokens and other signed capabilities. Required on every replica in production; local development may fall back to a key generated on disk. + +### `ENCRYPTION_SECRET_KEY` + +Type `str`, default `default-docsgpt-encryption-key`. + +Key used to encrypt stored credentials such as tool and connector secrets. + +### `INTERNAL_KEY` + +Type `str`, default unset. + +Internal API key for worker-to-backend authentication. + +### `OIDC_ISSUER` + +Type `str`, default unset. + +OIDC issuer URL with discovery, e.g. https://auth.example.com/application/o/docsgpt/. + +### `OIDC_CLIENT_ID` + +Type `str`, default unset. + +OIDC client id. + +### `OIDC_CLIENT_SECRET` + +Type `str`, default unset. + +OIDC client secret. Optional; PKCE is always used. + +### `OIDC_SCOPES` + +Type `str`, default `openid profile email`. + +Scopes requested from the IdP. + +### `OIDC_USER_ID_CLAIM` + +Type `str`, default `sub`. + +ID-token claim mapped to the DocsGPT user id. + +### `OIDC_FRONTEND_URL` + +Type `str`, default unset. + +Browser-facing app origin, e.g. http://localhost:5173. + +### `OIDC_REDIRECT_URI` + +Type `str`, default unset. + +Override for the callback URL; default is <request host>/api/auth/oidc/callback. + +### `OIDC_SESSION_LIFETIME_SECONDS` + +Type `int`, default `28800`. + +Lifetime of the minted session JWT in seconds (8h). + +### `OIDC_PROVIDER_NAME` + +Type `str`, default unset. + +Sign-in button label, e.g. "Acme SSO". + +### `OIDC_ALLOWED_GROUPS` + +Type `str`, default unset. + +Comma-separated group allowlist; unset admits any authenticated user. + +### `OIDC_GROUPS_CLAIM` + +Type `str`, default `groups`. + +ID-token/userinfo claim carrying group membership. + +### `OIDC_ADMIN_GROUPS` + +Type `str`, default unset. + +Comma-separated groups granted admin; unset means no OIDC admin mapping. + +### `LOCAL_MODE_ADMIN` + +Type `bool`, default `false`. + +Grant admin without a database role. Persisted admin grants live in user_roles (AUTH_TYPE=oidc only); this is the only non-DB admin path, for AUTH_TYPE=None self-host. MUST stay False if networked. + +### `SCIM_ENABLED` + +Type `bool`, default `false`. + +Enable SCIM 2.0 provisioning at /scim/v2. + +### `SCIM_TOKEN` + +Type `str`, default unset. + +Bearer token for IdP SCIM clients (required when SCIM is enabled). + + +## LLM providers + +Which model answers, how it is reached, and provider-specific behaviour. + +### `LLM_PROVIDER` + +Type `str`, default `docsgpt`. + +LLM provider key, e.g. openai, anthropic, docsgpt. + +### `LLM_NAME` + +Type `str`, default unset. + +Model name for the provider; with openai, e.g. gpt-4 or gpt-3.5-turbo. + +### `API_KEY` + +Type `str`, default unset. + +LLM API key used by LLM_PROVIDER. + +### `OPENAI_API_KEY` + +Type `str`, default unset. + +OpenAI API key. + +### `ANTHROPIC_API_KEY` + +Type `str`, default unset. + +Anthropic API key. + +### `GOOGLE_API_KEY` + +Type `str`, default unset. + +Google AI API key. + +### `GROQ_API_KEY` + +Type `str`, default unset. + +Groq API key. + +### `HUGGINGFACE_API_KEY` + +Type `str`, default unset. + +Hugging Face API key. + +### `OPEN_ROUTER_API_KEY` + +Type `str`, default unset. + +OpenRouter API key. + +### `NOVITA_API_KEY` + +Type `str`, default unset. + +Novita API key. + +### `OPENAI_API_BASE` + +Type `str`, default unset. + +Azure OpenAI API base URL. + +### `OPENAI_API_VERSION` + +Type `str`, default unset. + +Azure OpenAI API version. + +### `AZURE_DEPLOYMENT_NAME` + +Type `str`, default unset. + +Azure deployment name for answering. + +### `AZURE_EMBEDDINGS_DEPLOYMENT_NAME` + +Type `str`, default unset. + +Azure deployment name for embeddings. + +### `OPENAI_BASE_URL` + +Type `str`, default unset. + +Base URL for OpenAI-compatible model servers. + +### `LLM_PATH` + +Type `str`, default `/models/docsgpt-7b-f16.gguf`. + +Path to the local GGUF model used by the llama.cpp provider. + +### `FALLBACK_LLM_PROVIDER` + +Type `str`, default unset. + +Provider for the fallback LLM. + +### `FALLBACK_LLM_NAME` + +Type `str`, default unset. + +Model name for the fallback LLM. + +### `FALLBACK_LLM_API_KEY` + +Type `str`, default unset. + +API key for the fallback LLM. + +### `TITLE_MODEL_ID` + +Type `str`, default unset. + +Optional cheaper model for conversation titles; unset reuses the answer model. + +### `MODELS_CONFIG_DIR` + +Type `str`, default unset. + +Directory of operator-supplied model YAMLs, loaded after the built-in catalog; later wins on duplicate model id. See docsgpt/core/models/README.md. + +### `DEFAULT_LLM_TOKEN_LIMIT` + +Type `int`, default `128000`. + +Context window assumed when the model is not found in the registry. + +### `RESERVED_TOKENS` + +Type `dict`, default `{"system_prompt": 500, "current_query": 500, "safety_buffer": 1000}`. + +Tokens held back from the context window for the system prompt, the query and a safety buffer. + +### `CACHE_REDIS_URL` + +Type `str`, default `redis://localhost:6379/2`. + +Redis URL for the LLM cache. + +### `OPENAI_RESPONSES_STORE` + +Type `bool`, default `false`. + +True persists Responses API calls server-side so previous_response_id can chain turns. False keeps them stateless, carrying reasoning across the tool loop as encrypted items. + +### `OPENAI_RESPONSES_CHAIN_ACROSS_TURNS` + +Type `bool`, default `true`. + +Cross-turn previous_response_id chaining (store mode only). The chained transcript lives on the provider and is invisible to every local guard, so it is bounded: a turn starts from the local history when the previous turn's reported prompt already reached the budget (default: the model's context window) or when the conversation was compressed after that turn was produced. + +### `OPENAI_RESPONSES_CHAIN_BUDGET_TOKENS` + +Type `int`, default unset. + +Prompt-token budget for cross-turn chaining; unset uses the model's context window. + +### `OPENAI_RESPONSES_TRUNCATION_AUTO` + +Type `bool`, default `false`. + +Send truncation: "auto" so the provider drops the oldest input items instead of failing every request once a chain exceeds the model's window. + +### `OPENAI_PROMPT_CACHE_KEY` + +Type `bool`, default `true`. + +Route a user's Responses API calls to the same prompt-cache shard with an opaque per-user key. + +### `OPENAI_PROMPT_CACHE_RETENTION` + +Type `str`, default unset. + +Request extended prompt-cache retention where the provider offers it. + +### `OPENAI_REASONING_SUMMARY` + +Type `str`, default `auto`. + +Reasoning summary mode requested from the Responses API. + + +## Embeddings + +The embedding model, remote or local, and the batching around it. + +### `EMBEDDINGS_NAME` + +Type `str`, default `huggingface_sentence-transformers/all-mpnet-base-v2`. + +Embedding model. The legacy model is the default on purpose: an install that never pinned this has vectors from it, and granite is the same width so a swap would fail silently. New installs get granite from .env-template; existing ones switch by setting this and running docsgpt.scripts.reembed. + +### `EMBEDDINGS_BASE_URL` + +Type `str`, default unset. + +Remote embeddings API URL (OpenAI-compatible). + +### `EMBEDDINGS_KEY` + +Type `str`, default unset. + +API key for embeddings (with OpenAI, the same value as API_KEY). + +### `EMBEDDINGS_MAX_INPUT_TOKENS` + +Type `int`, default unset. + +Truncate each remote embed input to N tokens (overflow is lost). + +### `EMBEDDINGS_BATCH_SIZE` + +Type `int`, default `32`. + +Chunks per store transaction and per remote embed request. + +### `EMBEDDINGS_MODEL_BATCH_SIZE` + +Type `int`, default `1`. + +Documents per local ONNX forward pass. Each pass pads to its longest input, and that waste grows with the square of chunk length: at 1250 tokens, 32 peaked at 6.6 GB, 1 at 2.9 GB. + +### `EMBEDDINGS_THREADS` + +Type `int`, default unset. + +Intra-op threads for the local ONNX runner; unset uses every core. It scales sub-linearly, so several single-threaded workers beat one many-threaded process on the same cores. + +### `EMBEDDINGS_CACHE_DIR` + +Type `str`, default `/models`. + +Where embedding models and their tokenizers are cached. Persistent by default: FastEmbed's own default is the temp dir. + +### `EMBEDDINGS_POOLING` + +Type `str`, default unset. + +Pooling strategy ("cls" or "mean"). Read from the model's own repository; set only for a repository that declares none, or to override what it declares. + +### `EMBEDDINGS_NORMALIZE` + +Type `bool`, default unset. + +L2-normalise embeddings. Read from the model's own repository; set only for a repository that declares nothing, or to override what it declares. + +### `EMBEDDINGS_DELEGATE_TO_WORKER` + +Type `bool`, default `true`. + +Embed on the worker so the API holds no model (~890 MB), at one broker round trip per query. Ignored when EMBEDDINGS_BASE_URL is set, which is the better answer for production. + +### `EMBEDDINGS_QUEUE` + +Type `str`, default `embeddings`. + +Celery queue the embed task is routed to. + +### `EMBEDDINGS_DELEGATE_TIMEOUT` + +Type `int`, default `60`. + +Seconds the API waits for the worker to return an embedding. + + +## Retrieval + +Which vector store answers searches and how retrieval fans out across sources. + +### `VECTOR_STORE` + +Type `str`, default `faiss`. + +Vector store backend: faiss, elasticsearch, mongodb, qdrant, milvus or pgvector. + +### `RETRIEVERS_ENABLED` + +Type `list`, default `["classic", "default"]`. + +Retriever keys an agent may use; must match RetrieverCreator.retrievers registry keys, NOT the legacy classic_rag label which never matched the registry. + +### `RETRIEVAL_MAX_PARALLEL_SOURCES` + +Type `int`, default `4`. + +Concurrent per-source searches in one retrieval; the query is embedded once and shared. + +### `PER_SOURCE_RETRIEVAL_ENABLED` + +Type `bool`, default `true`. + +Kill-switch for per-source retrieval dispatch; False collapses to a single retriever. + +### `GRAPHRAG_ENABLED` + +Type `bool`, default `false`. + +Gates graph-aware ingestion and retrieval. + +### `GRAPHRAG_EXTRACTION_MODEL` + +Type `str`, default unset. + +Model for ingest-time graph extraction; unset reuses LLM_PROVIDER/LLM_NAME. + +### `GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION` + +Type `int`, default `2000`. + +Hard cap on chunks extracted per source (cost control). + + +## Vector stores + +Per-backend connection details; only the backend named by VECTOR_STORE is read. + +### `MONGO_URI` + +Type `str`, default unset. + +Only consulted when VECTOR_STORE=mongodb or when running scripts/db/backfill.py; user data lives in Postgres. + +### `ELASTIC_CLOUD_ID` + +Type `str`, default unset. + +Elastic Cloud id. + +### `ELASTIC_USERNAME` + +Type `str`, default unset. + +Elasticsearch username. + +### `ELASTIC_PASSWORD` + +Type `str`, default unset. + +Elasticsearch password. + +### `ELASTIC_URL` + +Type `str`, default unset. + +Elasticsearch URL. + +### `ELASTIC_INDEX` + +Type `str`, default `docsgpt`. + +Elasticsearch index name. + +### `QDRANT_COLLECTION_NAME` + +Type `str`, default `docsgpt`. + +Qdrant collection name. + +### `QDRANT_LOCATION` + +Type `str`, default unset. + +Qdrant location (':memory:' or a URL). + +### `QDRANT_URL` + +Type `str`, default unset. + +Qdrant server URL. + +### `QDRANT_PORT` + +Type `int`, default `6333`. + +Qdrant REST port. + +### `QDRANT_GRPC_PORT` + +Type `int`, default `6334`. + +Qdrant gRPC port. + +### `QDRANT_PREFER_GRPC` + +Type `bool`, default `false`. + +Use gRPC instead of REST where possible. + +### `QDRANT_HTTPS` + +Type `bool`, default unset. + +Use HTTPS for the Qdrant connection. + +### `QDRANT_API_KEY` + +Type `str`, default unset. + +Qdrant API key. + +### `QDRANT_PREFIX` + +Type `str`, default unset. + +URL prefix for a Qdrant behind a proxy. + +### `QDRANT_TIMEOUT` + +Type `float`, default unset. + +Qdrant request timeout in seconds. + +### `QDRANT_HOST` + +Type `str`, default unset. + +Qdrant host (alternative to QDRANT_URL). + +### `QDRANT_PATH` + +Type `str`, default unset. + +Path for an embedded on-disk Qdrant. + +### `QDRANT_DISTANCE_FUNC` + +Type `str`, default `Cosine`. + +Qdrant distance function. + +### `PGVECTOR_CONNECTION_STRING` + +Type `str`, default unset. + +pgvector connection string. postgres://, postgresql:// and postgresql+psycopg:// are all accepted and normalized internally for psycopg.connect(). Unset falls back to POSTGRES_URI. + +### `PGVECTOR_POOL_MAX_SIZE` + +Type `int`, default `8`. + +Per-process connection pool size; 0 uses one direct connection per store. + +### `PGVECTOR_IVFFLAT_PROBES` + +Type `int`, default unset. + +IVFFlat probes; unset derives sqrt(lists) from the index. Higher means better recall, more scan. + +### `MILVUS_COLLECTION_NAME` + +Type `str`, default `docsgpt`. + +Milvus collection name. + +### `MILVUS_URI` + +Type `str`, default `/milvus_local.db`. + +Milvus server URI. The default is a milvus-lite (embedded) database file under the data home, like the other local stores. + +### `MILVUS_TOKEN` + +Type `str`, default `""`. + +Milvus auth token. + +### `LANCEDB_PATH` + +Type `str`, default `/data/lancedb`. + +LanceDB local data directory. + +### `LANCEDB_TABLE_NAME` + +Type `str`, default `docsgpts`. + +LanceDB table for stored vectors. + + +## User-data database + +The Postgres database holding users, conversations and sources, and what startup may do to it. + +### `POSTGRES_URI` + +Type `str`, default unset. + +User-data Postgres connection URI. + +### `AUTO_MIGRATE` + +Type `bool`, default `true`. + +On startup, apply pending Alembic migrations. Disable if you manage schema out-of-band. + +### `AUTO_CREATE_DB` + +Type `bool`, default `true`. + +On startup, create the target Postgres database if missing (needs CREATEDB privilege). + +### `AUTO_VECTOR_SCHEMA` + +Type `bool`, default `true`. + +On startup, create the pgvector/graph tables and verify the embedding dimension. No Alembic migration covers the vector DB (it may be a separate cluster); set False to manage it yourself. + + +## Workers + +How background tasks are queued and how worker processes are recycled. + +### `CELERY_BROKER_URL` + +Type `str`, default `redis://localhost:6379/0`. + +Celery broker URL. + +### `CELERY_RESULT_BACKEND` + +Type `str`, default `redis://localhost:6379/1`. + +Celery result backend URL. + +### `CELERY_WORKER_PREFETCH_MULTIPLIER` + +Type `int`, default `1`. + +Tasks prefetched per worker process; 1 caps SIGKILL loss to one task. + +### `CELERY_VISIBILITY_TIMEOUT` + +Type `int`, default `3600`. + +Broker visibility timeout in seconds. Must exceed the longest legitimate task runtime but stay short enough that SIGKILLed tasks redeliver promptly. + +### `CELERY_WORKER_MAX_MEMORY_PER_CHILD` + +Type `int`, default `4194304`. + +Recycle a prefork child past this resident size in KB; backstops docling/torch heap growth. Checked between tasks, so it does not bound the peak within one. 0 disables. + +### `CELERY_WORKER_MAX_TASKS_PER_CHILD` + +Type `int`, default `0`. + +Recycle a worker child after N tasks; 0 disables. + +### `API_URL` + +Type `str`, default `http://localhost:7091`. + +Backend URL the Celery worker calls back into. + + +## Ingestion and parsing + +Upload limits, the parser engine, and per-format byte caps for ingestion and attachments. + +### `UPLOAD_FOLDER` + +Type `str`, default `inputs`. + +Directory under the data home for uploaded sources. + +### `UPLOAD_MAX_REQUEST_BYTES` + +Type `int`, default `268435456`, must be > 0. + +Cap on an upload request body; applied by Flask before multipart parsing. + +### `UPLOAD_MAX_FILE_BYTES` + +Type `int`, default `104857600`, must be > 0. + +Cap on a single uploaded file; also enforced while copying. + +### `PARSE_SPEC_MAX_BYTES` + +Type `int`, default `10485760`, must be > 0. + +Cap on an OpenAPI/tool spec file accepted for parsing. + +### `UPLOAD_MAX_ARCHIVE_BYTES` + +Type `int`, default `262144000`, must be > 0. + +Cap on total bytes extracted from one uploaded archive. + +### `UPLOAD_MAX_ARCHIVE_FILES` + +Type `int`, default `10000`, must be > 0. + +Cap on files extracted from one uploaded archive. + +### `UPLOAD_MAX_ARCHIVE_RATIO` + +Type `int`, default `1000`, must be > 0. + +Maximum decompressed-to-compressed ratio before an archive is rejected. + +### `UPLOAD_MAX_ARCHIVE_DEPTH` + +Type `int`, default `3`, must be >= 0. + +Maximum nesting depth of archives inside archives. + +### `PARSE_PDF_AS_IMAGE` + +Type `bool`, default `false`. + +Render PDF pages to images before parsing. + +### `PARSE_IMAGE_REMOTE` + +Type `bool`, default `false`. + +Send images to a remote parser. + +### `DOC_PARSER_ENGINE` + +Type `str`, default `anydoc`. + +Document parser for source ingestion, chat attachments and the read_document tool. "anydoc" (default): firecrawl-anydoc, a Rust converter with no ML models; milliseconds per file, ~100 MB peak RSS. "docling": the layout/table-model pipeline (optional install; needed for read_document's structured output and the docling OCR backend). Files anydoc cannot convert (scanned PDFs, malformed input) fall back to docling when it is installed, otherwise to the native OCR parsers (OCR on) or the legacy parsers. Rollback to the previous behaviour is this one variable. + +### `DOCLING_PIPELINE_QUEUE_MAX_SIZE` + +Type `int`, default `2`. + +Pages docling's threaded pipeline buffers in flight; the library default (100) drives worker RSS to ~3 GB on a mid-size PDF. + +### `DOCLING_COMPILE_TORCH_MODELS` + +Type `bool`, default `false`. + +Let docling torch.compile its models (slower start, faster pages). + +### `DOCLING_TABULAR_MAX_BYTES` + +Type `int`, default `2000000`. + +Largest CSV/XLSX docling will parse, in bytes. + +### `DOCLING_MARKUP_MAX_BYTES` + +Type `int`, default `8000000`. + +Largest HTML/XML docling will parse, in bytes. + +### `MARKUP_MAX_BYTES` + +Type `int`, default `8000000`. + +HTML/XHTML larger than this (bytes) are head-truncated before the markdownify parser runs (the anydoc engine's HTML path). The tree that path builds costs ~50x the input (30 MB of HTML measured at 1.6 GB RSS) and the upload cap is 100 MB, so the gate is what keeps one upload from taking the ingest worker down. 0 disables it. + +### `PDF_TRUST_CHECK` + +Type `bool`, default `true`. + +Trust-check anydoc's PDF output (docsgpt/parser/file/pdf_trust.py): flag composite (Type0) fonts without a ToUnicode map, and CJK-declaring PDFs whose extracted text has almost no CJK, the two classes where anydoc drops text silently. A flagged file re-parses on the docling fallback when docling is installed; otherwise the anydoc output is kept and the document gets extra_info["parse_warnings"]. ~30 ms per scanned MB. + +### `ANYDOC_TABLEIZE` + +Type `bool`, default `false`. + +Rewrite dot-leader / whitespace-aligned table runs in anydoc's PDF markdown into GFM tables (docsgpt/parser/file/tableize.py). Off by default: it rewrites content on a heuristic (>=3 uniform label+numbers lines) validated only on a small corpus so far. + +### `ATTACHMENT_PDF_TEXT_FAST_PATH` + +Type `bool`, default `true`. + +Read PDF attachments via their embedded text layer (pypdfium2) instead of docling, falling back to docling when there is no text layer. Attachments go into a prompt, so docling's structural markdown earns far less than the tens of seconds per file it costs; source ingestion is unaffected because chunking and retrieval do depend on that structure. + +### `ATTACHMENT_PDF_TEXT_MIN_MEDIAN_CHARS` + +Type `int`, default `32`. + +Median chars per sampled page below which a PDF attachment is treated as a scan and handed to docling. Measured on real uploads: scans at 0-17 chars/page, text-layer documents at 433-6834. + +### `ATTACHMENT_TEXT_MAX_BYTES` + +Type `int`, default `5000000`. + +Cap on extracted attachment text. + +### `AGENT_IMAGE_MAX_BYTES` + +Type `int`, default `5000000`. + +Cap on an image passed to an agent. + +### `AGENT_IMAGE_MAX_PIXELS` + +Type `int`, default `16777216`. + +Cap on the pixel count of an image passed to an agent. + +### `GITHUB_INGEST_MAX_FILE_BYTES` + +Type `int`, default `1048576`. + +Skip GitHub repo blobs larger than this (0 = no cap). + +### `GITHUB_INGEST_MAX_WORKERS` + +Type `int`, default `8`. + +Parallel file fetches per GitHub repo ingest. + +### `DOCUMENT_PARSE_QUEUE` + +Type `str`, default `parsing`. + +Celery queue the parse_document task is routed to. + +### `DOCUMENT_PARSE_TIMEOUT` + +Type `int`, default `120`. + +Seconds the read_document tool awaits the enqueued parse before degrading. + +### `DOCUMENT_PARSE_TIMEOUT_PER_MB` + +Type `int`, default `60`. + +Extra seconds of parse window per MiB of input. The base timeout is a FLOOR: the window grows with document size because OCR cost scales with pages. Without this a large scan is silently dropped at the base window. + +### `DOCUMENT_PARSE_TIMEOUT_MAX` + +Type `int`, default `900`. + +Absolute ceiling on the size-scaled parse window, in seconds. + +### `DOCUMENT_PARSE_MAX_BYTES` + +Type `int`, default `0`. + +Cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES). + +### `DOCUMENT_MAX_DECOMPRESSED_BYTES` + +Type `int`, default `314572800`. + +Cap on bytes decompressed from an archive handed to read_document. + +### `DOCUMENT_MAX_ARCHIVE_ENTRIES` + +Type `int`, default `10000`. + +Cap on entries in an archive handed to read_document. + + +## OCR + +Whether OCR runs, which stack performs it, and which engine it uses. + +### `OCR_ENABLED` + +Type `bool`, default `false`, also read from `DOCLING_OCR_ENABLED`. + +OCR scanned PDFs and images during source ingestion. + +### `OCR_ATTACHMENTS_ENABLED` + +Type `bool`, default `false`, also read from `DOCLING_OCR_ATTACHMENTS_ENABLED`. + +OCR scanned PDFs and images attached to a chat. + +### `OCR_BACKEND` + +Type `str`, default `auto`. + +Which stack runs OCR when it is on. auto: docling when installed, otherwise native. docling: the layout-model pipeline (hybrid region OCR, reading order, table structure); needs the optional docling extra. native: pypdfium2/Pillow page rendering straight into tesseract or a DeepSeek-OCR endpoint (docsgpt/parser/file/ocr_parser.py); no ML models in the worker, tables come out as text lines under tesseract. + +### `OCR_ENGINE` + +Type `str`, default `tesseract`. + +OCR engine used when OCR is on. Benched 2026-08 on EN/ZH/table/degraded scans (docs/Guides/ocr has the menu). tesseract (recommended): best classic-engine accuracy (perfect EN word recall, 0.000 bilingual CER, 100% table cells), ~35 MB, CPU-only; needs the system binary and language packs, an optional install like every OCR dependency (build with INSTALL_TESSERACT=true, or apt/brew install tesseract-ocr for a local run); both backends. deepseek: DeepSeek-OCR against an Ollama/vLLM endpoint (OCR_DEEPSEEK_*); best table/CJK quality, the worker stays light (no layout models) but each page costs seconds on the model server; both backends. auto: docling's pick, ocrmac on macOS (excellent), rapidocr on Linux (silently shreds some long text lines; avoid as a server default). ocrmac | rapidocr: force one of those. auto/ocrmac/rapidocr exist only inside docling; the native backend runs tesseract for them. An engine that is not installed degrades (docling: to auto) with a warning instead of failing the parse. + +### `OCR_LANGS` + +Type `str`, default `eng`. + +Tesseract language packs, "+"-separated (e.g. "eng+chi_sim+deu"). Other engines keep their own defaults; their language codes differ. + +### `OCR_DEEPSEEK_URL` + +Type `str`, default `http://localhost:11434/v1/chat/completions`. + +Chat-completions URL of the DeepSeek-OCR endpoint (Ollama or vLLM). + +### `OCR_DEEPSEEK_MODEL` + +Type `str`, default `deepseek-ocr:3b`. + +Model name at the DeepSeek-OCR endpoint. + +### `OCR_DEEPSEEK_TIMEOUT` + +Type `float`, default `300.0`. + +Seconds allowed per page request to the DeepSeek endpoint, on both backends (native sends pages one at a time; docling's VLM pipeline keeps its own concurrency). A 3B model on a laptop needs minutes; a vLLM GPU deployment, seconds. + +### `OCR_RENDER_DPI` + +Type `int`, default `200`. + +Native backend only: resolution at which pages without a text layer are rendered before OCR. 200 suits tesseract; clamped to 72-600. + +### `OCR_MIN_CHARS_PER_PAGE` + +Type `int`, default `20`, also read from `DOCLING_OCR_MIN_CHARS_PER_PAGE`. + +Chars-per-page floor below which an OCR'd PDF/image parse is treated as an OCR dropout rather than as content (long-running docling workers were observed returning zero characters for every scanned page after a long scanned PDF, with no error). docling retries once on a fresh full-page-OCR converter; both backends then fail loudly instead of indexing an empty document. 0 disables the guard. + + +## File storage + +Local disk or an S3-compatible bucket, and how download URLs are produced. + +### `STORAGE_TYPE` + +Type `str`, default `local`. + +File storage backend: local or s3. + +### `URL_STRATEGY` + +Type `str`, default `backend`. + +How download links are produced: backend (streamed through the API) or s3 (presigned URLs). + +### `S3_BUCKET_NAME` + +Type `str`, default `docsgpt-test-bucket`. + +Bucket name. + +### `S3_ENDPOINT_URL` + +Type `str`, default unset. + +Custom endpoint for S3-compatible services (MinIO, R2, B2, Spaces); omit for AWS. + +### `S3_ACCESS_KEY_ID` + +Type `str`, default unset. + +Access key id. + +### `S3_SECRET_ACCESS_KEY` + +Type `str`, default unset. + +Secret access key. + +### `S3_REGION` + +Type `str`, default unset. + +AWS region; use "auto" for Cloudflare R2. + +### `S3_PATH_STYLE` + +Type `bool`, default `false`. + +Path-style addressing (required by most non-AWS services). + +### `SAGEMAKER_REGION` + +Type `str`, default unset. + +Legacy AWS region from the retired SageMaker provider; deprecated fallback for S3_REGION. + +### `SAGEMAKER_ACCESS_KEY` + +Type `str`, default unset. + +Legacy AWS access key from the retired SageMaker provider; deprecated fallback for S3_ACCESS_KEY_ID. + +### `SAGEMAKER_SECRET_KEY` + +Type `str`, default unset. + +Legacy AWS secret key from the retired SageMaker provider; deprecated fallback for S3_SECRET_ACCESS_KEY. + + +## Connectors + +Client credentials and callback URLs for Google Drive, Microsoft, Confluence, GitHub and MCP. + +### `GOOGLE_CLIENT_ID` + +Type `str`, default unset. + +Google OAuth client id. + +### `GOOGLE_CLIENT_SECRET` + +Type `str`, default unset. + +Google OAuth client secret. + +### `CONNECTOR_REDIRECT_BASE_URI` + +Type `str`, default `http://127.0.0.1:7091/api/connectors/callback`. + +OAuth callback URL; register it as-is in your provider's console (e.g. GCP). + +### `CONNECTOR_ALLOWED_ORIGINS` + +Type `str`, default unset. + +Comma-separated frontend origins allowed to receive connector OAuth results, e.g. https://docsgpt.example.com. The callback origin and OIDC_FRONTEND_URL are always allowed; a loopback callback also allows localhost:5173. + +### `MICROSOFT_CLIENT_ID` + +Type `str`, default unset. + +Azure AD application (client) id. + +### `MICROSOFT_CLIENT_SECRET` + +Type `str`, default unset. + +Azure AD application client secret. + +### `MICROSOFT_TENANT_ID` + +Type `str`, default `common`. + +Azure AD tenant id, or 'common' for multi-tenant. + +### `MICROSOFT_AUTHORITY` + +Type `str`, default unset. + +Authority URL override, e.g. "https://login.microsoftonline.com/\{tenant_id\}". + +### `CONFLUENCE_CLIENT_ID` + +Type `str`, default unset. + +Confluence Cloud OAuth client id. + +### `CONFLUENCE_CLIENT_SECRET` + +Type `str`, default unset. + +Confluence Cloud OAuth client secret. + +### `GITHUB_ACCESS_TOKEN` + +Type `str`, default unset. + +GitHub PAT with read access to repositories. + +### `MCP_OAUTH_REDIRECT_URI` + +Type `str`, default unset. + +Public callback URL for MCP OAuth; unset derives it from CONNECTOR_REDIRECT_BASE_URI. + + +## Server + +Serving the UI, public URLs, and process-level knobs of the API server. + +### `SERVE_UI` + +Type `bool`, default `true`. + +Serve the web UI shipped in the package (docsgpt/static) from the API process. + +### `FLASK_DEBUG_MODE` + +Type `bool`, default `false`. + +Run Flask in debug mode. + +### `VERSION_CHECK` + +Type `bool`, default `true`. + +Anonymous startup version check for security issues. + +### `PUBLIC_API_BASE_URL` + +Type `str`, default unset. + +Public base URL for user-facing endpoint references in prompts. + +### `GRACEFUL_SHUTDOWN_TIMEOUT_SECONDS` + +Type `int`, default `30`. + +Bounds uvicorn's shutdown drain (uvicorn_worker doesn't forward --graceful-timeout). Keep below the gunicorn --timeout (180) watchdog. Used by BoundedDrainUvicornWorker. + +### `WSGI_THREADPOOL_WORKERS` + +Type `int`, default `96`. + +Threads serving the WSGI (Flask) part of the app under the ASGI server. + +### `V1_SESSION_TTL_SECONDS` + +Type `int`, default `86400`. + +Lets OpenAI-compatible clients identify a logical chat by session header, which chat-completions itself has no field for; TTL of that session mapping. + + +## Events and devices + +The internal push channel (notifications and durable replay) and the Redis pool behind it. + +### `ENABLE_SSE_PUSH` + +Type `bool`, default `true`. + +Internal SSE push channel (notifications and durable replay journal). False makes /api/events emit "push_disabled" and return; clients fall back to polling. + +### `EVENTS_STREAM_MAXLEN` + +Type `int`, default `1000`. + +Per-user durable backlog cap in entries; ~24h of replay at typical rates. + +### `SSE_KEEPALIVE_SECONDS` + +Type `int`, default `15`, must be >= 1. + +Interval between SSE keepalive comments. + +### `SSE_MAX_CONCURRENT_PER_USER` + +Type `int`, default `8`. + +Simultaneous SSE connections per user; each holds a pooled async Redis connection for its lifetime. 8 covers multi-tab use without one user starving the pool. 0 disables. + +### `ASYNC_REDIS_MAX_CONNECTIONS` + +Type `int`, default `2000`, must be >= 1. + +Pool size of the async Redis client behind the event-loop routes, per process. Every open notification tab, chat reconnect and device session holds one connection, so this caps concurrent streams per worker (redis-py's own default is 100). Keep the total across workers below the Redis server's maxclients (10000 by default). + +### `EVENTS_REPLAY_MAX_PER_REQUEST` + +Type `int`, default `200`. + +Backlog entries XRANGE returns per /api/events snapshot. Bounds what one replay moves from Redis to the wire: a client looping Last-Event-ID reconnects enumerates at most this many per round-trip. + +### `EVENTS_REPLAY_MAX_AGE_HOURS` + +Type `int`, default `48`. + +Oldest backlog entry a replay will return. + +### `EVENTS_REPLAY_BUDGET_REQUESTS_PER_WINDOW` + +Type `int`, default `30`. + +Sliding-window cap on snapshot replays per user; exhausting it returns 429 with the cursor pinned so the client backs off until the window rolls over. + +### `EVENTS_REPLAY_BUDGET_WINDOW_SECONDS` + +Type `int`, default `60`. + +Length of the replay budget window. + +### `MESSAGE_EVENTS_RETENTION_DAYS` + +Type `int`, default `14`. + +Retention for the message_events journal, enforced by the cleanup_message_events beat task. Replay only needs streams a client could still be tailing. + +### `REMOTE_DEVICE_SESSION_IDLE_SECONDS` + +Type `int`, default `60`. + +Seconds without a heartbeat before a remote-device session is considered idle. + +### `REMOTE_DEVICE_REQUIRE_SIGNATURE` + +Type `bool`, default `false`. + +Require signed commands from remote devices. + +### `REMOTE_DEVICE_PAIRING_TTL_SECONDS` + +Type `int`, default `600`. + +Lifetime of a pairing code. + +### `REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS` + +Type `int`, default `900`. + +Redis TTL of the per-device command queue, routing invocations cross-process so a scheduled run reaches the web-held device session. Must exceed the max drain deadline (605s) so a command for a briefly-offline device isn't evicted before its own drain gives up. + +### `REMOTE_DEVICE_INVOCATION_TTL_SECONDS` + +Type `int`, default `900`. + +Redis TTL of a pending remote-device invocation. + +### `REMOTE_DEVICE_OUTPUT_STREAM_MAXLEN` + +Type `int`, default `10000`. + +Cap on buffered output entries per remote-device invocation stream. + + +## Agents + +What an agent may do per turn and how its context is kept within budget. + +### `AGENT_NAME` + +Type `str`, default `classic`. + +Default agent type for agentless chats. + +### `DEFAULT_MAX_HISTORY` + +Type `int`, default `150`. + +Default number of history messages kept. + +### `DEFAULT_AGENT_LIMITS` + +Type `dict`, default `{"token_limit": 50000, "request_limit": 500}`. + +Per-agent default quotas: tokens and requests. + +### `DEFAULT_CHAT_TOOLS` + +Type `list`, default `["memory", "read_webpage", "scheduler"]`. + +Config-free tools on by default in agentless chats. scheduler is dual-registered in BUILTIN_AGENT_TOOLS so one synthetic id resolves via defaults or the agent picker. Add code_executor and artifact_generator once a sandbox runner is configured; both execute through it and would fail on every call without one. + +### `ENABLE_TOOL_PREFETCH` + +Type `bool`, default `true`. + +Pre-fetch retrieval before the agent's first turn. + +### `TOOL_RESULT_MAX_TOKENS` + +Type `int`, default `20000`. + +Cap on one tool result entering the LLM context (0 disables); journal and DB keep it whole. + +### `ENABLE_CONVERSATION_COMPRESSION` + +Type `bool`, default `true`. + +Compress long conversations once they approach the context window. + +### `COMPRESSION_THRESHOLD_PERCENTAGE` + +Type `float`, default `0.8`. + +Fraction of the context window at which compression triggers. + +### `COMPRESSION_MODEL_OVERRIDE` + +Type `str`, default unset. + +Use a different model for compression; unset reuses the answer model. + +### `COMPRESSION_PROMPT_VERSION` + +Type `str`, default `v1.0`. + +Tracks compression prompt iterations. + +### `COMPRESSION_MAX_HISTORY_POINTS` + +Type `int`, default `3`. + +Keep only the last N compression points to prevent DB bloat. + +### `COMPRESSION_RECENT_FIELD_MAX_TOKENS` + +Type `int`, default `8000`. + +Per-field cap on the verbatim tail kept after a compression point (0 disables). + +### `WORKFLOW_NODE_NATIVE_MAX_FILES` + +Type `int`, default `5`. + +Files per node passed natively to the LLM; past the cap they are extracted to text or dropped, to bound context and cost. Re-uses SANDBOX_MAX_INPUT_BYTES per file. + +### `WORKFLOW_NODE_EXTRACT_MAX_FILES` + +Type `int`, default `5`. + +Documents per node extracted via the parsing worker. Each issues a separate blocking parse; past the cap they are skipped with a truncation note. + +### `WORKFLOW_NODE_EXTRACT_BUDGET_SECONDS` + +Type `int`, default `900`. + +Wall clock one node may spend on blocking parses, shared across all of them. Without it a node could serialize WORKFLOW_NODE_EXTRACT_MAX_FILES full windows on a web threadpool slot. + +### `WORKFLOW_RUN_STALE_SECONDS` + +Type `int`, default `3600`. + +A run row is pre-created as running; a disconnect or crash can strand it there. The beat reaper fails runs still running past this. Generous so a long run is never cut off. + +### `ARTIFACT_MAX_BYTES` + +Type `int`, default `52428800`. + +Cap on a single stored artifact version's bytes (0 disables). + +### `ARTIFACT_MAX_COUNT_PER_USER` + +Type `int`, default `5000`. + +Cap on artifacts a user may own (0 disables). + +### `ARTIFACT_MAX_TOTAL_BYTES_PER_USER` + +Type `int`, default `5368709120`. + +Cap on a user's total stored artifact bytes (0 disables). + + +## Guardrails + +Input/output checks every agent runs, and the floor no agent may weaken. + +### `GUARDRAILS_ENABLED` + +Type `bool`, default `true`. + +Master switch; False disables every stage. + +### `GUARDRAILS_CHECKS_ENABLED` + +Type `list`, default `[]`. + +Allowlist of GuardrailCreator.checks keys; empty means every registered check. + +### `GUARDRAILS_FLOOR` + +Type `dict`, default `{}`. + +A GuardrailsConfig fragment every agent inherits and cannot weaken; agents may add controls or make an action stricter, never looser. "enabled" is required; without it the floor parses but applies to nothing. Example: \{"enabled": true, "mode": "scan_all", "controls": [\{"check": "secrets", "stage": "output", "action": "redact"\}]\} + +### `GUARDRAILS_JUDGE_MODEL` + +Type `str`, default unset. + +Judge model for the topic/policy checks; unset reuses the request's model. + +### `GUARDRAILS_STORE_SCANNED_TEXT` + +Type `bool`, default `false`. + +Persist scanned text alongside guardrail_events. Off by default: pre-redaction text is exactly the material a PII control exists to keep out of storage. + +### `GUARDRAILS_EVENTS_RETENTION_DAYS` + +Type `int`, default `30`, must be >= 1. + +Days guardrail events are kept before the cleanup task removes them. + + +## Scheduler + +Cadence, quotas and timeouts of scheduled runs. + +### `SCHEDULE_DISPATCHER_INTERVAL` + +Type `int`, default `30`. + +Seconds between dispatcher passes that enqueue due schedules. + +### `SCHEDULE_MIN_INTERVAL` + +Type `int`, default `900`. + +Smallest allowed recurrence interval in seconds. + +### `SCHEDULE_MAX_PER_USER` + +Type `int`, default `50`. + +Cap on schedules a user may own. + +### `SCHEDULE_RUN_TIMEOUT` + +Type `int`, default `600`. + +Wall-clock cap on one scheduled run, in seconds. + +### `SCHEDULE_MISFIRE_GRACE` + +Type `int`, default `60`. + +Seconds past the due time within which a missed run still fires. + +### `SCHEDULE_AUTOPAUSE_FAILURES` + +Type `int`, default `3`. + +Consecutive failures after which a schedule is paused automatically. + +### `SCHEDULE_ONCE_MAX_HORIZON` + +Type `int`, default `31536000`. + +How far ahead a one-off run may be scheduled, in seconds (one year). + +### `SCHEDULE_RUN_OUTPUT_RETENTION_DAYS` + +Type `int`, default `90`. + +Days scheduled-run output is kept. + + +## Sandbox + +The app is a CLIENT of an always-on runner; defaults are safe so app import never fails unconfigured. + +### `SANDBOX_BACKEND` + +Type `str`, default `jupyter`. + +Sandbox backend: jupyter (self-host) or daytona (Daytona Cloud). + +### `SANDBOX_GATEWAY_URL` + +Type `str`, default `http://localhost:8888`. + +URL of the Jupyter Kernel Gateway runner (the docsgpt-sandbox service). + +### `SANDBOX_GATEWAY_AUTH_TOKEN` + +Type `str`, default unset. + +Gateway auth token, if set. + +### `SANDBOX_KERNEL_NAME` + +Type `str`, default `docsgpt-python`. + +Kernelspec per session. The env-scrubbing docsgpt-python spec keeps kernel code from reading the gateway token or operator secrets from os.environ; the stock python3 spec inherits the gateway env verbatim and must not be used with untrusted code. + +### `SANDBOX_MAX_TTL` + +Type `int`, default `1200`. + +Hard cap (s) on agent-selectable keep-alive TTL. + +### `SANDBOX_MAX_SESSIONS` + +Type `int`, default `32`. + +Concurrent live sessions per process, backend-agnostic; at the cap an LRU-idle session is evicted. 0 or negative disables the cap. + +### `SANDBOX_EXEC_TIMEOUT` + +Type `int`, default `60`. + +Default wall-clock cap (s) per exec call. + +### `SANDBOX_HTTP_TIMEOUT` + +Type `int`, default `10`. + +Fixed cap (s) for REST control calls (create/delete/alive/interrupt). + +### `SANDBOX_MAX_OUTPUT_BYTES` + +Type `int`, default `8388608`. + +Cap on buffered stdout+stderr per exec. + +### `SANDBOX_MAX_FILE_BYTES` + +Type `int`, default `10485760`. + +Cap on get_file size routed through stdout. + +### `SANDBOX_MAX_INPUT_BYTES` + +Type `int`, default `26214400`. + +Cap on an input document staged into a sandbox session. + +### `SANDBOX_MEMORY` + +Type `str`, default `1g`. + +Docker mem_limit for the runner container. Consumed by the docsgpt-sandbox compose service, not the app; part of the untrusted-code security boundary. + +### `SANDBOX_CPUS` + +Type `str`, default `1.0`. + +Docker CPU quota for the runner container. Consumed by the docsgpt-sandbox compose service, not the app; part of the untrusted-code security boundary. + +### `DAYTONA_API_KEY` + +Type `str`, default unset. + +Daytona Cloud API key (secret). + +### `DAYTONA_API_URL` + +Type `str`, default unset. + +Override for the Daytona API base URL, if self-targeting. + +### `DAYTONA_TARGET` + +Type `str`, default unset. + +Daytona region/target, e.g. "us". + +### `DAYTONA_SNAPSHOT` + +Type `str`, default unset. + +Image for new sandboxes; render libs via scripts/build_daytona_snapshot.py. + +### `DAYTONA_LANGUAGE` + +Type `str`, default `python`. + +Default runtime language for created sandboxes. + +### `DAYTONA_AUTO_STOP_INTERVAL` + +Type `int`, default `15`. + +Minutes idle before Daytona auto-stops a sandbox (0 disables). + +### `DAYTONA_AUTO_DELETE_INTERVAL` + +Type `int`, default `60`. + +Minutes after stop before Daytona auto-deletes a sandbox (-1 disables). + +### `DAYTONA_MAX_SANDBOXES` + +Type `int`, default `50`. + +Cap on concurrent live Daytona sandboxes (cost-DoS guard). + + +## Speech + +Voice providers and transcription options. + +### `TTS_PROVIDER` + +Type `str`, default `google_tts`. + +Text-to-speech provider: google_tts, elevenlabs, or none to switch it off. + +### `ELEVENLABS_API_KEY` + +Type `str`, default unset. + +ElevenLabs API key. + +### `STT_PROVIDER` + +Type `str`, default `openai`. + +Speech-to-text provider: openai, faster_whisper, or none to switch it off. + +### `OPENAI_STT_MODEL` + +Type `str`, default `gpt-4o-mini-transcribe`. + +OpenAI transcription model. + +### `STT_LANGUAGE` + +Type `str`, default unset. + +Language hint for transcription; unset auto-detects. + +### `STT_MAX_FILE_SIZE_MB` + +Type `int`, default `50`. + +Cap on an audio file accepted for transcription. + +### `STT_ENABLE_TIMESTAMPS` + +Type `bool`, default `false`. + +Return word/segment timestamps. + +### `STT_ENABLE_DIARIZATION` + +Type `bool`, default `false`. + +Label speakers in the transcript. diff --git a/docs/content/Deploying/_meta.js b/docs/content/Deploying/_meta.js index 4865a40e..105b9efd 100644 --- a/docs/content/Deploying/_meta.js +++ b/docs/content/Deploying/_meta.js @@ -3,6 +3,10 @@ export default { "title": "⚙️ App Configuration", "href": "/Deploying/DocsGPT-Settings" }, + "Settings-Reference": { + "title": "📖 Settings Reference", + "href": "/Deploying/Settings-Reference" + }, "OIDC-SSO": { "title": "🔐 SSO with OIDC", "href": "/Deploying/OIDC-SSO" diff --git a/docsgpt/core/settings/reference.py b/docsgpt/core/settings/reference.py new file mode 100644 index 00000000..02d3721d --- /dev/null +++ b/docsgpt/core/settings/reference.py @@ -0,0 +1,169 @@ +"""Render the settings reference page from the ``Settings`` definitions. + +The page under ``docs/content/Deploying/Settings-Reference.mdx`` is generated +from the field types, defaults and descriptions in this package, so the model +is the single source of truth. Regenerate it after changing a setting:: + + python -m docsgpt.core.settings.reference --write + +``--check`` exits non-zero when the checked-in page is stale; the test suite +runs the same comparison. +""" + +from __future__ import annotations + +import argparse +import inspect +import json +import sys +import types +import typing +from pathlib import Path +from typing import Any, Literal, Optional, Union + +from pydantic import AliasChoices +from pydantic.fields import FieldInfo + +from docsgpt.core.paths import home_dir +from docsgpt.core.settings import SETTINGS_GROUPS, Settings + +REFERENCE_PATH = Path("docs") / "content" / "Deploying" / "Settings-Reference.mdx" + +HOME_PLACEHOLDER = "" + +_HEADER = """\ +--- +title: Settings Reference +description: Every DocsGPT setting, grouped by domain, with its type, default and purpose. +--- + +{/* GENERATED FILE. Do not edit by hand: run `python -m docsgpt.core.settings.reference --write`. */} + +# Settings Reference + +Every setting DocsGPT reads, generated from `docsgpt/core/settings/`. Each one is +an environment variable of the same name, set in `.env` or the process +environment; see [App Configuration](/Deploying/DocsGPT-Settings) for how the +file is found and for worked examples. `` below is the data home +described there. +""" + + +def _mdx(text: str) -> str: + """Escape prose for MDX, where braces open expressions and ``<`` opens JSX.""" + return text.replace("{", "\\{").replace("}", "\\}").replace("<", "<") + + +def _type_name(annotation: Any) -> str: + origin = typing.get_origin(annotation) + if origin in (Union, types.UnionType): + args = [a for a in typing.get_args(annotation) if a is not type(None)] + return " | ".join(_type_name(a) for a in args) + if origin is Literal: + return " | ".join(json.dumps(v) for v in typing.get_args(annotation)) + if origin is not None: + return getattr(origin, "__name__", str(origin)) + return getattr(annotation, "__name__", str(annotation)) + + +def _default_text(field: FieldInfo) -> str: + value = field.default_factory() if field.default_factory is not None else field.default + if value is None: + return "unset" + if isinstance(value, bool): + return "`true`" if value else "`false`" + if isinstance(value, str): + value = value.replace(str(home_dir()), HOME_PLACEHOLDER) + return "`\"\"`" if value == "" else f"`{value}`" + if isinstance(value, (list, dict)): + return f"`{json.dumps(value)}`" + return f"`{value}`" + + +def _constraints(field: FieldInfo) -> list[str]: + out = [] + for item in field.metadata: + for attr, symbol in (("gt", ">"), ("ge", ">="), ("lt", "<"), ("le", "<=")): + if hasattr(item, attr): + out.append(f"{symbol} {getattr(item, attr)}") + return out + + +def _aliases(name: str, field: FieldInfo) -> list[str]: + alias = field.validation_alias + if isinstance(alias, AliasChoices): + return [str(c) for c in alias.choices if str(c) != name] + if isinstance(alias, str) and alias != name: + return [alias] + return [] + + +def _render_field(name: str, field: FieldInfo) -> str: + facts = [f"Type `{_type_name(field.annotation)}`", f"default {_default_text(field)}"] + constraints = _constraints(field) + if constraints: + facts.append("must be " + " and ".join(constraints)) + aliases = _aliases(name, field) + if aliases: + facts.append("also read from " + ", ".join(f"`{a}`" for a in aliases)) + # The facts are code spans, which MDX leaves alone; only prose needs escaping. + lines = [f"### `{name}`", "", ", ".join(facts) + "."] + if field.deprecated: + note = field.deprecated if isinstance(field.deprecated, str) else "This setting is deprecated." + lines += ["", f"**Deprecated.** {_mdx(str(note))}"] + if field.description: + lines += ["", _mdx(field.description)] + return "\n".join(lines) + + +def _group_intro(group: type) -> str: + doc = inspect.getdoc(group) or "" + return doc.split("\n\n", 1)[0].replace("\n", " ").strip() + + +def render_reference() -> str: + """The full reference page as MDX text.""" + parts = [_HEADER] + for title, group in SETTINGS_GROUPS: + parts.append(f"\n## {title}\n") + intro = _group_intro(group) + if intro: + parts.append(_mdx(intro) + "\n") + for name in group.model_fields: + parts.append(_render_field(name, Settings.model_fields[name]) + "\n") + return "\n".join(parts).rstrip("\n") + "\n" + + +def reference_path(root: Optional[Path] = None) -> Path: + """Where the generated page lives in a checkout; ``root`` defaults to the repository root.""" + if root is None: + root = Path(__file__).resolve().parents[3] + return root / REFERENCE_PATH + + +def main(argv: Optional[list[str]] = None) -> int: + parser = argparse.ArgumentParser(description=__doc__.split("\n\n", 1)[0]) + action = parser.add_mutually_exclusive_group() + action.add_argument("--write", action="store_true", help="write the page into the docs tree") + action.add_argument("--check", action="store_true", help="exit 1 if the checked-in page is stale") + args = parser.parse_args(argv) + + rendered = render_reference() + path = reference_path() + if args.write: + path.write_text(rendered, encoding="utf-8") + print(f"wrote {path}") + return 0 + if args.check: + current = path.read_text(encoding="utf-8") if path.exists() else "" + if current != rendered: + print(f"{path} is stale; run: python -m docsgpt.core.settings.reference --write", file=sys.stderr) + return 1 + print(f"{path} is up to date") + return 0 + sys.stdout.write(rendered) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py new file mode 100644 index 00000000..9c2f7d88 --- /dev/null +++ b/tests/core/test_settings.py @@ -0,0 +1,109 @@ +"""Contract tests for ``docsgpt.core.settings``. + +``Settings`` is composed from one ``SettingsGroup`` per domain; these tests pin +the properties that composition must keep (flat names, no duplicate fields, +validators from every group applied) and that the generated reference page +tracks the definitions. +""" + +from pathlib import Path + +import pytest + +from docsgpt.core.settings import SETTINGS_GROUPS, Settings, settings +from docsgpt.core.settings.reference import reference_path, render_reference + +SECRET_FIELDS = ( + "API_KEY", + "OPENAI_API_KEY", + "ANTHROPIC_API_KEY", + "GOOGLE_API_KEY", + "GROQ_API_KEY", + "HUGGINGFACE_API_KEY", + "NOVITA_API_KEY", + "EMBEDDINGS_KEY", + "FALLBACK_LLM_API_KEY", + "QDRANT_API_KEY", + "ELEVENLABS_API_KEY", + "INTERNAL_KEY", +) + + +@pytest.mark.unit +class TestComposition: + def test_every_group_field_is_a_flat_settings_attribute(self): + for _, group in SETTINGS_GROUPS: + for name in group.model_fields: + assert name in Settings.model_fields, name + assert hasattr(settings, name), name + + def test_no_field_is_defined_in_two_groups(self): + owners: dict[str, str] = {} + for title, group in SETTINGS_GROUPS: + for name in group.model_fields: + assert name not in owners, f"{name} is defined in both {owners[name]} and {title}" + owners[name] = title + assert len(owners) == len(Settings.model_fields) + + def test_every_field_has_a_description(self): + missing = [name for name, field in Settings.model_fields.items() if not field.description] + assert not missing, f"settings without a description: {missing}" + + def test_defaults_load_without_an_env_file(self, monkeypatch): + for name in Settings.model_fields: + monkeypatch.delenv(name, raising=False) + fresh = Settings(_env_file=None) + assert fresh.LLM_PROVIDER == "docsgpt" + assert fresh.VECTOR_STORE == "faiss" + + +@pytest.mark.unit +class TestValidators: + """Validators live on the group that owns the field; composition must keep all of them. + + Two groups defining a validator under the same method name would silently + keep only one, so this checks every secret field, across every group. + """ + + @pytest.mark.parametrize("name", SECRET_FIELDS) + @pytest.mark.parametrize("raw", ["None", "none", "", " "]) + def test_unset_secret_spellings_become_none(self, name, raw): + assert getattr(Settings.model_validate({name: raw}), name) is None + + @pytest.mark.parametrize("name", SECRET_FIELDS) + def test_secret_is_stripped(self, name): + assert getattr(Settings.model_validate({name: " k3y "}), name) == "k3y" + + def test_normalize_api_key_classmethod_is_kept(self): + assert Settings.normalize_api_key("None") is None + assert Settings.normalize_api_key(" x ") == "x" + assert Settings.normalize_api_key(42) == 42 + + def test_postgres_uris_are_normalized(self): + loaded = Settings.model_validate( + {"POSTGRES_URI": "postgres://u:p@h/db", "PGVECTOR_CONNECTION_STRING": "postgresql+psycopg://u:p@h/v"} + ) + assert loaded.POSTGRES_URI.startswith("postgresql+psycopg://") + assert loaded.PGVECTOR_CONNECTION_STRING.startswith("postgresql://") + + def test_legacy_docling_ocr_aliases_are_read(self): + loaded = Settings.model_validate({"DOCLING_OCR_ENABLED": "true", "DOCLING_OCR_MIN_CHARS_PER_PAGE": "7"}) + assert loaded.OCR_ENABLED is True + assert loaded.OCR_MIN_CHARS_PER_PAGE == 7 + + +@pytest.mark.unit +class TestReference: + def test_reference_lists_every_setting_once(self): + page = render_reference() + for name in Settings.model_fields: + assert page.count(f"### `{name}`") == 1, name + + def test_checked_in_reference_is_current(self): + path: Path = reference_path() + if not path.exists(): + pytest.skip("docs tree not present (installed package, not a checkout)") + assert path.read_text(encoding="utf-8") == render_reference(), ( + "docs/content/Deploying/Settings-Reference.mdx is stale; " + "run: python -m docsgpt.core.settings.reference --write" + ) From 95d0799494c579dc3df3973bf1696c589f997528 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:08:50 +0100 Subject: [PATCH 053/130] refactor(settings): tighten types on closed choices, containers and bounds Enum-like settings whose allowed values were only listed in a comment are now Literal types, so a typo fails at startup with a message naming the allowed values instead of falling through to a default with a warning (or, for VECTOR_STORE, failing on first use): AUTH_TYPE, VECTOR_STORE, STORAGE_TYPE, URL_STRATEGY, OCR_BACKEND, OCR_ENGINE, SANDBOX_BACKEND, DOC_PARSER_ENGINE, TTS_PROVIDER, STT_PROVIDER Each keeps a before-validator that strips and lower-cases the value, since the registries that consume them already lower-cased at the use site, and AUTH_TYPE maps the "None"/"none"/"" spellings a .env file carries to None (it was the string "None" before, which only worked because nothing compared against it). An empty TTS/STT provider still means "off". LLM_PROVIDER stays a plain str because providers are plugin-extensible. Containers are typed (dict[str, int], list[str], dict[str, Any]) instead of bare dict/list, six fields that were Optional with a non-None default are plain, and integer settings whose description already states a range carry it as a constraint (ge=0 for "0 disables", ge=1 for counts that cannot be zero, 0 < threshold <= 1). --- docs/content/Deploying/Settings-Reference.mdx | 72 +++++++++---------- docsgpt/core/settings/_shared.py | 9 ++- docsgpt/core/settings/agents.py | 9 +-- docsgpt/core/settings/auth.py | 16 +++-- docsgpt/core/settings/connectors.py | 4 +- docsgpt/core/settings/embeddings.py | 3 +- docsgpt/core/settings/events.py | 4 +- docsgpt/core/settings/guardrails.py | 6 +- docsgpt/core/settings/ingestion.py | 20 ++++-- docsgpt/core/settings/llm.py | 2 +- docsgpt/core/settings/ocr.py | 16 +++-- docsgpt/core/settings/retrieval.py | 19 +++-- docsgpt/core/settings/sandbox.py | 17 +++-- docsgpt/core/settings/server.py | 2 +- docsgpt/core/settings/speech.py | 19 +++-- docsgpt/core/settings/storage.py | 15 ++-- docsgpt/core/settings/vectorstores.py | 14 ++-- docsgpt/core/settings/workers.py | 3 +- tests/core/test_settings.py | 53 ++++++++++++++ 19 files changed, 206 insertions(+), 97 deletions(-) diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 6b935de9..59b20430 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -20,9 +20,9 @@ How users authenticate: none, a shared token, per-session JWTs, or OIDC SSO. ### `AUTH_TYPE` -Type `str`, default unset. +Type `"simple_jwt" | "session_jwt" | "oidc"`, default unset. -Authentication mode: simple_jwt, session_jwt, oidc, or unset for no authentication. +Authentication mode: simple_jwt, session_jwt, oidc, or unset (None) for no authentication. ### `JWT_SECRET_KEY` @@ -86,7 +86,7 @@ Override for the callback URL; default is <request host>/api/auth/oidc/callba ### `OIDC_SESSION_LIFETIME_SECONDS` -Type `int`, default `28800`. +Type `int`, default `28800`, must be > 0. Lifetime of the minted session JWT in seconds (8h). @@ -354,13 +354,13 @@ Truncate each remote embed input to N tokens (overflow is lost). ### `EMBEDDINGS_BATCH_SIZE` -Type `int`, default `32`. +Type `int`, default `32`, must be >= 1. Chunks per store transaction and per remote embed request. ### `EMBEDDINGS_MODEL_BATCH_SIZE` -Type `int`, default `1`. +Type `int`, default `1`, must be >= 1. Documents per local ONNX forward pass. Each pass pads to its longest input, and that waste grows with the square of chunk length: at 1250 tokens, 32 peaked at 6.6 GB, 1 at 2.9 GB. @@ -413,9 +413,9 @@ Which vector store answers searches and how retrieval fans out across sources. ### `VECTOR_STORE` -Type `str`, default `faiss`. +Type `"faiss" | "elasticsearch" | "mongodb" | "qdrant" | "milvus" | "pgvector"`, default `faiss`. -Vector store backend: faiss, elasticsearch, mongodb, qdrant, milvus or pgvector. +Vector store backend. ### `RETRIEVERS_ENABLED` @@ -425,7 +425,7 @@ Retriever keys an agent may use; must match RetrieverCreator.retrievers registry ### `RETRIEVAL_MAX_PARALLEL_SOURCES` -Type `int`, default `4`. +Type `int`, default `4`, must be >= 1. Concurrent per-source searches in one retrieval; the query is embedded once and shared. @@ -580,7 +580,7 @@ pgvector connection string. postgres://, postgresql:// and postgresql+psycopg:// ### `PGVECTOR_POOL_MAX_SIZE` -Type `int`, default `8`. +Type `int`, default `8`, must be >= 0. Per-process connection pool size; 0 uses one direct connection per store. @@ -680,13 +680,13 @@ Broker visibility timeout in seconds. Must exceed the longest legitimate task ru ### `CELERY_WORKER_MAX_MEMORY_PER_CHILD` -Type `int`, default `4194304`. +Type `int`, default `4194304`, must be >= 0. Recycle a prefork child past this resident size in KB; backstops docling/torch heap growth. Checked between tasks, so it does not bound the peak within one. 0 disables. ### `CELERY_WORKER_MAX_TASKS_PER_CHILD` -Type `int`, default `0`. +Type `int`, default `0`, must be >= 0. Recycle a worker child after N tasks; 0 disables. @@ -763,7 +763,7 @@ Send images to a remote parser. ### `DOC_PARSER_ENGINE` -Type `str`, default `anydoc`. +Type `"anydoc" | "docling"`, default `anydoc`. Document parser for source ingestion, chat attachments and the read_document tool. "anydoc" (default): firecrawl-anydoc, a Rust converter with no ML models; milliseconds per file, ~100 MB peak RSS. "docling": the layout/table-model pipeline (optional install; needed for read_document's structured output and the docling OCR backend). Files anydoc cannot convert (scanned PDFs, malformed input) fall back to docling when it is installed, otherwise to the native OCR parsers (OCR on) or the legacy parsers. Rollback to the previous behaviour is this one variable. @@ -793,7 +793,7 @@ Largest HTML/XML docling will parse, in bytes. ### `MARKUP_MAX_BYTES` -Type `int`, default `8000000`. +Type `int`, default `8000000`, must be >= 0. HTML/XHTML larger than this (bytes) are head-truncated before the markdownify parser runs (the anydoc engine's HTML path). The tree that path builds costs ~50x the input (30 MB of HTML measured at 1.6 GB RSS) and the upload cap is 100 MB, so the gate is what keeps one upload from taking the ingest worker down. 0 disables it. @@ -841,13 +841,13 @@ Cap on the pixel count of an image passed to an agent. ### `GITHUB_INGEST_MAX_FILE_BYTES` -Type `int`, default `1048576`. +Type `int`, default `1048576`, must be >= 0. Skip GitHub repo blobs larger than this (0 = no cap). ### `GITHUB_INGEST_MAX_WORKERS` -Type `int`, default `8`. +Type `int`, default `8`, must be >= 1. Parallel file fetches per GitHub repo ingest. @@ -877,7 +877,7 @@ Absolute ceiling on the size-scaled parse window, in seconds. ### `DOCUMENT_PARSE_MAX_BYTES` -Type `int`, default `0`. +Type `int`, default `0`, must be >= 0. Cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES). @@ -912,13 +912,13 @@ OCR scanned PDFs and images attached to a chat. ### `OCR_BACKEND` -Type `str`, default `auto`. +Type `"auto" | "docling" | "native"`, default `auto`. Which stack runs OCR when it is on. auto: docling when installed, otherwise native. docling: the layout-model pipeline (hybrid region OCR, reading order, table structure); needs the optional docling extra. native: pypdfium2/Pillow page rendering straight into tesseract or a DeepSeek-OCR endpoint (docsgpt/parser/file/ocr_parser.py); no ML models in the worker, tables come out as text lines under tesseract. ### `OCR_ENGINE` -Type `str`, default `tesseract`. +Type `"tesseract" | "deepseek" | "auto" | "ocrmac" | "rapidocr"`, default `tesseract`. OCR engine used when OCR is on. Benched 2026-08 on EN/ZH/table/degraded scans (docs/Guides/ocr has the menu). tesseract (recommended): best classic-engine accuracy (perfect EN word recall, 0.000 bilingual CER, 100% table cells), ~35 MB, CPU-only; needs the system binary and language packs, an optional install like every OCR dependency (build with INSTALL_TESSERACT=true, or apt/brew install tesseract-ocr for a local run); both backends. deepseek: DeepSeek-OCR against an Ollama/vLLM endpoint (OCR_DEEPSEEK_*); best table/CJK quality, the worker stays light (no layout models) but each page costs seconds on the model server; both backends. auto: docling's pick, ocrmac on macOS (excellent), rapidocr on Linux (silently shreds some long text lines; avoid as a server default). ocrmac | rapidocr: force one of those. auto/ocrmac/rapidocr exist only inside docling; the native backend runs tesseract for them. An engine that is not installed degrades (docling: to auto) with a warning instead of failing the parse. @@ -954,7 +954,7 @@ Native backend only: resolution at which pages without a text layer are rendered ### `OCR_MIN_CHARS_PER_PAGE` -Type `int`, default `20`, also read from `DOCLING_OCR_MIN_CHARS_PER_PAGE`. +Type `int`, default `20`, must be >= 0, also read from `DOCLING_OCR_MIN_CHARS_PER_PAGE`. Chars-per-page floor below which an OCR'd PDF/image parse is treated as an OCR dropout rather than as content (long-running docling workers were observed returning zero characters for every scanned page after a long scanned PDF, with no error). docling retries once on a fresh full-page-OCR converter; both backends then fail loudly instead of indexing an empty document. 0 disables the guard. @@ -965,13 +965,13 @@ Local disk or an S3-compatible bucket, and how download URLs are produced. ### `STORAGE_TYPE` -Type `str`, default `local`. +Type `"local" | "s3"`, default `local`. -File storage backend: local or s3. +File storage backend. ### `URL_STRATEGY` -Type `str`, default `backend`. +Type `"backend" | "s3"`, default `backend`. How download links are produced: backend (streamed through the API) or s3 (presigned URLs). @@ -1143,7 +1143,7 @@ Bounds uvicorn's shutdown drain (uvicorn_worker doesn't forward --graceful-timeo ### `WSGI_THREADPOOL_WORKERS` -Type `int`, default `96`. +Type `int`, default `96`, must be >= 1. Threads serving the WSGI (Flask) part of the app under the ASGI server. @@ -1166,7 +1166,7 @@ Internal SSE push channel (notifications and durable replay journal). False make ### `EVENTS_STREAM_MAXLEN` -Type `int`, default `1000`. +Type `int`, default `1000`, must be >= 1. Per-user durable backlog cap in entries; ~24h of replay at typical rates. @@ -1178,7 +1178,7 @@ Interval between SSE keepalive comments. ### `SSE_MAX_CONCURRENT_PER_USER` -Type `int`, default `8`. +Type `int`, default `8`, must be >= 0. Simultaneous SSE connections per user; each holds a pooled async Redis connection for its lifetime. 8 covers multi-tab use without one user starving the pool. 0 disables. @@ -1190,7 +1190,7 @@ Pool size of the async Redis client behind the event-loop routes, per process. E ### `EVENTS_REPLAY_MAX_PER_REQUEST` -Type `int`, default `200`. +Type `int`, default `200`, must be >= 1. Backlog entries XRANGE returns per /api/events snapshot. Bounds what one replay moves from Redis to the wire: a client looping Last-Event-ID reconnects enumerates at most this many per round-trip. @@ -1291,7 +1291,7 @@ Pre-fetch retrieval before the agent's first turn. ### `TOOL_RESULT_MAX_TOKENS` -Type `int`, default `20000`. +Type `int`, default `20000`, must be >= 0. Cap on one tool result entering the LLM context (0 disables); journal and DB keep it whole. @@ -1303,7 +1303,7 @@ Compress long conversations once they approach the context window. ### `COMPRESSION_THRESHOLD_PERCENTAGE` -Type `float`, default `0.8`. +Type `float`, default `0.8`, must be > 0 and <= 1. Fraction of the context window at which compression triggers. @@ -1327,7 +1327,7 @@ Keep only the last N compression points to prevent DB bloat. ### `COMPRESSION_RECENT_FIELD_MAX_TOKENS` -Type `int`, default `8000`. +Type `int`, default `8000`, must be >= 0. Per-field cap on the verbatim tail kept after a compression point (0 disables). @@ -1474,7 +1474,7 @@ The app is a CLIENT of an always-on runner; defaults are safe so app import neve ### `SANDBOX_BACKEND` -Type `str`, default `jupyter`. +Type `"jupyter" | "daytona"`, default `jupyter`. Sandbox backend: jupyter (self-host) or daytona (Daytona Cloud). @@ -1582,13 +1582,13 @@ Default runtime language for created sandboxes. ### `DAYTONA_AUTO_STOP_INTERVAL` -Type `int`, default `15`. +Type `int`, default `15`, must be >= 0. Minutes idle before Daytona auto-stops a sandbox (0 disables). ### `DAYTONA_AUTO_DELETE_INTERVAL` -Type `int`, default `60`. +Type `int`, default `60`, must be >= -1. Minutes after stop before Daytona auto-deletes a sandbox (-1 disables). @@ -1605,9 +1605,9 @@ Voice providers and transcription options. ### `TTS_PROVIDER` -Type `str`, default `google_tts`. +Type `"google_tts" | "elevenlabs" | "none"`, default `google_tts`. -Text-to-speech provider: google_tts, elevenlabs, or none to switch it off. +Text-to-speech provider; none switches it off. ### `ELEVENLABS_API_KEY` @@ -1617,9 +1617,9 @@ ElevenLabs API key. ### `STT_PROVIDER` -Type `str`, default `openai`. +Type `"openai" | "faster_whisper" | "none"`, default `openai`. -Speech-to-text provider: openai, faster_whisper, or none to switch it off. +Speech-to-text provider; none switches it off. ### `OPENAI_STT_MODEL` diff --git a/docsgpt/core/settings/_shared.py b/docsgpt/core/settings/_shared.py index 4eb5c48d..235ee07b 100644 --- a/docsgpt/core/settings/_shared.py +++ b/docsgpt/core/settings/_shared.py @@ -9,7 +9,7 @@ domain's definitions live in their own module. from __future__ import annotations -from typing import Optional +from typing import Any, Optional from pydantic_settings import BaseSettings, SettingsConfigDict @@ -20,6 +20,13 @@ class SettingsGroup(BaseSettings): model_config = SettingsConfigDict(extra="ignore") +def normalize_choice(value: Any) -> Any: + """Case-fold a closed-choice setting so ``PGVector`` and ``pgvector`` are the same choice.""" + if isinstance(value, str): + return value.strip().lower() + return value + + def normalize_secret(value: Optional[str]) -> Optional[str]: """Map the ways an unset secret reaches us from ``.env`` to ``None``. diff --git a/docsgpt/core/settings/agents.py b/docsgpt/core/settings/agents.py index fdee1cc3..9b682786 100644 --- a/docsgpt/core/settings/agents.py +++ b/docsgpt/core/settings/agents.py @@ -14,11 +14,11 @@ class AgentSettings(SettingsGroup): AGENT_NAME: str = Field(default="classic", description="Default agent type for agentless chats.") DEFAULT_MAX_HISTORY: int = Field(default=150, description="Default number of history messages kept.") - DEFAULT_AGENT_LIMITS: dict = Field( + DEFAULT_AGENT_LIMITS: dict[str, int] = Field( default={"token_limit": 50000, "request_limit": 500}, description="Per-agent default quotas: tokens and requests.", ) - DEFAULT_CHAT_TOOLS: list = Field( + DEFAULT_CHAT_TOOLS: list[str] = Field( default=["memory", "read_webpage", "scheduler"], description=( "Config-free tools on by default in agentless chats. scheduler is dual-registered in " @@ -30,6 +30,7 @@ class AgentSettings(SettingsGroup): ENABLE_TOOL_PREFETCH: bool = Field(default=True, description="Pre-fetch retrieval before the agent's first turn.") TOOL_RESULT_MAX_TOKENS: int = Field( default=20000, + ge=0, description="Cap on one tool result entering the LLM context (0 disables); journal and DB keep it whole.", ) @@ -38,7 +39,7 @@ class AgentSettings(SettingsGroup): default=True, description="Compress long conversations once they approach the context window." ) COMPRESSION_THRESHOLD_PERCENTAGE: float = Field( - default=0.8, description="Fraction of the context window at which compression triggers." + default=0.8, gt=0, le=1, description="Fraction of the context window at which compression triggers." ) COMPRESSION_MODEL_OVERRIDE: Optional[str] = Field( default=None, description="Use a different model for compression; unset reuses the answer model." @@ -48,7 +49,7 @@ class AgentSettings(SettingsGroup): default=3, description="Keep only the last N compression points to prevent DB bloat." ) COMPRESSION_RECENT_FIELD_MAX_TOKENS: int = Field( - default=8000, description="Per-field cap on the verbatim tail kept after a compression point (0 disables)." + default=8000, ge=0, description="Per-field cap on the verbatim tail kept after a compression point (0 disables)." ) # Workflows. diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py index d0763934..10e69801 100644 --- a/docsgpt/core/settings/auth.py +++ b/docsgpt/core/settings/auth.py @@ -2,19 +2,19 @@ from __future__ import annotations -from typing import Optional +from typing import Literal, Optional from pydantic import Field, field_validator -from docsgpt.core.settings._shared import SettingsGroup, normalize_secret +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice, normalize_secret class AuthSettings(SettingsGroup): """How users authenticate: none, a shared token, per-session JWTs, or OIDC SSO.""" - AUTH_TYPE: Optional[str] = Field( + AUTH_TYPE: Optional[Literal["simple_jwt", "session_jwt", "oidc"]] = Field( default=None, - description="Authentication mode: simple_jwt, session_jwt, oidc, or unset for no authentication.", + description="Authentication mode: simple_jwt, session_jwt, oidc, or unset (None) for no authentication.", ) JWT_SECRET_KEY: str = Field( default="", @@ -51,7 +51,7 @@ class AuthSettings(SettingsGroup): default=None, description="Override for the callback URL; default is /api/auth/oidc/callback." ) OIDC_SESSION_LIFETIME_SECONDS: int = Field( - default=28800, description="Lifetime of the minted session JWT in seconds (8h)." + default=28800, gt=0, description="Lifetime of the minted session JWT in seconds (8h)." ) OIDC_PROVIDER_NAME: Optional[str] = Field( default=None, description='Sign-in button label, e.g. "Acme SSO".' @@ -85,3 +85,9 @@ class AuthSettings(SettingsGroup): @classmethod def _normalize_auth_secrets(cls, v): return normalize_secret(v) + + @field_validator("AUTH_TYPE", mode="before") + @classmethod + def _normalize_auth_type(cls, v): + # ``AUTH_TYPE=None`` and ``AUTH_TYPE=`` in .env both mean "no authentication". + return normalize_choice(normalize_secret(v)) diff --git a/docsgpt/core/settings/connectors.py b/docsgpt/core/settings/connectors.py index 303c00b6..676ffe64 100644 --- a/docsgpt/core/settings/connectors.py +++ b/docsgpt/core/settings/connectors.py @@ -15,7 +15,7 @@ class ConnectorSettings(SettingsGroup): # Google Drive integration. GOOGLE_CLIENT_ID: Optional[str] = Field(default=None, description="Google OAuth client id.") GOOGLE_CLIENT_SECRET: Optional[str] = Field(default=None, description="Google OAuth client secret.") - CONNECTOR_REDIRECT_BASE_URI: Optional[str] = Field( + CONNECTOR_REDIRECT_BASE_URI: str = Field( default="http://127.0.0.1:7091/api/connectors/callback", description="OAuth callback URL; register it as-is in your provider's console (e.g. GCP).", ) @@ -31,7 +31,7 @@ class ConnectorSettings(SettingsGroup): # Microsoft Entra ID (Azure AD) integration. MICROSOFT_CLIENT_ID: Optional[str] = Field(default=None, description="Azure AD application (client) id.") MICROSOFT_CLIENT_SECRET: Optional[str] = Field(default=None, description="Azure AD application client secret.") - MICROSOFT_TENANT_ID: Optional[str] = Field( + MICROSOFT_TENANT_ID: str = Field( default="common", description="Azure AD tenant id, or 'common' for multi-tenant." ) MICROSOFT_AUTHORITY: Optional[str] = Field( diff --git a/docsgpt/core/settings/embeddings.py b/docsgpt/core/settings/embeddings.py index 85de4d15..3ac749aa 100644 --- a/docsgpt/core/settings/embeddings.py +++ b/docsgpt/core/settings/embeddings.py @@ -32,10 +32,11 @@ class EmbeddingsSettings(SettingsGroup): default=None, description="Truncate each remote embed input to N tokens (overflow is lost)." ) EMBEDDINGS_BATCH_SIZE: int = Field( - default=32, description="Chunks per store transaction and per remote embed request." + default=32, ge=1, description="Chunks per store transaction and per remote embed request." ) EMBEDDINGS_MODEL_BATCH_SIZE: int = Field( default=1, + ge=1, description=( "Documents per local ONNX forward pass. Each pass pads to its longest input, and that waste grows " "with the square of chunk length: at 1250 tokens, 32 peaked at 6.6 GB, 1 at 2.9 GB." diff --git a/docsgpt/core/settings/events.py b/docsgpt/core/settings/events.py index 0ab01d4c..f8ca6742 100644 --- a/docsgpt/core/settings/events.py +++ b/docsgpt/core/settings/events.py @@ -18,11 +18,12 @@ class EventsSettings(SettingsGroup): ), ) EVENTS_STREAM_MAXLEN: int = Field( - default=1000, description="Per-user durable backlog cap in entries; ~24h of replay at typical rates." + default=1000, ge=1, description="Per-user durable backlog cap in entries; ~24h of replay at typical rates." ) SSE_KEEPALIVE_SECONDS: int = Field(default=15, ge=1, description="Interval between SSE keepalive comments.") SSE_MAX_CONCURRENT_PER_USER: int = Field( default=8, + ge=0, description=( "Simultaneous SSE connections per user; each holds a pooled async Redis connection for its lifetime. " "8 covers multi-tab use without one user starving the pool. 0 disables." @@ -40,6 +41,7 @@ class EventsSettings(SettingsGroup): ) EVENTS_REPLAY_MAX_PER_REQUEST: int = Field( default=200, + ge=1, description=( "Backlog entries XRANGE returns per /api/events snapshot. Bounds what one replay moves from Redis to " "the wire: a client looping Last-Event-ID reconnects enumerates at most this many per round-trip." diff --git a/docsgpt/core/settings/guardrails.py b/docsgpt/core/settings/guardrails.py index 81ac60f2..871f748b 100644 --- a/docsgpt/core/settings/guardrails.py +++ b/docsgpt/core/settings/guardrails.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import Optional +from typing import Any, Optional from pydantic import Field @@ -13,10 +13,10 @@ class GuardrailSettings(SettingsGroup): """Input/output checks every agent runs, and the floor no agent may weaken.""" GUARDRAILS_ENABLED: bool = Field(default=True, description="Master switch; False disables every stage.") - GUARDRAILS_CHECKS_ENABLED: list = Field( + GUARDRAILS_CHECKS_ENABLED: list[str] = Field( default=[], description="Allowlist of GuardrailCreator.checks keys; empty means every registered check." ) - GUARDRAILS_FLOOR: dict = Field( + GUARDRAILS_FLOOR: dict[str, Any] = Field( default={}, description=( "A GuardrailsConfig fragment every agent inherits and cannot weaken; agents may add controls or " diff --git a/docsgpt/core/settings/ingestion.py b/docsgpt/core/settings/ingestion.py index 38da376d..d44a1360 100644 --- a/docsgpt/core/settings/ingestion.py +++ b/docsgpt/core/settings/ingestion.py @@ -2,9 +2,11 @@ from __future__ import annotations -from pydantic import Field +from typing import Literal -from docsgpt.core.settings._shared import SettingsGroup +from pydantic import Field, field_validator + +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice class IngestionSettings(SettingsGroup): @@ -37,7 +39,7 @@ class IngestionSettings(SettingsGroup): ) PARSE_PDF_AS_IMAGE: bool = Field(default=False, description="Render PDF pages to images before parsing.") PARSE_IMAGE_REMOTE: bool = Field(default=False, description="Send images to a remote parser.") - DOC_PARSER_ENGINE: str = Field( + DOC_PARSER_ENGINE: Literal["anydoc", "docling"] = Field( default="anydoc", description=( 'Document parser for source ingestion, chat attachments and the read_document tool. "anydoc" ' @@ -66,6 +68,7 @@ class IngestionSettings(SettingsGroup): ) MARKUP_MAX_BYTES: int = Field( default=8_000_000, + ge=0, description=( "HTML/XHTML larger than this (bytes) are head-truncated before the markdownify parser runs (the " "anydoc engine's HTML path). The tree that path builds costs ~50x the input (30 MB of HTML measured " @@ -113,9 +116,9 @@ class IngestionSettings(SettingsGroup): default=16_777_216, description="Cap on the pixel count of an image passed to an agent." ) GITHUB_INGEST_MAX_FILE_BYTES: int = Field( - default=1048576, description="Skip GitHub repo blobs larger than this (0 = no cap)." + default=1048576, ge=0, description="Skip GitHub repo blobs larger than this (0 = no cap)." ) - GITHUB_INGEST_MAX_WORKERS: int = Field(default=8, description="Parallel file fetches per GitHub repo ingest.") + GITHUB_INGEST_MAX_WORKERS: int = Field(default=8, ge=1, description="Parallel file fetches per GitHub repo ingest.") # read_document parsing on a dedicated Celery queue (backend parser). DOCUMENT_PARSE_QUEUE: str = Field(default="parsing", description="Celery queue the parse_document task is routed to.") @@ -134,7 +137,7 @@ class IngestionSettings(SettingsGroup): default=900, description="Absolute ceiling on the size-scaled parse window, in seconds." ) DOCUMENT_PARSE_MAX_BYTES: int = Field( - default=0, description="Cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES)." + default=0, ge=0, description="Cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES)." ) DOCUMENT_MAX_DECOMPRESSED_BYTES: int = Field( default=300 * 1024 * 1024, description="Cap on bytes decompressed from an archive handed to read_document." @@ -142,3 +145,8 @@ class IngestionSettings(SettingsGroup): DOCUMENT_MAX_ARCHIVE_ENTRIES: int = Field( default=10000, description="Cap on entries in an archive handed to read_document." ) + + @field_validator("DOC_PARSER_ENGINE", mode="before") + @classmethod + def _normalize_parser_engine(cls, v): + return normalize_choice(v) diff --git a/docsgpt/core/settings/llm.py b/docsgpt/core/settings/llm.py index 232ad3be..11c17cac 100644 --- a/docsgpt/core/settings/llm.py +++ b/docsgpt/core/settings/llm.py @@ -59,7 +59,7 @@ class LLMSettings(SettingsGroup): DEFAULT_LLM_TOKEN_LIMIT: int = Field( default=128000, description="Context window assumed when the model is not found in the registry." ) - RESERVED_TOKENS: dict = Field( + RESERVED_TOKENS: dict[str, int] = Field( default={"system_prompt": 500, "current_query": 500, "safety_buffer": 1000}, description="Tokens held back from the context window for the system prompt, the query and a safety buffer.", ) diff --git a/docsgpt/core/settings/ocr.py b/docsgpt/core/settings/ocr.py index 264b21c5..2ffc0432 100644 --- a/docsgpt/core/settings/ocr.py +++ b/docsgpt/core/settings/ocr.py @@ -2,9 +2,11 @@ from __future__ import annotations -from pydantic import AliasChoices, Field +from typing import Literal -from docsgpt.core.settings._shared import SettingsGroup +from pydantic import AliasChoices, Field, field_validator + +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice class OCRSettings(SettingsGroup): @@ -25,7 +27,7 @@ class OCRSettings(SettingsGroup): validation_alias=AliasChoices("OCR_ATTACHMENTS_ENABLED", "DOCLING_OCR_ATTACHMENTS_ENABLED"), description="OCR scanned PDFs and images attached to a chat.", ) - OCR_BACKEND: str = Field( + OCR_BACKEND: Literal["auto", "docling", "native"] = Field( default="auto", description=( "Which stack runs OCR when it is on. auto: docling when installed, otherwise native. docling: the " @@ -35,7 +37,7 @@ class OCRSettings(SettingsGroup): "lines under tesseract." ), ) - OCR_ENGINE: str = Field( + OCR_ENGINE: Literal["tesseract", "deepseek", "auto", "ocrmac", "rapidocr"] = Field( default="tesseract", description=( "OCR engine used when OCR is on. Benched 2026-08 on EN/ZH/table/degraded scans (docs/Guides/ocr has " @@ -80,6 +82,7 @@ class OCRSettings(SettingsGroup): ) OCR_MIN_CHARS_PER_PAGE: int = Field( default=20, + ge=0, validation_alias=AliasChoices("OCR_MIN_CHARS_PER_PAGE", "DOCLING_OCR_MIN_CHARS_PER_PAGE"), description=( "Chars-per-page floor below which an OCR'd PDF/image parse is treated as an OCR dropout rather than " @@ -89,3 +92,8 @@ class OCRSettings(SettingsGroup): "guard." ), ) + + @field_validator("OCR_BACKEND", "OCR_ENGINE", mode="before") + @classmethod + def _normalize_ocr_choices(cls, v): + return normalize_choice(v) diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py index 39672d29..a0bf87f9 100644 --- a/docsgpt/core/settings/retrieval.py +++ b/docsgpt/core/settings/retrieval.py @@ -2,21 +2,20 @@ from __future__ import annotations -from typing import Optional +from typing import Literal, Optional -from pydantic import Field +from pydantic import Field, field_validator -from docsgpt.core.settings._shared import SettingsGroup +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice class RetrievalSettings(SettingsGroup): """Which vector store answers searches and how retrieval fans out across sources.""" - VECTOR_STORE: str = Field( - default="faiss", - description="Vector store backend: faiss, elasticsearch, mongodb, qdrant, milvus or pgvector.", + VECTOR_STORE: Literal["faiss", "elasticsearch", "mongodb", "qdrant", "milvus", "pgvector"] = Field( + default="faiss", description="Vector store backend." ) - RETRIEVERS_ENABLED: list = Field( + RETRIEVERS_ENABLED: list[str] = Field( default=["classic", "default"], description=( "Retriever keys an agent may use; must match RetrieverCreator.retrievers registry keys, NOT the " @@ -25,6 +24,7 @@ class RetrievalSettings(SettingsGroup): ) RETRIEVAL_MAX_PARALLEL_SOURCES: int = Field( default=4, + ge=1, description="Concurrent per-source searches in one retrieval; the query is embedded once and shared.", ) PER_SOURCE_RETRIEVAL_ENABLED: bool = Field( @@ -38,3 +38,8 @@ class RetrievalSettings(SettingsGroup): GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field( default=2000, description="Hard cap on chunks extracted per source (cost control)." ) + + @field_validator("VECTOR_STORE", mode="before") + @classmethod + def _normalize_vector_store(cls, v): + return normalize_choice(v) diff --git a/docsgpt/core/settings/sandbox.py b/docsgpt/core/settings/sandbox.py index 958734bc..96e4d6e0 100644 --- a/docsgpt/core/settings/sandbox.py +++ b/docsgpt/core/settings/sandbox.py @@ -2,17 +2,17 @@ from __future__ import annotations -from typing import Optional +from typing import Literal, Optional -from pydantic import Field +from pydantic import Field, field_validator -from docsgpt.core.settings._shared import SettingsGroup +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice class SandboxSettings(SettingsGroup): """The app is a CLIENT of an always-on runner; defaults are safe so app import never fails unconfigured.""" - SANDBOX_BACKEND: str = Field( + SANDBOX_BACKEND: Literal["jupyter", "daytona"] = Field( default="jupyter", description="Sandbox backend: jupyter (self-host) or daytona (Daytona Cloud)." ) SANDBOX_GATEWAY_URL: str = Field( @@ -78,11 +78,16 @@ class SandboxSettings(SettingsGroup): ) DAYTONA_LANGUAGE: str = Field(default="python", description="Default runtime language for created sandboxes.") DAYTONA_AUTO_STOP_INTERVAL: int = Field( - default=15, description="Minutes idle before Daytona auto-stops a sandbox (0 disables)." + default=15, ge=0, description="Minutes idle before Daytona auto-stops a sandbox (0 disables)." ) DAYTONA_AUTO_DELETE_INTERVAL: int = Field( - default=60, description="Minutes after stop before Daytona auto-deletes a sandbox (-1 disables)." + default=60, ge=-1, description="Minutes after stop before Daytona auto-deletes a sandbox (-1 disables)." ) DAYTONA_MAX_SANDBOXES: int = Field( default=50, description="Cap on concurrent live Daytona sandboxes (cost-DoS guard)." ) + + @field_validator("SANDBOX_BACKEND", mode="before") + @classmethod + def _normalize_sandbox_backend(cls, v): + return normalize_choice(v) diff --git a/docsgpt/core/settings/server.py b/docsgpt/core/settings/server.py index 84bb2510..afb79121 100644 --- a/docsgpt/core/settings/server.py +++ b/docsgpt/core/settings/server.py @@ -28,7 +28,7 @@ class ServerSettings(SettingsGroup): ), ) WSGI_THREADPOOL_WORKERS: int = Field( - default=96, description="Threads serving the WSGI (Flask) part of the app under the ASGI server." + default=96, ge=1, description="Threads serving the WSGI (Flask) part of the app under the ASGI server." ) V1_SESSION_TTL_SECONDS: int = Field( default=24 * 60 * 60, diff --git a/docsgpt/core/settings/speech.py b/docsgpt/core/settings/speech.py index 465c8fc5..da5c245e 100644 --- a/docsgpt/core/settings/speech.py +++ b/docsgpt/core/settings/speech.py @@ -2,22 +2,22 @@ from __future__ import annotations -from typing import Optional +from typing import Literal, Optional from pydantic import Field, field_validator -from docsgpt.core.settings._shared import SettingsGroup, normalize_secret +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice, normalize_secret class SpeechSettings(SettingsGroup): """Voice providers and transcription options.""" - TTS_PROVIDER: str = Field( - default="google_tts", description="Text-to-speech provider: google_tts, elevenlabs, or none to switch it off." + TTS_PROVIDER: Literal["google_tts", "elevenlabs", "none"] = Field( + default="google_tts", description="Text-to-speech provider; none switches it off." ) ELEVENLABS_API_KEY: Optional[str] = Field(default=None, description="ElevenLabs API key.") - STT_PROVIDER: str = Field( - default="openai", description="Speech-to-text provider: openai, faster_whisper, or none to switch it off." + STT_PROVIDER: Literal["openai", "faster_whisper", "none"] = Field( + default="openai", description="Speech-to-text provider; none switches it off." ) OPENAI_STT_MODEL: str = Field(default="gpt-4o-mini-transcribe", description="OpenAI transcription model.") STT_LANGUAGE: Optional[str] = Field(default=None, description="Language hint for transcription; unset auto-detects.") @@ -29,3 +29,10 @@ class SpeechSettings(SettingsGroup): @classmethod def _normalize_speech_secrets(cls, v): return normalize_secret(v) + + @field_validator("TTS_PROVIDER", "STT_PROVIDER", mode="before") + @classmethod + def _normalize_speech_providers(cls, v): + # An empty value has always meant "off"; keep that spelling working. + v = normalize_choice(v) + return "none" if v == "" else v diff --git a/docsgpt/core/settings/storage.py b/docsgpt/core/settings/storage.py index ac970a54..2fe75347 100644 --- a/docsgpt/core/settings/storage.py +++ b/docsgpt/core/settings/storage.py @@ -2,18 +2,18 @@ from __future__ import annotations -from typing import Optional +from typing import Literal, Optional -from pydantic import Field +from pydantic import Field, field_validator -from docsgpt.core.settings._shared import SettingsGroup +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice class StorageSettings(SettingsGroup): """Local disk or an S3-compatible bucket, and how download URLs are produced.""" - STORAGE_TYPE: str = Field(default="local", description="File storage backend: local or s3.") - URL_STRATEGY: str = Field( + STORAGE_TYPE: Literal["local", "s3"] = Field(default="local", description="File storage backend.") + URL_STRATEGY: Literal["backend", "s3"] = Field( default="backend", description="How download links are produced: backend (streamed through the API) or s3 (presigned URLs).", ) @@ -47,3 +47,8 @@ class StorageSettings(SettingsGroup): "S3_SECRET_ACCESS_KEY." ), ) + + @field_validator("STORAGE_TYPE", "URL_STRATEGY", mode="before") + @classmethod + def _normalize_storage_choices(cls, v): + return normalize_choice(v) diff --git a/docsgpt/core/settings/vectorstores.py b/docsgpt/core/settings/vectorstores.py index 13c504df..067d1df0 100644 --- a/docsgpt/core/settings/vectorstores.py +++ b/docsgpt/core/settings/vectorstores.py @@ -27,13 +27,13 @@ class VectorStoreSettings(SettingsGroup): ELASTIC_USERNAME: Optional[str] = Field(default=None, description="Elasticsearch username.") ELASTIC_PASSWORD: Optional[str] = Field(default=None, description="Elasticsearch password.") ELASTIC_URL: Optional[str] = Field(default=None, description="Elasticsearch URL.") - ELASTIC_INDEX: Optional[str] = Field(default="docsgpt", description="Elasticsearch index name.") + ELASTIC_INDEX: str = Field(default="docsgpt", description="Elasticsearch index name.") # Qdrant. - QDRANT_COLLECTION_NAME: Optional[str] = Field(default="docsgpt", description="Qdrant collection name.") + QDRANT_COLLECTION_NAME: str = Field(default="docsgpt", description="Qdrant collection name.") QDRANT_LOCATION: Optional[str] = Field(default=None, description="Qdrant location (':memory:' or a URL).") QDRANT_URL: Optional[str] = Field(default=None, description="Qdrant server URL.") - QDRANT_PORT: Optional[int] = Field(default=6333, description="Qdrant REST port.") + QDRANT_PORT: int = Field(default=6333, description="Qdrant REST port.") QDRANT_GRPC_PORT: int = Field(default=6334, description="Qdrant gRPC port.") QDRANT_PREFER_GRPC: bool = Field(default=False, description="Use gRPC instead of REST where possible.") QDRANT_HTTPS: Optional[bool] = Field(default=None, description="Use HTTPS for the Qdrant connection.") @@ -53,7 +53,7 @@ class VectorStoreSettings(SettingsGroup): ), ) PGVECTOR_POOL_MAX_SIZE: int = Field( - default=8, description="Per-process connection pool size; 0 uses one direct connection per store." + default=8, ge=0, description="Per-process connection pool size; 0 uses one direct connection per store." ) PGVECTOR_IVFFLAT_PROBES: Optional[int] = Field( default=None, @@ -61,7 +61,7 @@ class VectorStoreSettings(SettingsGroup): ) # Milvus. - MILVUS_COLLECTION_NAME: Optional[str] = Field(default="docsgpt", description="Milvus collection name.") + MILVUS_COLLECTION_NAME: str = Field(default="docsgpt", description="Milvus collection name.") MILVUS_URI: Optional[str] = Field( default_factory=lambda: str(home_dir() / "milvus_local.db"), description=( @@ -69,14 +69,14 @@ class VectorStoreSettings(SettingsGroup): "like the other local stores." ), ) - MILVUS_TOKEN: Optional[str] = Field(default="", description="Milvus auth token.") + MILVUS_TOKEN: str = Field(default="", description="Milvus auth token.") # LanceDB. LANCEDB_PATH: str = Field( default_factory=lambda: str(home_dir() / "data" / "lancedb"), description="LanceDB local data directory.", ) - LANCEDB_TABLE_NAME: Optional[str] = Field(default="docsgpts", description="LanceDB table for stored vectors.") + LANCEDB_TABLE_NAME: str = Field(default="docsgpts", description="LanceDB table for stored vectors.") @field_validator("PGVECTOR_CONNECTION_STRING", mode="before") @classmethod diff --git a/docsgpt/core/settings/workers.py b/docsgpt/core/settings/workers.py index fbc7daae..86246cab 100644 --- a/docsgpt/core/settings/workers.py +++ b/docsgpt/core/settings/workers.py @@ -24,13 +24,14 @@ class WorkerSettings(SettingsGroup): ) CELERY_WORKER_MAX_MEMORY_PER_CHILD: int = Field( default=4194304, + ge=0, description=( "Recycle a prefork child past this resident size in KB; backstops docling/torch heap growth. " "Checked between tasks, so it does not bound the peak within one. 0 disables." ), ) CELERY_WORKER_MAX_TASKS_PER_CHILD: int = Field( - default=0, description="Recycle a worker child after N tasks; 0 disables." + default=0, ge=0, description="Recycle a worker child after N tasks; 0 disables." ) API_URL: str = Field( default="http://localhost:7091", description="Backend URL the Celery worker calls back into." diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py index 9c2f7d88..a4d64bf3 100644 --- a/tests/core/test_settings.py +++ b/tests/core/test_settings.py @@ -9,6 +9,7 @@ tracks the definitions. from pathlib import Path import pytest +from pydantic import ValidationError from docsgpt.core.settings import SETTINGS_GROUPS, Settings, settings from docsgpt.core.settings.reference import reference_path, render_reference @@ -107,3 +108,55 @@ class TestReference: "docs/content/Deploying/Settings-Reference.mdx is stale; " "run: python -m docsgpt.core.settings.reference --write" ) + + +@pytest.mark.unit +class TestClosedChoices: + """Enum-like settings are Literal types: a typo fails at startup instead of falling through.""" + + @pytest.mark.parametrize("raw", ["None", "none", "", " "]) + def test_auth_type_unset_spellings(self, raw): + assert Settings.model_validate({"AUTH_TYPE": raw}).AUTH_TYPE is None + + @pytest.mark.parametrize( + ("name", "raw", "expected"), + [ + ("AUTH_TYPE", " OIDC ", "oidc"), + ("VECTOR_STORE", "PGVector", "pgvector"), + ("STORAGE_TYPE", "S3", "s3"), + ("URL_STRATEGY", "Backend", "backend"), + ("OCR_BACKEND", "Native", "native"), + ("OCR_ENGINE", "Tesseract ", "tesseract"), + ("SANDBOX_BACKEND", "Daytona", "daytona"), + ("DOC_PARSER_ENGINE", "Docling", "docling"), + ("TTS_PROVIDER", "ElevenLabs", "elevenlabs"), + ("STT_PROVIDER", "", "none"), + ("TTS_PROVIDER", "NONE", "none"), + ], + ) + def test_choices_are_case_insensitive(self, name, raw, expected): + assert getattr(Settings.model_validate({name: raw}), name) == expected + + @pytest.mark.parametrize( + ("name", "raw"), + [ + ("AUTH_TYPE", "basic"), + ("VECTOR_STORE", "lancedb"), + ("STORAGE_TYPE", "gcs"), + ("OCR_BACKEND", "paddle"), + ("SANDBOX_BACKEND", "docker"), + ("DOC_PARSER_ENGINE", "fast"), + ("STT_PROVIDER", "whisper"), + ], + ) + def test_unknown_choice_is_rejected(self, name, raw): + with pytest.raises(ValidationError): + Settings.model_validate({name: raw}) + + @pytest.mark.parametrize( + ("name", "raw"), + [("EMBEDDINGS_BATCH_SIZE", 0), ("COMPRESSION_THRESHOLD_PERCENTAGE", 1.5), ("UPLOAD_MAX_FILE_BYTES", 0)], + ) + def test_out_of_range_numbers_are_rejected(self, name, raw): + with pytest.raises(ValidationError): + Settings.model_validate({name: raw}) From a8dab8864d8c2b23e990bf95ffaa48e5228cb080 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:09:09 +0100 Subject: [PATCH 054/130] refactor(settings): validate cross-field rules in the model The "AUTH_TYPE=oidc requires OIDC_ISSUER, OIDC_CLIENT_ID and OIDC_FRONTEND_URL" check lived in app.py, so it only ran when the Flask app was imported; a worker or script with the same misconfiguration started fine. It is now a model validator on the auth group and runs wherever Settings is loaded, with the same message. DEPLOYMENT_TYPE, which app.py read straight from the environment to decide whether a missing JWT_SECRET_KEY is fatal, is a documented setting on the server group now, so it shows up in the reference like every other variable the app reads. --- docs/content/Deploying/Settings-Reference.mdx | 6 ++++++ docsgpt/app.py | 11 +---------- docsgpt/core/settings/auth.py | 14 +++++++++++++- docsgpt/core/settings/server.py | 7 +++++++ tests/core/test_settings.py | 17 ++++++++++++++++- 5 files changed, 43 insertions(+), 12 deletions(-) diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 59b20430..88e4149d 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -1111,6 +1111,12 @@ Public callback URL for MCP OAuth; unset derives it from CONNECTOR_REDIRECT_BASE Serving the UI, public URLs, and process-level knobs of the API server. +### `DEPLOYMENT_TYPE` + +Type `str`, default unset. + +Deployment class, e.g. cloud or production. A production class refuses to run without a configured JWT_SECRET_KEY instead of generating a local one on disk. + ### `SERVE_UI` Type `bool`, default `true`. diff --git a/docsgpt/app.py b/docsgpt/app.py index 0913218f..4a11937e 100644 --- a/docsgpt/app.py +++ b/docsgpt/app.py @@ -1,5 +1,4 @@ import logging -import os import platform import uuid @@ -170,16 +169,8 @@ def enforce_document_upload_request_size_limit(): # only local development may use the atomic filesystem fallback. settings.JWT_SECRET_KEY = resolve_jwt_secret_key( settings.JWT_SECRET_KEY, - os.getenv("DEPLOYMENT_TYPE"), + settings.DEPLOYMENT_TYPE, ) -if settings.AUTH_TYPE == "oidc": - _missing_oidc = [ - name - for name in ("OIDC_ISSUER", "OIDC_CLIENT_ID", "OIDC_FRONTEND_URL") - if not getattr(settings, name) - ] - if _missing_oidc: - raise RuntimeError(f"AUTH_TYPE=oidc requires settings: {', '.join(_missing_oidc)}") SIMPLE_JWT_TOKEN = None if settings.AUTH_TYPE == "simple_jwt": payload = {"sub": "local"} diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py index 10e69801..04b039bc 100644 --- a/docsgpt/core/settings/auth.py +++ b/docsgpt/core/settings/auth.py @@ -4,11 +4,15 @@ from __future__ import annotations from typing import Literal, Optional -from pydantic import Field, field_validator +from pydantic import Field, field_validator, model_validator from docsgpt.core.settings._shared import SettingsGroup, normalize_choice, normalize_secret +#: Settings an OIDC deployment cannot run without; checked when AUTH_TYPE=oidc. +OIDC_REQUIRED = ("OIDC_ISSUER", "OIDC_CLIENT_ID", "OIDC_FRONTEND_URL") + + class AuthSettings(SettingsGroup): """How users authenticate: none, a shared token, per-session JWTs, or OIDC SSO.""" @@ -91,3 +95,11 @@ class AuthSettings(SettingsGroup): def _normalize_auth_type(cls, v): # ``AUTH_TYPE=None`` and ``AUTH_TYPE=`` in .env both mean "no authentication". return normalize_choice(normalize_secret(v)) + + @model_validator(mode="after") + def _require_oidc_settings(self): + if self.AUTH_TYPE == "oidc": + missing = [name for name in OIDC_REQUIRED if not getattr(self, name)] + if missing: + raise ValueError(f"AUTH_TYPE=oidc requires settings: {', '.join(missing)}") + return self diff --git a/docsgpt/core/settings/server.py b/docsgpt/core/settings/server.py index afb79121..407ef4aa 100644 --- a/docsgpt/core/settings/server.py +++ b/docsgpt/core/settings/server.py @@ -12,6 +12,13 @@ from docsgpt.core.settings._shared import SettingsGroup class ServerSettings(SettingsGroup): """Serving the UI, public URLs, and process-level knobs of the API server.""" + DEPLOYMENT_TYPE: Optional[str] = Field( + default=None, + description=( + "Deployment class, e.g. cloud or production. A production class refuses to run without a " + "configured JWT_SECRET_KEY instead of generating a local one on disk." + ), + ) SERVE_UI: bool = Field( default=True, description="Serve the web UI shipped in the package (docsgpt/static) from the API process." ) diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py index a4d64bf3..0757d8a4 100644 --- a/tests/core/test_settings.py +++ b/tests/core/test_settings.py @@ -110,6 +110,21 @@ class TestReference: ) +@pytest.mark.unit +class TestCrossFieldRules: + OIDC = {"OIDC_ISSUER": "https://idp.example/", "OIDC_CLIENT_ID": "docsgpt", "OIDC_FRONTEND_URL": "http://app"} + + def test_oidc_requires_issuer_client_and_frontend(self): + with pytest.raises(ValidationError, match="AUTH_TYPE=oidc requires settings: OIDC_CLIENT_ID, OIDC_FRONTEND_URL"): + Settings.model_validate({"AUTH_TYPE": "oidc", "OIDC_ISSUER": self.OIDC["OIDC_ISSUER"]}) + + def test_oidc_with_required_settings_loads(self): + assert Settings.model_validate({"AUTH_TYPE": "OIDC", **self.OIDC}).AUTH_TYPE == "oidc" + + def test_oidc_settings_are_not_required_for_other_modes(self): + assert Settings.model_validate({"AUTH_TYPE": "session_jwt"}).OIDC_ISSUER is None + + @pytest.mark.unit class TestClosedChoices: """Enum-like settings are Literal types: a typo fails at startup instead of falling through.""" @@ -121,7 +136,7 @@ class TestClosedChoices: @pytest.mark.parametrize( ("name", "raw", "expected"), [ - ("AUTH_TYPE", " OIDC ", "oidc"), + ("AUTH_TYPE", " Session_JWT ", "session_jwt"), ("VECTOR_STORE", "PGVector", "pgvector"), ("STORAGE_TYPE", "S3", "s3"), ("URL_STRATEGY", "Backend", "backend"), From d70644743b068fe56ccc1bb1ee7918c693130f6e Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:11:00 +0100 Subject: [PATCH 055/130] refactor(settings): deprecate SAGEMAKER_* and drop two unused settings SAGEMAKER_REGION, SAGEMAKER_ACCESS_KEY and SAGEMAKER_SECRET_KEY survive only as a fallback for the S3_* credentials. They carry Field(deprecated=...) now, so any read emits a DeprecationWarning naming the replacement and the generated reference shows the notice. The S3 store is the one sanctioned reader; it silences that warning locally because it already logs its own operator-facing one when the fallback is actually used. DEFAULT_MAX_HISTORY was referenced nowhere. RETRIEVERS_ENABLED was read by no code at all, while two docs pages described it as an enforced allow-list; both the setting and those claims are removed. --- docs/content/Deploying/DocsGPT-Settings.mdx | 1 - docs/content/Deploying/Settings-Reference.mdx | 18 ++++++------------ .../Sources/Per-source-configuration.mdx | 2 -- docsgpt/core/settings/agents.py | 1 - docsgpt/core/settings/retrieval.py | 7 ------- docsgpt/core/settings/storage.py | 3 +++ docsgpt/storage/s3.py | 11 ++++++++--- tests/core/test_settings.py | 15 +++++++++++---- 8 files changed, 28 insertions(+), 30 deletions(-) diff --git a/docs/content/Deploying/DocsGPT-Settings.mdx b/docs/content/Deploying/DocsGPT-Settings.mdx index 534413d6..0e55740e 100644 --- a/docs/content/Deploying/DocsGPT-Settings.mdx +++ b/docs/content/Deploying/DocsGPT-Settings.mdx @@ -461,7 +461,6 @@ These control how sources are retrieved and whether the advanced RAG features ar | Setting | Default | Description | | --- | --- | --- | -| `RETRIEVERS_ENABLED` | `["classic", "default"]` | Allow-list of retrievers usable instance-wide. Valid keys: `classic`, `default`, `hybrid`, `graphrag`. A per-source `retriever` must be within this list. | | `PER_SOURCE_RETRIEVAL_ENABLED` | `true` | Master switch for per-source retrieval config. When `false`, all sources fall back to the classic retriever regardless of their stored config. | | `GRAPHRAG_ENABLED` | `false` | Enable [GraphRAG](/Sources/GraphRAG). Requires `VECTOR_STORE=pgvector`. | | `GRAPHRAG_EXTRACTION_MODEL` | unset | Model used for ingest-time graph extraction. Unset reuses the instance default model. | diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 88e4149d..efc94fc5 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -417,12 +417,6 @@ Type `"faiss" | "elasticsearch" | "mongodb" | "qdrant" | "milvus" | "pgvector"`, Vector store backend. -### `RETRIEVERS_ENABLED` - -Type `list`, default `["classic", "default"]`. - -Retriever keys an agent may use; must match RetrieverCreator.retrievers registry keys, NOT the legacy classic_rag label which never matched the registry. - ### `RETRIEVAL_MAX_PARALLEL_SOURCES` Type `int`, default `4`, must be >= 1. @@ -1015,18 +1009,24 @@ Path-style addressing (required by most non-AWS services). Type `str`, default unset. +**Deprecated.** Set S3_REGION instead; the SAGEMAKER_* fallback will be removed. + Legacy AWS region from the retired SageMaker provider; deprecated fallback for S3_REGION. ### `SAGEMAKER_ACCESS_KEY` Type `str`, default unset. +**Deprecated.** Set S3_ACCESS_KEY_ID instead; the SAGEMAKER_* fallback will be removed. + Legacy AWS access key from the retired SageMaker provider; deprecated fallback for S3_ACCESS_KEY_ID. ### `SAGEMAKER_SECRET_KEY` Type `str`, default unset. +**Deprecated.** Set S3_SECRET_ACCESS_KEY instead; the SAGEMAKER_* fallback will be removed. + Legacy AWS secret key from the retired SageMaker provider; deprecated fallback for S3_SECRET_ACCESS_KEY. @@ -1271,12 +1271,6 @@ Type `str`, default `classic`. Default agent type for agentless chats. -### `DEFAULT_MAX_HISTORY` - -Type `int`, default `150`. - -Default number of history messages kept. - ### `DEFAULT_AGENT_LIMITS` Type `dict`, default `{"token_limit": 50000, "request_limit": 500}`. diff --git a/docs/content/Sources/Per-source-configuration.mdx b/docs/content/Sources/Per-source-configuration.mdx index 4c23a710..a66ea285 100644 --- a/docs/content/Sources/Per-source-configuration.mdx +++ b/docs/content/Sources/Per-source-configuration.mdx @@ -115,8 +115,6 @@ Requests are also bounded: `chunks` is clamped to 0–500, and `0` still means Keyword search for the **hybrid** retriever is currently implemented only for the **pgvector** vector store. On other stores (FAISS, Qdrant, Milvus, etc.) the keyword half returns nothing, so `hybrid` quietly behaves like `classic` (vector-only). -Operators can restrict which retrievers are usable instance-wide with the `RETRIEVERS_ENABLED` setting; a per-source `retriever` value must be within that allow-list. - ### Exposure: prefetch vs. agentic tool `exposure` controls *how* a source's content is delivered to the model: diff --git a/docsgpt/core/settings/agents.py b/docsgpt/core/settings/agents.py index 9b682786..b15a4d62 100644 --- a/docsgpt/core/settings/agents.py +++ b/docsgpt/core/settings/agents.py @@ -13,7 +13,6 @@ class AgentSettings(SettingsGroup): """What an agent may do per turn and how its context is kept within budget.""" AGENT_NAME: str = Field(default="classic", description="Default agent type for agentless chats.") - DEFAULT_MAX_HISTORY: int = Field(default=150, description="Default number of history messages kept.") DEFAULT_AGENT_LIMITS: dict[str, int] = Field( default={"token_limit": 50000, "request_limit": 500}, description="Per-agent default quotas: tokens and requests.", diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py index a0bf87f9..8c716646 100644 --- a/docsgpt/core/settings/retrieval.py +++ b/docsgpt/core/settings/retrieval.py @@ -15,13 +15,6 @@ class RetrievalSettings(SettingsGroup): VECTOR_STORE: Literal["faiss", "elasticsearch", "mongodb", "qdrant", "milvus", "pgvector"] = Field( default="faiss", description="Vector store backend." ) - RETRIEVERS_ENABLED: list[str] = Field( - default=["classic", "default"], - description=( - "Retriever keys an agent may use; must match RetrieverCreator.retrievers registry keys, NOT the " - "legacy classic_rag label which never matched the registry." - ), - ) RETRIEVAL_MAX_PARALLEL_SOURCES: int = Field( default=4, ge=1, diff --git a/docsgpt/core/settings/storage.py b/docsgpt/core/settings/storage.py index 2fe75347..0212c050 100644 --- a/docsgpt/core/settings/storage.py +++ b/docsgpt/core/settings/storage.py @@ -34,14 +34,17 @@ class StorageSettings(SettingsGroup): # Legacy AWS credentials from the retired SageMaker provider. SAGEMAKER_REGION: Optional[str] = Field( default=None, + deprecated="Set S3_REGION instead; the SAGEMAKER_* fallback will be removed.", description="Legacy AWS region from the retired SageMaker provider; deprecated fallback for S3_REGION.", ) SAGEMAKER_ACCESS_KEY: Optional[str] = Field( default=None, + deprecated="Set S3_ACCESS_KEY_ID instead; the SAGEMAKER_* fallback will be removed.", description="Legacy AWS access key from the retired SageMaker provider; deprecated fallback for S3_ACCESS_KEY_ID.", ) SAGEMAKER_SECRET_KEY: Optional[str] = Field( default=None, + deprecated="Set S3_SECRET_ACCESS_KEY instead; the SAGEMAKER_* fallback will be removed.", description=( "Legacy AWS secret key from the retired SageMaker provider; deprecated fallback for " "S3_SECRET_ACCESS_KEY." diff --git a/docsgpt/storage/s3.py b/docsgpt/storage/s3.py index 790b869d..38e20552 100644 --- a/docsgpt/storage/s3.py +++ b/docsgpt/storage/s3.py @@ -4,6 +4,7 @@ import io import logging import os import posixpath +import warnings from typing import BinaryIO, Callable, List, Optional, Tuple import boto3 @@ -30,9 +31,13 @@ class S3Storage(BaseStorage): secret_key = settings.S3_SECRET_ACCESS_KEY region = settings.S3_REGION - legacy_access = getattr(settings, "SAGEMAKER_ACCESS_KEY", None) - legacy_secret = getattr(settings, "SAGEMAKER_SECRET_KEY", None) - legacy_region = getattr(settings, "SAGEMAKER_REGION", None) + # The SAGEMAKER_* fields are marked deprecated on the model and warn on every read; + # this is the one sanctioned reader, and it raises its own operator-facing warning. + with warnings.catch_warnings(): + warnings.simplefilter("ignore", DeprecationWarning) + legacy_access = settings.SAGEMAKER_ACCESS_KEY + legacy_secret = settings.SAGEMAKER_SECRET_KEY + legacy_region = settings.SAGEMAKER_REGION used_legacy = ( (not access_key and legacy_access) diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py index 0757d8a4..df0e0e07 100644 --- a/tests/core/test_settings.py +++ b/tests/core/test_settings.py @@ -6,6 +6,7 @@ validators from every group applied) and that the generated reference page tracks the definitions. """ +import warnings from pathlib import Path import pytest @@ -33,10 +34,16 @@ SECRET_FIELDS = ( @pytest.mark.unit class TestComposition: def test_every_group_field_is_a_flat_settings_attribute(self): - for _, group in SETTINGS_GROUPS: - for name in group.model_fields: - assert name in Settings.model_fields, name - assert hasattr(settings, name), name + with warnings.catch_warnings(): + warnings.simplefilter("ignore", DeprecationWarning) # reading a deprecated field warns + for _, group in SETTINGS_GROUPS: + for name in group.model_fields: + assert name in Settings.model_fields, name + assert hasattr(settings, name), name + + def test_deprecated_fields_warn_on_read(self): + with pytest.warns(DeprecationWarning, match="S3_REGION"): + _ = Settings(_env_file=None).SAGEMAKER_REGION def test_no_field_is_defined_in_two_groups(self): owners: dict[str, str] = {} From f882ef49a71c01eea6f30e6011f133d7deaf30c0 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:14:34 +0100 Subject: [PATCH 056/130] refactor: read settings directly instead of getattr with a second default About 85 call sites read a setting as getattr(settings, "NAME", fallback), each carrying its own copy of the default. Every one of those names is a field with a default on the model, so the fallback could never apply to the real settings object; it only masked drift. Two had drifted: - OPENAI_PROMPT_CACHE_KEY defaults to True on the model but the reader fell back to False, and two test stubs relied on that. - SharePoint's MICROSOFT_AUTHORITY fallback to https://login.microsoftonline.com/ never fired, because the attribute always exists (as None), so MSAL got authority=None. The connector now derives the tenant authority when the setting is unset, as its test always assumed. Four places read EMBEDDINGS_KEY straight from os.environ, skipping the "None"/"" normalisation the model applies; they read the setting now. Test stubs that replaced a module's settings with a SimpleNamespace list every setting the code under test reads. --- docs/content/Deploying/Settings-Reference.mdx | 2 +- docsgpt/agents/base.py | 4 +-- docsgpt/agents/tool_executor.py | 2 +- docsgpt/agents/tools/artifact_generator.py | 2 +- docsgpt/agents/tools/attachment_bridge.py | 2 +- docsgpt/agents/tools/code_executor.py | 8 +++--- docsgpt/agents/tools/mcp_tool.py | 4 +-- docsgpt/agents/tools/read_document.py | 2 +- docsgpt/agents/workflow_agent.py | 2 +- docsgpt/agents/workflows/workflow_engine.py | 14 +++++----- .../answer/services/compression/service.py | 2 +- docsgpt/api/async_sse.py | 3 +-- docsgpt/api/user/artifacts/download.py | 2 +- docsgpt/api/user/base.py | 2 +- docsgpt/api/user/tasks.py | 6 ++--- docsgpt/core/model_registry.py | 2 +- docsgpt/core/settings/connectors.py | 3 ++- docsgpt/devices/broker.py | 6 ++--- docsgpt/graphrag/store.py | 6 ++--- docsgpt/guardrails/guardrail_creator.py | 2 +- docsgpt/guardrails/runtime.py | 8 +++--- docsgpt/llm/handlers/base.py | 2 +- docsgpt/llm/openai.py | 6 ++--- docsgpt/parser/connectors/share_point/auth.py | 2 +- docsgpt/parser/document_reader.py | 12 ++++----- docsgpt/parser/embedding_pipeline.py | 8 +++--- docsgpt/parser/file/bulk.py | 2 +- docsgpt/parser/file/docling_parser.py | 6 +---- docsgpt/parser/file/ocr_parser.py | 12 ++++----- docsgpt/parser/remote/github_loader.py | 4 +-- docsgpt/parser/tokenization.py | 2 +- docsgpt/retriever/dispatcher.py | 2 +- docsgpt/sandbox/artifacts_capture.py | 6 ++--- docsgpt/sandbox/sandbox_creator.py | 2 +- docsgpt/scripts/reembed.py | 4 +-- docsgpt/storage/db/bootstrap.py | 8 +++--- docsgpt/storage/storage_creator.py | 2 +- docsgpt/utils.py | 4 +-- docsgpt/vectorstore/base.py | 2 +- docsgpt/vectorstore/embeddings_delegated.py | 4 +-- docsgpt/vectorstore/embeddings_local.py | 10 +++---- docsgpt/vectorstore/pgconn.py | 2 +- docsgpt/vectorstore/pgvector.py | 4 +-- tests/llm/test_openai_responses.py | 3 +++ tests/llm/test_responses_chain_budget.py | 27 ++++++++++--------- .../connectors/test_share_point_auth.py | 4 +-- 46 files changed, 111 insertions(+), 113 deletions(-) diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index efc94fc5..02876714 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -1080,7 +1080,7 @@ Azure AD tenant id, or 'common' for multi-tenant. Type `str`, default unset. -Authority URL override, e.g. "https://login.microsoftonline.com/\{tenant_id\}". +Authority URL override; unset derives https://login.microsoftonline.com/<MICROSOFT_TENANT_ID>. ### `CONFLUENCE_CLIENT_ID` diff --git a/docsgpt/agents/base.py b/docsgpt/agents/base.py index 71caaa34..d731a415 100644 --- a/docsgpt/agents/base.py +++ b/docsgpt/agents/base.py @@ -382,7 +382,7 @@ class BaseAgent(ABC): when the conversation was compressed after that turn was produced — the compressed local history is the context then, not the server's. """ - if not getattr(settings, "OPENAI_RESPONSES_CHAIN_ACROSS_TURNS", True): + if not settings.OPENAI_RESPONSES_CHAIN_ACROSS_TURNS: return None if not self.chat_history: return None @@ -422,7 +422,7 @@ class BaseAgent(ABC): # No provider-reported usage on the previous turn (older rows, # estimate-only providers): nothing to bound against. return meta["response_id"] - budget = getattr(settings, "OPENAI_RESPONSES_CHAIN_BUDGET_TOKENS", None) + budget = settings.OPENAI_RESPONSES_CHAIN_BUDGET_TOKENS if not budget: from docsgpt.core.model_utils import get_token_limit diff --git a/docsgpt/agents/tool_executor.py b/docsgpt/agents/tool_executor.py index edec516f..9952e3d2 100644 --- a/docsgpt/agents/tool_executor.py +++ b/docsgpt/agents/tool_executor.py @@ -53,7 +53,7 @@ def _dedupable_tool_names() -> frozenset: """ from docsgpt.core.settings import settings - return frozenset(BUILTIN_AGENT_TOOLS) | frozenset(getattr(settings, "DEFAULT_CHAT_TOOLS", None) or []) + return frozenset(BUILTIN_AGENT_TOOLS) | frozenset(settings.DEFAULT_CHAT_TOOLS or []) def _requires_approval(tool: Dict, action: Dict) -> bool: diff --git a/docsgpt/agents/tools/artifact_generator.py b/docsgpt/agents/tools/artifact_generator.py index 1f3c7007..5c2da920 100644 --- a/docsgpt/agents/tools/artifact_generator.py +++ b/docsgpt/agents/tools/artifact_generator.py @@ -829,7 +829,7 @@ class ArtifactGeneratorTool(Tool): spec_path = f"{token_dir}/spec.json" out_path = f"{token_dir}/out.{_KIND_INFO[kind]['ext']}" program = _RENDERERS[kind].format(spec_path=spec_path, out_path=out_path) - timeout = float(getattr(settings, "SANDBOX_EXEC_TIMEOUT", 60)) + timeout = float(settings.SANDBOX_EXEC_TIMEOUT) manager = SandboxCreator.get_manager() try: diff --git a/docsgpt/agents/tools/attachment_bridge.py b/docsgpt/agents/tools/attachment_bridge.py index 7e1400d1..23c56bfb 100644 --- a/docsgpt/agents/tools/attachment_bridge.py +++ b/docsgpt/agents/tools/attachment_bridge.py @@ -113,7 +113,7 @@ def bridge_attachment( # Reject oversize attachments BEFORE buffering them: the authoritative ``size`` # column lets us avoid pulling a multi-hundred-MB file fully into worker memory, # and the bounded read below backstops a missing/lying ``size``. - max_bytes = int(getattr(settings, "ARTIFACT_MAX_BYTES", 0) or 0) + max_bytes = int(settings.ARTIFACT_MAX_BYTES or 0) declared_size = attachment.get("size") if max_bytes and isinstance(declared_size, (int, float)) and declared_size > max_bytes: raise AttachmentBridgeError( diff --git a/docsgpt/agents/tools/code_executor.py b/docsgpt/agents/tools/code_executor.py index 079d4a78..697411a7 100644 --- a/docsgpt/agents/tools/code_executor.py +++ b/docsgpt/agents/tools/code_executor.py @@ -87,9 +87,9 @@ class CodeExecutorTool(Tool): baked in. Keep the package lists in sync with deployment/sandbox/Dockerfile (jupyter) and scripts/build_daytona_snapshot.py (daytona snapshot). """ - backend = str(getattr(settings, "SANDBOX_BACKEND", "jupyter") or "jupyter").lower() + backend = str(settings.SANDBOX_BACKEND or "jupyter").lower() if backend == "daytona": - if getattr(settings, "DAYTONA_SNAPSHOT", None): + if settings.DAYTONA_SNAPSHOT: return ( "Preinstalled beyond the stdlib: python-pptx, python-docx, openpyxl, " "reportlab, lxml, pillow. pip install anything else from within the code " @@ -330,7 +330,7 @@ class CodeExecutorTool(Tool): # Reject an oversize input BEFORE buffering it: the declared ``size`` # avoids pulling a huge file into worker memory, and the bounded read # below backstops a missing/lying size column. - max_bytes = int(getattr(settings, "SANDBOX_MAX_INPUT_BYTES", 0) or 0) + max_bytes = int(settings.SANDBOX_MAX_INPUT_BYTES or 0) declared_size = version.get("size") if max_bytes and isinstance(declared_size, (int, float)) and declared_size > max_bytes: return {"error": f"input artifact {artifact_id} exceeds the {max_bytes}-byte sandbox input limit."} @@ -484,7 +484,7 @@ class CodeExecutorTool(Tool): @staticmethod def _exec_timeout() -> float: """Return the fixed per-run wall-clock cap (SANDBOX_EXEC_TIMEOUT; not caller-adjustable).""" - return float(getattr(settings, "SANDBOX_EXEC_TIMEOUT", 60)) + return float(settings.SANDBOX_EXEC_TIMEOUT) @staticmethod def _is_timeout(result: ExecResult) -> bool: diff --git a/docsgpt/agents/tools/mcp_tool.py b/docsgpt/agents/tools/mcp_tool.py index e8b4112b..061f4c6d 100644 --- a/docsgpt/agents/tools/mcp_tool.py +++ b/docsgpt/agents/tools/mcp_tool.py @@ -108,11 +108,11 @@ class MCPTool(Tool): if configured_redirect_uri: return configured_redirect_uri.rstrip("/") - explicit = getattr(settings, "MCP_OAUTH_REDIRECT_URI", None) + explicit = settings.MCP_OAUTH_REDIRECT_URI if explicit: return explicit.rstrip("/") - connector_base = getattr(settings, "CONNECTOR_REDIRECT_BASE_URI", None) + connector_base = settings.CONNECTOR_REDIRECT_BASE_URI if connector_base: parsed = urlparse(connector_base) if parsed.scheme and parsed.netloc: diff --git a/docsgpt/agents/tools/read_document.py b/docsgpt/agents/tools/read_document.py index 6e488937..ab9d2294 100644 --- a/docsgpt/agents/tools/read_document.py +++ b/docsgpt/agents/tools/read_document.py @@ -255,7 +255,7 @@ class ReadDocumentTool(Tool): # The task's per-call time limits are raised to match the awaited window: bound to # the base timeout at import, the worker would otherwise self-terminate a large # parse long before this await gives up. - queue = getattr(settings, "DOCUMENT_PARSE_QUEUE", "parsing") + queue = settings.DOCUMENT_PARSE_QUEUE try: async_result = parse_document.apply_async( args=[artifact_id, parent, self.user_id, options], diff --git a/docsgpt/agents/workflow_agent.py b/docsgpt/agents/workflow_agent.py index 6c1f71ff..2c08bc0c 100644 --- a/docsgpt/agents/workflow_agent.py +++ b/docsgpt/agents/workflow_agent.py @@ -331,7 +331,7 @@ class WorkflowAgent(BaseAgent): from docsgpt.storage.storage_creator import StorageCreator storage = StorageCreator.get_storage() - max_bytes = int(getattr(settings, "ARTIFACT_MAX_BYTES", 0) or 0) + max_bytes = int(settings.ARTIFACT_MAX_BYTES or 0) dropped: List[str] = [] if len(self.attachments) > _MAX_INPUT_DOCUMENTS: over = len(self.attachments) - _MAX_INPUT_DOCUMENTS diff --git a/docsgpt/agents/workflows/workflow_engine.py b/docsgpt/agents/workflows/workflow_engine.py index b1189511..54fa6e67 100644 --- a/docsgpt/agents/workflows/workflow_engine.py +++ b/docsgpt/agents/workflows/workflow_engine.py @@ -639,7 +639,7 @@ class WorkflowEngine: raw_ids = self._resolve_input_artifact_ids(inputs) if not raw_ids: return loaded - max_bytes = int(getattr(settings, "SANDBOX_MAX_INPUT_BYTES", 0) or 0) + max_bytes = int(settings.SANDBOX_MAX_INPUT_BYTES or 0) storage = StorageCreator.get_storage() # Two inputs whose current versions share a filename would clobber each other at the # same ``inputs/{name}`` path; track used paths and disambiguate deterministically. @@ -749,15 +749,15 @@ class WorkflowEngine: supported = set(supported_types) supports_images = any(t.startswith("image/") for t in supported) - max_files = int(getattr(settings, "WORKFLOW_NODE_NATIVE_MAX_FILES", 5)) - extract_max = int(getattr(settings, "WORKFLOW_NODE_EXTRACT_MAX_FILES", 5)) + max_files = int(settings.WORKFLOW_NODE_NATIVE_MAX_FILES) + extract_max = int(settings.WORKFLOW_NODE_EXTRACT_MAX_FILES) # One wall clock for every blocking parse this node issues. The cap # above bounds how MANY parses run; this bounds how LONG they take in # total, so N documents cannot serialize N size-scaled windows. parse_deadline = time.monotonic() + float( - getattr(settings, "WORKFLOW_NODE_EXTRACT_BUDGET_SECONDS", 900) + settings.WORKFLOW_NODE_EXTRACT_BUDGET_SECONDS ) - max_bytes = int(getattr(settings, "SANDBOX_MAX_INPUT_BYTES", 25 * 1024 * 1024)) + max_bytes = int(settings.SANDBOX_MAX_INPUT_BYTES) # One read-only connection for the whole batch; the resolved-version # rows are collected, then storage reads happen outside the DB context. @@ -976,7 +976,7 @@ class WorkflowEngine: if not user_id: return None options = {"output": "markdown", "include_tables": False, "persist": False} - queue = getattr(settings, "DOCUMENT_PARSE_QUEUE", "parsing") + queue = settings.DOCUMENT_PARSE_QUEUE # OCR cost scales with pages, so the window grows with the document's size # (floored at DOCUMENT_PARSE_TIMEOUT); the task's per-call time limits are # raised to match, else the worker would self-terminate mid-parse. @@ -1084,7 +1084,7 @@ class WorkflowEngine: """Return the stricter of the node's requested timeout and the sandbox cap.""" from docsgpt.core.settings import settings - cap = float(getattr(settings, "SANDBOX_EXEC_TIMEOUT", 60)) + cap = float(settings.SANDBOX_EXEC_TIMEOUT) if requested is None: return cap try: diff --git a/docsgpt/api/answer/services/compression/service.py b/docsgpt/api/answer/services/compression/service.py index eb47f5d0..6b6b14a5 100644 --- a/docsgpt/api/answer/services/compression/service.py +++ b/docsgpt/api/answer/services/compression/service.py @@ -367,7 +367,7 @@ class CompressionService: never mutated. """ max_tokens = int( - getattr(settings, "COMPRESSION_RECENT_FIELD_MAX_TOKENS", 8000) or 0 + settings.COMPRESSION_RECENT_FIELD_MAX_TOKENS or 0 ) if max_tokens <= 0: return queries diff --git a/docsgpt/api/async_sse.py b/docsgpt/api/async_sse.py index d3fa3437..a5064e2c 100644 --- a/docsgpt/api/async_sse.py +++ b/docsgpt/api/async_sse.py @@ -33,7 +33,6 @@ from docsgpt.streaming.async_event_replay import ( ) from docsgpt.streaming.async_redis import get_async_redis_instance from docsgpt.streaming.event_replay import ( - DEFAULT_KEEPALIVE_SECONDS, DEFAULT_POLL_TIMEOUT_SECONDS, ) from docsgpt.streaming.sse_leases import StreamCapExceeded, acquire_stream_lease @@ -127,7 +126,7 @@ async def stream_message_events(request: Request) -> Response: ) last_event_id = _normalise_last_event_id(raw_cursor) keepalive_seconds = float( - getattr(settings, "SSE_KEEPALIVE_SECONDS", DEFAULT_KEEPALIVE_SECONDS) + settings.SSE_KEEPALIVE_SECONDS ) logger.info( diff --git a/docsgpt/api/user/artifacts/download.py b/docsgpt/api/user/artifacts/download.py index a57ae727..3e23895d 100644 --- a/docsgpt/api/user/artifacts/download.py +++ b/docsgpt/api/user/artifacts/download.py @@ -154,7 +154,7 @@ async def download_artifact(request: Request) -> Response: # URL. If the active backend can't mint one, that's a config error: # surface a 500 rather than silently proxying bytes from a backend # the operator expected to be off the hot path. - if getattr(settings, "URL_STRATEGY", "backend") == "s3": + if settings.URL_STRATEGY == "s3": try: url = await anyio.to_thread.run_sync( partial(storage.generate_presigned_url, storage_path, expires_in=_PRESIGNED_URL_TTL) diff --git a/docsgpt/api/user/base.py b/docsgpt/api/user/base.py index f6cd6f58..9579a73f 100644 --- a/docsgpt/api/user/base.py +++ b/docsgpt/api/user/base.py @@ -235,7 +235,7 @@ def get_vector_store(source_id): store = VectorCreator.create_vectorstore( settings.VECTOR_STORE, source_id=source_id, - embeddings_key=os.getenv("EMBEDDINGS_KEY"), + embeddings_key=settings.EMBEDDINGS_KEY, ) return store diff --git a/docsgpt/api/user/tasks.py b/docsgpt/api/user/tasks.py index d4307b33..82d247e7 100644 --- a/docsgpt/api/user/tasks.py +++ b/docsgpt/api/user/tasks.py @@ -346,9 +346,9 @@ def parse_timeout_for_size(size_bytes: Optional[int]) -> float: """ from docsgpt.core.settings import settings - base = float(getattr(settings, "DOCUMENT_PARSE_TIMEOUT", 120) or 120) - per_mib = float(getattr(settings, "DOCUMENT_PARSE_TIMEOUT_PER_MB", 0) or 0) - ceiling = float(getattr(settings, "DOCUMENT_PARSE_TIMEOUT_MAX", base) or base) + base = float(settings.DOCUMENT_PARSE_TIMEOUT or 120) + per_mib = float(settings.DOCUMENT_PARSE_TIMEOUT_PER_MB or 0) + ceiling = float(settings.DOCUMENT_PARSE_TIMEOUT_MAX or base) size = float(size_bytes) if isinstance(size_bytes, (int, float)) else 0.0 scaled = base + per_mib * max(size, 0.0) / (1024 * 1024) return min(ceiling, max(base, scaled)) diff --git a/docsgpt/core/model_registry.py b/docsgpt/core/model_registry.py index 07a30a02..cc71e206 100644 --- a/docsgpt/core/model_registry.py +++ b/docsgpt/core/model_registry.py @@ -140,7 +140,7 @@ class ModelRegistry: from docsgpt.llm.providers import ALL_PROVIDERS directories = [BUILTIN_MODELS_DIR] - operator_dir = getattr(settings, "MODELS_CONFIG_DIR", None) + operator_dir = settings.MODELS_CONFIG_DIR if operator_dir: op_path = Path(operator_dir) if not op_path.exists(): diff --git a/docsgpt/core/settings/connectors.py b/docsgpt/core/settings/connectors.py index 676ffe64..b4199300 100644 --- a/docsgpt/core/settings/connectors.py +++ b/docsgpt/core/settings/connectors.py @@ -35,7 +35,8 @@ class ConnectorSettings(SettingsGroup): default="common", description="Azure AD tenant id, or 'common' for multi-tenant." ) MICROSOFT_AUTHORITY: Optional[str] = Field( - default=None, description='Authority URL override, e.g. "https://login.microsoftonline.com/{tenant_id}".' + default=None, + description="Authority URL override; unset derives https://login.microsoftonline.com/.", ) # Confluence Cloud integration. diff --git a/docsgpt/devices/broker.py b/docsgpt/devices/broker.py index 09d55360..79ac2287 100644 --- a/docsgpt/devices/broker.py +++ b/docsgpt/devices/broker.py @@ -631,15 +631,15 @@ class DeviceBroker: @staticmethod def _inv_ttl() -> int: - return int(getattr(settings, "REMOTE_DEVICE_INVOCATION_TTL_SECONDS", 900)) + return int(settings.REMOTE_DEVICE_INVOCATION_TTL_SECONDS) @staticmethod def _cmd_ttl() -> int: - return int(getattr(settings, "REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS", 900)) + return int(settings.REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS) @staticmethod def _out_maxlen() -> int: - return int(getattr(settings, "REMOTE_DEVICE_OUTPUT_STREAM_MAXLEN", 10_000)) + return int(settings.REMOTE_DEVICE_OUTPUT_STREAM_MAXLEN) def _to_int(value: Optional[str]) -> Optional[int]: diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index e9ee39b9..818004c0 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -77,11 +77,9 @@ class GraphStore: """Stores and queries a per-source knowledge graph in the pgvector DB.""" def __init__(self, connection_string: Optional[str] = None): - self._connection_string = connection_string or getattr( - settings, "PGVECTOR_CONNECTION_STRING", None - ) + self._connection_string = connection_string or settings.PGVECTOR_CONNECTION_STRING - if not self._connection_string and getattr(settings, "POSTGRES_URI", None): + if not self._connection_string and settings.POSTGRES_URI: from docsgpt.core.db_uri import normalize_pgvector_connection_string self._connection_string = normalize_pgvector_connection_string( diff --git a/docsgpt/guardrails/guardrail_creator.py b/docsgpt/guardrails/guardrail_creator.py index d891e065..f3357f67 100644 --- a/docsgpt/guardrails/guardrail_creator.py +++ b/docsgpt/guardrails/guardrail_creator.py @@ -52,7 +52,7 @@ class GuardrailCreator: does not require an operator to also edit their env. """ cls._ensure_builtin() - allowlist = getattr(settings, "GUARDRAILS_CHECKS_ENABLED", None) or [] + allowlist = settings.GUARDRAILS_CHECKS_ENABLED or [] if not allowlist: return sorted(cls.checks) return sorted(k for k in cls.checks if k in set(allowlist)) diff --git a/docsgpt/guardrails/runtime.py b/docsgpt/guardrails/runtime.py index 16c20638..3c73a53b 100644 --- a/docsgpt/guardrails/runtime.py +++ b/docsgpt/guardrails/runtime.py @@ -38,7 +38,7 @@ def _merge_mode(agent_mode: str, floor_mode: str) -> str: def instance_floor() -> Optional[GuardrailsConfig]: """The operator-set minimum, or None when unset/invalid.""" - raw = getattr(settings, "GUARDRAILS_FLOOR", None) + raw = settings.GUARDRAILS_FLOOR if not raw: return None try: @@ -104,7 +104,7 @@ def floor_keys() -> set: def resolve_config(raw_agent_config: Optional[dict]) -> GuardrailsConfig: """Parse ``agents.config`` and apply the instance floor.""" - if not getattr(settings, "GUARDRAILS_ENABLED", True): + if not settings.GUARDRAILS_ENABLED: return GuardrailsConfig() agent = AgentConfig.parse(raw_agent_config).guardrails return merge_floor(agent, instance_floor()) @@ -123,7 +123,7 @@ def _judge_factory(agent): decoded_token=agent.decoded_token, model_id=( model_override - or getattr(settings, "GUARDRAILS_JUDGE_MODEL", None) + or settings.GUARDRAILS_JUDGE_MODEL or agent.upstream_model_id ), agent_id=agent.agent_id, @@ -169,7 +169,7 @@ class GuardrailRecorder: self._seen: set = set() def __call__(self, decision: StageDecision) -> None: - store_text = bool(getattr(settings, "GUARDRAILS_STORE_SCANNED_TEXT", False)) + store_text = bool(settings.GUARDRAILS_STORE_SCANNED_TEXT) for verdict in decision.verdicts: if not verdict.outcome.triggered and verdict.outcome.evaluated: continue diff --git a/docsgpt/llm/handlers/base.py b/docsgpt/llm/handlers/base.py index 0cd4c84c..039d7007 100644 --- a/docsgpt/llm/handlers/base.py +++ b/docsgpt/llm/handlers/base.py @@ -35,7 +35,7 @@ def _bound_tool_response_for_llm(tool_response: Any) -> Any: from docsgpt.core.settings import settings from docsgpt.utils import num_tokens_from_string - max_tokens = int(getattr(settings, "TOOL_RESULT_MAX_TOKENS", 20000) or 0) + max_tokens = int(settings.TOOL_RESULT_MAX_TOKENS or 0) if max_tokens <= 0: return tool_response text = tool_response if isinstance(tool_response, str) else str(tool_response) diff --git a/docsgpt/llm/openai.py b/docsgpt/llm/openai.py index a901f0a1..ad2f1f4b 100644 --- a/docsgpt/llm/openai.py +++ b/docsgpt/llm/openai.py @@ -1307,14 +1307,14 @@ class OpenAILLM(BaseLLM): params["include"] = ["reasoning.encrypted_content"] # Backstop against a chain that outgrows the model's native window: # the provider drops the oldest input items instead of failing. - if getattr(settings, "OPENAI_RESPONSES_TRUNCATION_AUTO", False): + if settings.OPENAI_RESPONSES_TRUNCATION_AUTO: params["truncation"] = "auto" # Prompt-cache hints. The key pins a conversation to one cache shard; # retention asks for the extended tier where the deployment offers it. cache_key = getattr(self, "_prompt_cache_key", None) - if cache_key and getattr(settings, "OPENAI_PROMPT_CACHE_KEY", False): + if cache_key and settings.OPENAI_PROMPT_CACHE_KEY: params["prompt_cache_key"] = str(cache_key) - retention = getattr(settings, "OPENAI_PROMPT_CACHE_RETENTION", None) + retention = settings.OPENAI_PROMPT_CACHE_RETENTION if retention: params["prompt_cache_retention"] = retention return params diff --git a/docsgpt/parser/connectors/share_point/auth.py b/docsgpt/parser/connectors/share_point/auth.py index 9a4264f6..ec006740 100644 --- a/docsgpt/parser/connectors/share_point/auth.py +++ b/docsgpt/parser/connectors/share_point/auth.py @@ -41,7 +41,7 @@ class SharePointAuth(BaseConnectorAuth): self.redirect_uri = settings.CONNECTOR_REDIRECT_BASE_URI self.tenant_id = settings.MICROSOFT_TENANT_ID - self.authority = getattr(settings, "MICROSOFT_AUTHORITY", f"https://login.microsoftonline.com/{self.tenant_id}") + self.authority = settings.MICROSOFT_AUTHORITY or f"https://login.microsoftonline.com/{self.tenant_id}" self.auth_app = ConfidentialClientApplication( client_id=self.client_id, diff --git a/docsgpt/parser/document_reader.py b/docsgpt/parser/document_reader.py index 3a031d19..ed3446e7 100644 --- a/docsgpt/parser/document_reader.py +++ b/docsgpt/parser/document_reader.py @@ -97,10 +97,10 @@ def bound_parse_payload(payload: Dict[str, Any], max_chars: Optional[int] = None def _max_input_bytes() -> int: """Return the size cap for a parsed document (its own setting, else the sandbox cap).""" - explicit = int(getattr(settings, "DOCUMENT_PARSE_MAX_BYTES", 0) or 0) + explicit = int(settings.DOCUMENT_PARSE_MAX_BYTES or 0) if explicit > 0: return explicit - return int(getattr(settings, "SANDBOX_MAX_INPUT_BYTES", 25 * 1024 * 1024)) + return int(settings.SANDBOX_MAX_INPUT_BYTES) # Every zip-packaged format a parser map can route: OOXML and its macro/ @@ -127,8 +127,8 @@ def _zip_bomb_reason(source: Union[bytes, str, Path], suffix: str) -> Optional[s """ if suffix not in _ZIP_CONTAINER_EXTENSIONS: return None - max_entries = int(getattr(settings, "DOCUMENT_MAX_ARCHIVE_ENTRIES", 10000)) - cap = int(getattr(settings, "DOCUMENT_MAX_DECOMPRESSED_BYTES", 300 * 1024 * 1024)) + max_entries = int(settings.DOCUMENT_MAX_ARCHIVE_ENTRIES) + cap = int(settings.DOCUMENT_MAX_DECOMPRESSED_BYTES) opened = io.BytesIO(source) if isinstance(source, bytes) else source try: with zipfile.ZipFile(opened) as zf: @@ -169,14 +169,14 @@ def _resolve_ocr_enabled(ocr: str) -> bool: return True if ocr == "off": return False - return bool(getattr(settings, "OCR_ENABLED", False)) + return bool(settings.OCR_ENABLED) def _effective_engine(engine: str) -> str: """Resolve ``auto`` to the server's ``DOC_PARSER_ENGINE``; other values pass through.""" if engine != "auto": return engine - configured = getattr(settings, "DOC_PARSER_ENGINE", None) or "anydoc" + configured = settings.DOC_PARSER_ENGINE or "anydoc" return str(configured).strip().lower() diff --git a/docsgpt/parser/embedding_pipeline.py b/docsgpt/parser/embedding_pipeline.py index a0dbd02b..893c77f6 100755 --- a/docsgpt/parser/embedding_pipeline.py +++ b/docsgpt/parser/embedding_pipeline.py @@ -53,7 +53,7 @@ def _resolve_batch_size() -> int: Returns: Chunks per embed request, always >= 1. """ - raw = getattr(settings, "EMBEDDINGS_BATCH_SIZE", None) + raw = settings.EMBEDDINGS_BATCH_SIZE # Explicit type check rather than a bare ``int(raw)``: ``int(MagicMock())`` # succeeds and yields 1, which would silently drop ingest back to the # per-chunk behaviour this batching replaces. @@ -326,7 +326,7 @@ def embed_and_store_documents( store = VectorCreator.create_vectorstore( settings.VECTOR_STORE, source_id=source_id, - embeddings_key=os.getenv("EMBEDDINGS_KEY"), + embeddings_key=settings.EMBEDDINGS_KEY, ) loop_start = resume_index else: @@ -336,7 +336,7 @@ def embed_and_store_documents( settings.VECTOR_STORE, docs_init=[docs[0]], source_id=source_id, - embeddings_key=os.getenv("EMBEDDINGS_KEY"), + embeddings_key=settings.EMBEDDINGS_KEY, ) # Record the seeded chunk so single-doc ingests don't fail # ``assert_index_complete`` — the loop never runs for @@ -351,7 +351,7 @@ def embed_and_store_documents( store = VectorCreator.create_vectorstore( settings.VECTOR_STORE, source_id=source_id, - embeddings_key=os.getenv("EMBEDDINGS_KEY"), + embeddings_key=settings.EMBEDDINGS_KEY, ) # Only wipe the index on a fresh run — a resume must keep the # chunks that earlier attempts already embedded. diff --git a/docsgpt/parser/file/bulk.py b/docsgpt/parser/file/bulk.py index 4828f893..10ccfb1f 100644 --- a/docsgpt/parser/file/bulk.py +++ b/docsgpt/parser/file/bulk.py @@ -344,7 +344,7 @@ def get_default_file_extractor( """ if ocr_enabled is None: ocr_enabled = settings.OCR_ENABLED - selected = (engine or getattr(settings, "DOC_PARSER_ENGINE", None) or "anydoc") + selected = (engine or settings.DOC_PARSER_ENGINE or "anydoc") selected = str(selected).strip().lower() if selected == "docling": return _docling_file_extractor(ocr_enabled, pdf_text_fast_path) diff --git a/docsgpt/parser/file/docling_parser.py b/docsgpt/parser/file/docling_parser.py index e1efd546..577f43d3 100644 --- a/docsgpt/parser/file/docling_parser.py +++ b/docsgpt/parser/file/docling_parser.py @@ -335,11 +335,7 @@ def _ocr_min_chars_per_page() -> int: try: return int( - getattr( - settings, - "OCR_MIN_CHARS_PER_PAGE", - _DEFAULT_OCR_MIN_CHARS_PER_PAGE, - ) + settings.OCR_MIN_CHARS_PER_PAGE ) except (TypeError, ValueError): return _DEFAULT_OCR_MIN_CHARS_PER_PAGE diff --git a/docsgpt/parser/file/ocr_parser.py b/docsgpt/parser/file/ocr_parser.py index 28b67a7c..3a72f049 100644 --- a/docsgpt/parser/file/ocr_parser.py +++ b/docsgpt/parser/file/ocr_parser.py @@ -124,7 +124,7 @@ def resolve_ocr_backend(requested: Optional[str] = None) -> str: """ from docsgpt.core.settings import settings - backend = str(requested or getattr(settings, "OCR_BACKEND", None) or "auto").strip().lower() + backend = str(requested or settings.OCR_BACKEND or "auto").strip().lower() if backend not in VALID_OCR_BACKENDS: logger.warning(f"Unknown OCR_BACKEND {backend!r}; using auto") backend = "auto" @@ -154,7 +154,7 @@ def resolve_native_ocr_engine(requested: Optional[str] = None) -> str: """ from docsgpt.core.settings import settings - engine = str(requested or getattr(settings, "OCR_ENGINE", None) or "tesseract").strip().lower() + engine = str(requested or settings.OCR_ENGINE or "tesseract").strip().lower() if engine in NATIVE_OCR_ENGINES: return engine if engine in VALID_OCR_ENGINES: @@ -172,7 +172,7 @@ def ocr_min_chars_per_page() -> int: from docsgpt.core.settings import settings try: - return int(getattr(settings, "OCR_MIN_CHARS_PER_PAGE", _DEFAULT_MIN_CHARS_PER_PAGE)) + return int(settings.OCR_MIN_CHARS_PER_PAGE) except (TypeError, ValueError): return _DEFAULT_MIN_CHARS_PER_PAGE @@ -182,7 +182,7 @@ def render_dpi() -> int: from docsgpt.core.settings import settings try: - dpi = int(getattr(settings, "OCR_RENDER_DPI", _DEFAULT_RENDER_DPI)) + dpi = int(settings.OCR_RENDER_DPI) except (TypeError, ValueError): dpi = _DEFAULT_RENDER_DPI return max(_MIN_RENDER_DPI, min(_MAX_RENDER_DPI, dpi)) @@ -299,7 +299,7 @@ class TesseractEngine: else: from docsgpt.core.settings import settings - configured = str(getattr(settings, "OCR_LANGS", "") or "eng") + configured = str(settings.OCR_LANGS or "eng") langs = [lang.strip() for lang in configured.split("+") if lang.strip()] return "+".join(langs) or "eng" @@ -388,7 +388,7 @@ class DeepseekOcrEngine: self.url = url or settings.OCR_DEEPSEEK_URL self.model = model or settings.OCR_DEEPSEEK_MODEL - self.timeout = float(timeout if timeout is not None else getattr(settings, "OCR_DEEPSEEK_TIMEOUT", 300)) + self.timeout = float(timeout if timeout is not None else settings.OCR_DEEPSEEK_TIMEOUT) self.prompt = prompt self.max_tokens = max_tokens diff --git a/docsgpt/parser/remote/github_loader.py b/docsgpt/parser/remote/github_loader.py index dbf75d65..3ba4a55c 100644 --- a/docsgpt/parser/remote/github_loader.py +++ b/docsgpt/parser/remote/github_loader.py @@ -129,7 +129,7 @@ class GitHubLoader(BaseRemote): def _max_file_bytes(self) -> int: """Resolve the per-blob size cap; ``0`` disables it.""" - raw = getattr(settings, "GITHUB_INGEST_MAX_FILE_BYTES", None) + raw = settings.GITHUB_INGEST_MAX_FILE_BYTES if isinstance(raw, bool) or not isinstance(raw, (int, str)): return 1048576 try: @@ -139,7 +139,7 @@ class GitHubLoader(BaseRemote): def _max_workers(self) -> int: """Resolve the parallel-fetch width, clamped to a sane range.""" - raw = getattr(settings, "GITHUB_INGEST_MAX_WORKERS", None) + raw = settings.GITHUB_INGEST_MAX_WORKERS if isinstance(raw, bool) or not isinstance(raw, (int, str)): return 8 try: diff --git a/docsgpt/parser/tokenization.py b/docsgpt/parser/tokenization.py index 6391461d..0f45475f 100644 --- a/docsgpt/parser/tokenization.py +++ b/docsgpt/parser/tokenization.py @@ -267,7 +267,7 @@ def get_token_counter(embeddings_name: Optional[str] = None) -> TokenCounter: A :class:`HuggingFaceCounter` for a model whose tokenizer could be loaded, else a :class:`TiktokenCounter`. """ - name = embeddings_name or getattr(settings, "EMBEDDINGS_NAME", None) + name = embeddings_name or settings.EMBEDDINGS_NAME key = name or "__default__" with _cache_lock: if key in _cache: diff --git a/docsgpt/retriever/dispatcher.py b/docsgpt/retriever/dispatcher.py index d96b0037..a14fc02f 100644 --- a/docsgpt/retriever/dispatcher.py +++ b/docsgpt/retriever/dispatcher.py @@ -344,6 +344,6 @@ def build_dispatcher(create_classic: Callable[[], BaseRetriever], **kwargs): Returns: A ``Dispatcher`` or the legacy retriever from ``create_classic``. """ - if not getattr(settings, "PER_SOURCE_RETRIEVAL_ENABLED", True): + if not settings.PER_SOURCE_RETRIEVAL_ENABLED: return create_classic() return Dispatcher(**kwargs) diff --git a/docsgpt/sandbox/artifacts_capture.py b/docsgpt/sandbox/artifacts_capture.py index 9315f87d..697ca718 100644 --- a/docsgpt/sandbox/artifacts_capture.py +++ b/docsgpt/sandbox/artifacts_capture.py @@ -277,7 +277,7 @@ def _cleanup_orphan(storage: Any, saved_key: Optional[str]) -> None: def _check_single_artifact_size(size: int) -> None: """Reject a single artifact version whose byte size exceeds ``ARTIFACT_MAX_BYTES``.""" - max_bytes = int(getattr(settings, "ARTIFACT_MAX_BYTES", 0) or 0) + max_bytes = int(settings.ARTIFACT_MAX_BYTES or 0) if max_bytes > 0 and size > max_bytes: raise QuotaExceeded(f"artifact is too large: {size} bytes exceeds the {max_bytes}-byte per-file cap") @@ -292,10 +292,10 @@ def _enforce_user_quota(repo: ArtifactsRepository, user_id: str, added_bytes: in existing identity. """ _check_single_artifact_size(added_bytes) - max_count = int(getattr(settings, "ARTIFACT_MAX_COUNT_PER_USER", 0) or 0) + max_count = int(settings.ARTIFACT_MAX_COUNT_PER_USER or 0) if new_artifact and max_count > 0 and repo.count_for_user(user_id) >= max_count: raise QuotaExceeded(f"artifact count quota reached ({max_count}); delete artifacts to free space") - max_total = int(getattr(settings, "ARTIFACT_MAX_TOTAL_BYTES_PER_USER", 0) or 0) + max_total = int(settings.ARTIFACT_MAX_TOTAL_BYTES_PER_USER or 0) if max_total > 0 and repo.total_bytes_for_user(user_id) + added_bytes > max_total: raise QuotaExceeded(f"artifact storage quota reached ({max_total} bytes); delete artifacts to free space") diff --git a/docsgpt/sandbox/sandbox_creator.py b/docsgpt/sandbox/sandbox_creator.py index 4cf46b7e..0b34628d 100644 --- a/docsgpt/sandbox/sandbox_creator.py +++ b/docsgpt/sandbox/sandbox_creator.py @@ -64,7 +64,7 @@ class SandboxCreator: def get_manager(cls) -> SandboxManager: """Return the process-wide ``SandboxManager``, building it on first use.""" if cls._instance is None: - backend = cls.create_backend(getattr(settings, "SANDBOX_BACKEND", "jupyter")) + backend = cls.create_backend(settings.SANDBOX_BACKEND) cls._instance = SandboxManager( backend=backend, max_ttl=float(settings.SANDBOX_MAX_TTL), diff --git a/docsgpt/scripts/reembed.py b/docsgpt/scripts/reembed.py index 63c6e56e..eae8c71c 100644 --- a/docsgpt/scripts/reembed.py +++ b/docsgpt/scripts/reembed.py @@ -309,7 +309,7 @@ def reembed_pgvector(source_id: str, batch_size: int, dry_run: bool) -> Tuple[in # The graph seeds every traversal from its own vectors, so leaving them # in the old model's space is the same silent mismatch this script # exists to remove -- and at equal widths nothing would report it. - if getattr(settings, "GRAPHRAG_ENABLED", False): + if settings.GRAPHRAG_ENABLED: nodes = reembed_graph_nodes(store, conn, source_id, batch_size, dry_run) if nodes: logger.info( @@ -512,7 +512,7 @@ def main(argv: Optional[Sequence[str]] = None) -> int: # latency and a dependency on a worker running. Loading the model here also # means the script reports a real failure for a model it cannot load, # instead of timing out against an empty queue. - if getattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False): + if settings.EMBEDDINGS_DELEGATE_TO_WORKER: logger.info("Embedding in-process; worker delegation does not apply here.") settings.EMBEDDINGS_DELEGATE_TO_WORKER = False diff --git a/docsgpt/storage/db/bootstrap.py b/docsgpt/storage/db/bootstrap.py index 3ebf677e..fcfac3f6 100644 --- a/docsgpt/storage/db/bootstrap.py +++ b/docsgpt/storage/db/bootstrap.py @@ -96,7 +96,7 @@ def _release_boot_only_embeddings(log: logging.Logger) -> None: if settings.EMBEDDINGS_BASE_URL: return - if getattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False) is not True: + if settings.EMBEDDINGS_DELEGATE_TO_WORKER is not True: return import gc @@ -144,8 +144,8 @@ def ensure_vector_schema(*, logger: Optional[logging.Logger] = None) -> None: ) return - dsn = getattr(settings, "PGVECTOR_CONNECTION_STRING", None) - if not dsn and getattr(settings, "POSTGRES_URI", None): + dsn = settings.PGVECTOR_CONNECTION_STRING + if not dsn and settings.POSTGRES_URI: from docsgpt.core.db_uri import normalize_pgvector_connection_string dsn = normalize_pgvector_connection_string(settings.POSTGRES_URI) @@ -172,7 +172,7 @@ def ensure_vector_schema(*, logger: Optional[logging.Logger] = None) -> None: dim: Optional[int] = dimension_for(settings.EMBEDDINGS_NAME) - graph_enabled = bool(getattr(settings, "GRAPHRAG_ENABLED", False)) + graph_enabled = bool(settings.GRAPHRAG_ENABLED) started = time.monotonic() # A plain connection, never the store's pool: this can run pre-fork under # ``gunicorn --preload``, and an inherited pooled socket is a broken one. diff --git a/docsgpt/storage/storage_creator.py b/docsgpt/storage/storage_creator.py index 1c57db64..8604dbf1 100644 --- a/docsgpt/storage/storage_creator.py +++ b/docsgpt/storage/storage_creator.py @@ -18,7 +18,7 @@ class StorageCreator: @classmethod def get_storage(cls) -> BaseStorage: if cls._instance is None: - storage_type = getattr(settings, "STORAGE_TYPE", "local") + storage_type = settings.STORAGE_TYPE cls._instance = cls.create_storage(storage_type) return cls._instance diff --git a/docsgpt/utils.py b/docsgpt/utils.py index 233e3f7b..2b7496f5 100644 --- a/docsgpt/utils.py +++ b/docsgpt/utils.py @@ -347,7 +347,7 @@ def generate_agent_image_capability( agent_id: object, image_path: object, user_id: object ) -> str: """Create an HMAC capability for one agent's current internal image.""" - secret = getattr(settings, "JWT_SECRET_KEY", "") + secret = settings.JWT_SECRET_KEY if not isinstance(secret, str) or not secret: return "" try: @@ -389,7 +389,7 @@ def generate_image_url(image_path, agent_id=None, user_id=None): if not capability: return "" canonical_agent_id = str(uuid.UUID(str(agent_id))) - base_url = getattr(settings, "API_URL", "http://localhost:7091").rstrip("/") + base_url = settings.API_URL.rstrip("/") return f"{base_url}/api/images/{canonical_agent_id}/{capability}" diff --git a/docsgpt/vectorstore/base.py b/docsgpt/vectorstore/base.py index cad18833..d0c63c04 100644 --- a/docsgpt/vectorstore/base.py +++ b/docsgpt/vectorstore/base.py @@ -277,7 +277,7 @@ def _delegation_enabled() -> bool: with a ``MagicMock``, whose every attribute is a truthy object, and ``bool()`` on that would silently route them through the broker. """ - return getattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False) is True + return settings.EMBEDDINGS_DELEGATE_TO_WORKER is True def get_embeddings( diff --git a/docsgpt/vectorstore/embeddings_delegated.py b/docsgpt/vectorstore/embeddings_delegated.py index 0677c850..75f730b8 100644 --- a/docsgpt/vectorstore/embeddings_delegated.py +++ b/docsgpt/vectorstore/embeddings_delegated.py @@ -173,8 +173,8 @@ class DelegatedEmbeddings: the full timeout at once -- at the shipped 60s and 96 WSGI threads, an API that serves nothing at all, health checks included. """ - queue = getattr(settings, "EMBEDDINGS_QUEUE", "embeddings") - timeout = getattr(settings, "EMBEDDINGS_DELEGATE_TIMEOUT", 60) + queue = settings.EMBEDDINGS_QUEUE + timeout = settings.EMBEDDINGS_DELEGATE_TIMEOUT remaining = self._cooldown_remaining() if remaining > 0: diff --git a/docsgpt/vectorstore/embeddings_local.py b/docsgpt/vectorstore/embeddings_local.py index 259240d4..5a236e3b 100644 --- a/docsgpt/vectorstore/embeddings_local.py +++ b/docsgpt/vectorstore/embeddings_local.py @@ -188,8 +188,8 @@ def _describe_from_repo(repo: str) -> Optional[EmbeddingModel]: def _apply_overrides(spec: EmbeddingModel) -> EmbeddingModel: """Let ``EMBEDDINGS_POOLING``/``EMBEDDINGS_NORMALIZE`` win over any source.""" - pooling = getattr(settings, "EMBEDDINGS_POOLING", None) - normalize = getattr(settings, "EMBEDDINGS_NORMALIZE", None) + pooling = settings.EMBEDDINGS_POOLING + normalize = settings.EMBEDDINGS_NORMALIZE changes = {} if isinstance(pooling, str) and pooling.strip().lower() in ("cls", "mean"): changes["pooling"] = pooling.strip().lower() @@ -289,10 +289,10 @@ class EmbeddingsWrapper: try: _register(self.spec) init_kwargs = {"model_name": self.spec.repo} - threads = getattr(settings, "EMBEDDINGS_THREADS", None) + threads = settings.EMBEDDINGS_THREADS if isinstance(threads, int) and threads > 0: init_kwargs["threads"] = threads - cache_dir = getattr(settings, "EMBEDDINGS_CACHE_DIR", None) + cache_dir = settings.EMBEDDINGS_CACHE_DIR if cache_dir: init_kwargs["cache_dir"] = cache_dir self.model = TextEmbedding(**init_kwargs) @@ -331,7 +331,7 @@ class EmbeddingsWrapper: if not documents: return [] batch_size: Optional[int] = None - raw = getattr(settings, "EMBEDDINGS_MODEL_BATCH_SIZE", None) + raw = settings.EMBEDDINGS_MODEL_BATCH_SIZE if isinstance(raw, int) and not isinstance(raw, bool) and raw > 0: batch_size = raw diff --git a/docsgpt/vectorstore/pgconn.py b/docsgpt/vectorstore/pgconn.py index 5164629d..7c4aa33f 100644 --- a/docsgpt/vectorstore/pgconn.py +++ b/docsgpt/vectorstore/pgconn.py @@ -59,7 +59,7 @@ def resolve_pool_max_size() -> int: """ from docsgpt.core.settings import settings - value = getattr(settings, "PGVECTOR_POOL_MAX_SIZE", DEFAULT_POOL_MAX_SIZE) + value = settings.PGVECTOR_POOL_MAX_SIZE if isinstance(value, int) and not isinstance(value, bool) and value >= 0: return value return DEFAULT_POOL_MAX_SIZE diff --git a/docsgpt/vectorstore/pgvector.py b/docsgpt/vectorstore/pgvector.py index d586073d..60bb5262 100644 --- a/docsgpt/vectorstore/pgvector.py +++ b/docsgpt/vectorstore/pgvector.py @@ -54,9 +54,9 @@ class PGVectorStore(BaseVectorStore): # Use provided connection string or fall back to settings. # If PGVECTOR_CONNECTION_STRING is not set but POSTGRES_URI is, # reuse the same cluster — normalize from SQLAlchemy dialect to libpq form. - self._connection_string = connection_string or getattr(settings, 'PGVECTOR_CONNECTION_STRING', None) + self._connection_string = connection_string or settings.PGVECTOR_CONNECTION_STRING - if not self._connection_string and getattr(settings, 'POSTGRES_URI', None): + if not self._connection_string and settings.POSTGRES_URI: from docsgpt.core.db_uri import normalize_pgvector_connection_string self._connection_string = normalize_pgvector_connection_string(settings.POSTGRES_URI) diff --git a/tests/llm/test_openai_responses.py b/tests/llm/test_openai_responses.py index 620d3172..c02e08f2 100644 --- a/tests/llm/test_openai_responses.py +++ b/tests/llm/test_openai_responses.py @@ -29,6 +29,9 @@ def _make_llm(monkeypatch, capabilities=None, store_responses=False): AZURE_DEPLOYMENT_NAME="dep", OPENAI_RESPONSES_STORE=store_responses, OPENAI_REASONING_SUMMARY="auto", + OPENAI_RESPONSES_TRUNCATION_AUTO=False, + OPENAI_PROMPT_CACHE_KEY=False, + OPENAI_PROMPT_CACHE_RETENTION=None, ), ) from docsgpt.llm.openai import OpenAILLM diff --git a/tests/llm/test_responses_chain_budget.py b/tests/llm/test_responses_chain_budget.py index 12025f61..d41ef822 100644 --- a/tests/llm/test_responses_chain_budget.py +++ b/tests/llm/test_responses_chain_budget.py @@ -24,18 +24,19 @@ def _make_llm(monkeypatch, store_responses=True, **extra_settings): "docsgpt.llm.openai.StorageCreator", types.SimpleNamespace(get_storage=lambda: None), ) - monkeypatch.setattr( - "docsgpt.llm.openai.settings", - types.SimpleNamespace( - OPENAI_API_KEY="k", - API_KEY="k", - OPENAI_BASE_URL="", - AZURE_DEPLOYMENT_NAME="dep", - OPENAI_RESPONSES_STORE=store_responses, - OPENAI_REASONING_SUMMARY="auto", - **extra_settings, - ), - ) + # Every setting the Responses path reads, with the hints off; tests opt in per case. + stub = { + "OPENAI_API_KEY": "k", + "API_KEY": "k", + "OPENAI_BASE_URL": "", + "AZURE_DEPLOYMENT_NAME": "dep", + "OPENAI_RESPONSES_STORE": store_responses, + "OPENAI_REASONING_SUMMARY": "auto", + "OPENAI_RESPONSES_TRUNCATION_AUTO": False, + "OPENAI_PROMPT_CACHE_KEY": False, + "OPENAI_PROMPT_CACHE_RETENTION": None, + } + monkeypatch.setattr("docsgpt.llm.openai.settings", types.SimpleNamespace(**{**stub, **extra_settings})) from docsgpt.llm.openai import OpenAILLM llm = OpenAILLM(api_key="k") @@ -142,7 +143,7 @@ def _params(llm, **kwargs): @pytest.mark.unit -def test_build_responses_params_defaults_omit_truncation_and_cache_hints(monkeypatch): +def test_build_responses_params_omits_truncation_and_cache_hints_when_off(monkeypatch): llm = _make_llm(monkeypatch) llm._prompt_cache_key = "conv-123" params = _params(llm) diff --git a/tests/parser/connectors/test_share_point_auth.py b/tests/parser/connectors/test_share_point_auth.py index feab4c6b..ecf25a7d 100644 --- a/tests/parser/connectors/test_share_point_auth.py +++ b/tests/parser/connectors/test_share_point_auth.py @@ -14,8 +14,8 @@ def mock_settings(): s.MICROSOFT_TENANT_ID = "tenant-id-123" s.CONNECTOR_REDIRECT_BASE_URI = "https://redirect.example.com/callback" s.MONGO_DB_NAME = "test_db" - # Delete MICROSOFT_AUTHORITY so getattr falls back to default - del s.MICROSOFT_AUTHORITY + # Unset, as in a real Settings object, so the tenant-derived authority is used. + s.MICROSOFT_AUTHORITY = None return s From a20c83f4688384973be42ebd95b912f456b4f2f2 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 11:32:50 +0100 Subject: [PATCH 057/130] feat: a development loop in one command MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `docsgpt up --native` installs services meant to outlive the shell. Development wants the opposite, and until now it meant three terminals from the guide: uvicorn, celery, and vite. `docsgpt dev` runs this checkout's API and worker as children of one terminal, both restarting when a file is saved, their output interleaved and labelled, and Ctrl-C stopping them together. `--ui` adds the Vite dev server, `--mock-llm` runs the bundled mock model so no API key is needed, and `--no-worker` leaves the worker to your editor's debugger. Celery has no reloader of its own, so the worker is wrapped in watchfiles when it is installed, and runs plain when it is not. Alongside it, the commands a dev loop keeps reaching for: - `docsgpt doctor` checks what usually breaks a new setup: PostgreSQL answering and its schema matching this version, Redis answering, a model provider being configured, and the port being free. - `docsgpt restart [api|worker]` bounces services without rewriting settings or rerunning migrations, which `down` plus `up` did. - `docsgpt logs -f` follows a native install instead of telling you to run `tail -f` yourself. - `docsgpt env set` applies itself to a running native install rather than asking you to run `docsgpt up` again to change one value. Two bugs found on the way, both older than this change: - `docsgpt api --reload` watched the working directory, which in a checkout is 178,425 files: .venv, node_modules, and the indexes/ and inputs/ the app writes to while ingesting, so the server restarted itself mid-request. It watches the package now — 1,217 files. - The VS Code "Flask Debugger" ran `flask run`, which serves only the WSGI app: /mcp, the SSE streams and artifact downloads 404 under it. The guide warned about this in prose while the debug config did it anyway. It runs uvicorn on the ASGI app now, like production. --- .vscode/launch.json | 51 ++-- .../Deploying/Development-Environment.mdx | 52 +++- docs/content/changelog.mdx | 11 + docsgpt/cli.py | 35 ++- docsgpt/deploy/commands.py | 234 +++++++++++++++++- docsgpt/deploy/dev.py | 214 ++++++++++++++++ tests/deploy/test_dev.py | 180 ++++++++++++++ tests/deploy/test_doctor.py | 167 +++++++++++++ tests/test_cli.py | 24 +- 9 files changed, 932 insertions(+), 36 deletions(-) create mode 100644 docsgpt/deploy/dev.py create mode 100644 tests/deploy/test_dev.py create mode 100644 tests/deploy/test_doctor.py diff --git a/.vscode/launch.json b/.vscode/launch.json index 30700f70..28d4fb57 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -2,39 +2,36 @@ "version": "0.2.0", "configurations": [ { - "name": "Frontend Debug (npm)", + "name": "Frontend (npm)", "type": "node-terminal", "request": "launch", "command": "npm run dev", "cwd": "${workspaceFolder}/frontend" }, { - "name": "Flask Debugger", - "type": "debugpy", - "request": "launch", - "module": "flask", - "env": { - "FLASK_APP": "docsgpt/app.py", - "PYTHONPATH": "${workspaceFolder}", - "FLASK_ENV": "development", - "FLASK_DEBUG": "1", - "FLASK_RUN_PORT": "7091", - "FLASK_RUN_HOST": "0.0.0.0" - - }, - "args": [ - "run", - "--no-debugger" - ], - "cwd": "${workspaceFolder}", + "name": "API (uvicorn)", + "type": "debugpy", + "request": "launch", + "module": "uvicorn", + "env": { + "PYTHONPATH": "${workspaceFolder}" + }, + "args": [ + "docsgpt.asgi:asgi_app", + "--host", + "127.0.0.1", + "--port", + "7091" + ], + "cwd": "${workspaceFolder}" }, { - "name": "Celery Debugger", + "name": "Celery worker", "type": "debugpy", "request": "launch", "module": "celery", "env": { - "PYTHONPATH": "${workspaceFolder}", + "PYTHONPATH": "${workspaceFolder}" }, "args": [ "-A", @@ -47,10 +44,10 @@ "cwd": "${workspaceFolder}" }, { - "name": "Dev Containers (Mongo + Redis)", + "name": "Dev services (Postgres + Redis)", "type": "node-terminal", "request": "launch", - "command": "docker compose -f deployment/docker-compose-dev.yaml up --build", + "command": "docker compose -f deployment/docker-compose-dev.yaml up", "cwd": "${workspaceFolder}" } ], @@ -58,9 +55,9 @@ { "name": "DocsGPT: Full Stack", "configurations": [ - "Frontend Debug (npm)", - "Flask Debugger", - "Celery Debugger" + "Frontend (npm)", + "API (uvicorn)", + "Celery worker" ], "presentation": { "group": "DocsGPT", @@ -68,4 +65,4 @@ } } ] -} \ No newline at end of file +} diff --git a/docs/content/Deploying/Development-Environment.mdx b/docs/content/Deploying/Development-Environment.mdx index 6c01879f..133c5531 100644 --- a/docs/content/Deploying/Development-Environment.mdx +++ b/docs/content/Deploying/Development-Environment.mdx @@ -119,7 +119,34 @@ To run the DocsGPT backend locally, you'll need to set up a Python environment a 5. **Run the Backend:** - For local development, run the ASGI composition under uvicorn. It serves the **whole** application, hot-reloads on source changes, and matches the production runtime: + One command runs the API and the worker from this checkout, each restarting when you save a file: + + ```bash + docsgpt dev + ``` + + Both run as children of that terminal, with their output interleaved and labelled, and Ctrl-C stops + them together. Useful flags: + + | Flag | What it does | + | --- | --- | + | `--ui` | also start the Vite dev server, so the whole app runs from one command | + | `--mock-llm` | run `scripts/mock_llm.py` and point DocsGPT at it, so no API key is needed | + | `--no-worker` | leave the worker to you, for instance when debugging it in your editor | + | `--no-reload` | do not restart anything on save | + | `--port` | serve the API somewhere other than 7091 | + + `docsgpt dev` is for a checkout. `docsgpt up --native`, by contrast, installs supervised services + that outlive the shell — see [Run it as services](/Deploying/Pip-Install#run-it-as-services-without-docker). + + + `docsgpt doctor` checks the things that usually break a new setup: whether PostgreSQL answers and + its schema matches this version, whether Redis answers, whether a model provider is configured, + and whether the port is free. Run it first when something does not start. + + + To run the two processes yourself instead, start the ASGI composition under uvicorn. It serves the + **whole** application, hot-reloads on source changes, and matches the production runtime: ```bash uvicorn docsgpt.asgi:asgi_app --host 0.0.0.0 --port 7091 --reload @@ -135,7 +162,7 @@ To run the DocsGPT backend locally, you'll need to set up a Python environment a But it serves **only** the WSGI Flask app and omits the native-async routes mounted on the ASGI shell in `docsgpt/asgi.py`: the `/mcp` FastMCP endpoint, the chat reconnect reader `GET /api/messages//events`, the notification stream `GET /api/events`, the remote-device command stream `GET /api/devices/sessions//events`, and artifact downloads `GET /api/artifacts//download`. Under `flask run` those paths return 404 — chat still works (`POST /stream` is a Flask route), but live notifications, stream auto-resume, paired devices and artifact downloads don't. Use `flask run` only when you don't need them. -6. **Start the Celery Worker:** +6. **Start the Celery Worker** (not needed if you used `docsgpt dev`)**:** Open a new terminal window (and activate your virtual environment if you used one). Start the Celery worker to handle background tasks: @@ -153,10 +180,14 @@ To run the DocsGPT backend locally, you'll need to set up a Python environment a **Running in Debugger (VSCode):** -For easier debugging, you can launch the Flask app and Celery worker directly from VSCode's debugger. +For easier debugging, you can launch the API and the Celery worker directly from VSCode's debugger. * Press Shift + Cmd + D (macOS) or Shift + Windows + D (Windows) to open the Run and Debug view. -* You should see configurations named "Flask" and "Celery". Select the desired configuration and click the "Start Debugging" button (green play icon). +* You should see configurations named "API (uvicorn)" and "Celery worker", and a compound "DocsGPT: Full Stack" that starts them with the frontend. Select one and click the "Start Debugging" button (green play icon). + +The API configuration runs the same ASGI app as production, so the routes mounted on the ASGI shell +work under the debugger. It deliberately runs without `--reload`: the reloader restarts the server in +a child process, which your breakpoints would not be attached to. ## 3. Start the Frontend @@ -207,3 +238,16 @@ To run the DocsGPT frontend locally, you'll need Node.js and npm (Node Package M This command will start the Vite development server. The frontend application will typically be accessible at [http://localhost:5173/](http://localhost:5173/). The terminal will display the exact URL where the frontend is running. With both the backend and frontend running, you should now have a fully functional DocsGPT development environment. You can access the application in your browser at [http://localhost:5173/](http://localhost:5173/) and start developing! + +## Working on two branches at once + +Each install keeps its own directory and its own services, so a second branch can run beside the +first as long as it gets its own port: + +```bash +docsgpt dev --port 7092 # a second checkout, second terminal +docsgpt up --native --dir ~/.docsgpt/review --port 7092 # or a second installed copy +``` + +A native install in another directory gets its own service names, so the two never write over each +other's units. `docsgpt status --dir ~/.docsgpt/review` reports on that one alone. diff --git a/docs/content/changelog.mdx b/docs/content/changelog.mdx index ab78c4ad..59fa4ae8 100644 --- a/docs/content/changelog.mdx +++ b/docs/content/changelog.mdx @@ -13,6 +13,17 @@ request, and [Upgrading](/upgrading) covers the steps an existing deployment has ## Unreleased +### A development loop in one command + +`docsgpt dev` runs this checkout's API and worker as children of one terminal, both restarting when +you save, with their output interleaved and Ctrl-C stopping them together. `--ui` adds the Vite dev +server and `--mock-llm` runs the bundled mock model, so a working loop needs no API key. +`docsgpt doctor` checks PostgreSQL, its schema version, Redis, the model provider and the port; +`docsgpt restart` bounces the services without touching settings; `docsgpt logs -f` now follows a +native install; and `docsgpt env set` applies itself to a running native install instead of asking +you to run `docsgpt up` again. See +[Setting up a development environment](/Deploying/Development-Environment). + ### Run DocsGPT without Docker `docsgpt up --native` runs the API and the worker as services on the machine itself, launchd on diff --git a/docsgpt/cli.py b/docsgpt/cli.py index 1ca2a187..d837642d 100644 --- a/docsgpt/cli.py +++ b/docsgpt/cli.py @@ -89,7 +89,17 @@ def _api(args: argparse.Namespace) -> int: if args.reload or sys.platform == "win32": import uvicorn - uvicorn.run("docsgpt.asgi:asgi_app", host=args.host, port=args.port, reload=args.reload) + from docsgpt.core.paths import package_dir + + # Watch the package, not the working directory: a checkout also holds .venv, node_modules + # and the data the app writes (indexes/, inputs/), which restarts the server mid-ingest. + uvicorn.run( + "docsgpt.asgi:asgi_app", + host=args.host, + port=args.port, + reload=args.reload, + reload_dirs=[str(package_dir())] if args.reload else None, + ) return 0 _gunicorn_application(_gunicorn_options(args.host, args.port, args.workers)).run() @@ -246,12 +256,33 @@ def _add_deploy_commands(commands) -> None: restore.add_argument("--timeout", type=int, default=300, help="seconds to wait for the API afterwards (default: 300)") + doctor = stack_command("doctor", "doctor", "check what this machine needs to run DocsGPT") + doctor.add_argument("--postgres-uri", help="check this database instead of the one in .env") + doctor.add_argument("--redis-url", help="check this Redis instead of the one in .env") + + restart = stack_command("restart", "restart", "restart the services, changing nothing else") + restart.add_argument("services", nargs="*", help="services to restart, e.g. api worker") + env = stack_command("env", "env", "show, get or set the stack's settings") env_actions = env.add_subparsers(dest="env_action", metavar="") get = env_actions.add_parser("get", help="print one setting") get.add_argument("key") - set_ = env_actions.add_parser("set", help="set settings (KEY=VALUE ...); run `docsgpt up` to apply") + set_ = env_actions.add_parser("set", help="set settings (KEY=VALUE ...)") set_.add_argument("pairs", nargs="+", metavar="KEY=VALUE") + set_.add_argument("--no-restart", dest="restart", action="store_false", + help="do not restart a running native install afterwards") + + dev = commands.add_parser("dev", help="run this checkout's API, worker and UI with reload") + dev.add_argument("--host", default=DEFAULT_HOST, help="interface for the API (default: localhost)") + dev.add_argument("--port", type=int, default=DEFAULT_PORT, help="port for the API (default: 7091)") + dev.add_argument("--ui", action="store_true", help="also run the frontend dev server") + dev.add_argument("--mock-llm", action="store_true", + help="run the mock LLM and point DocsGPT at it, so no API key is needed") + dev.add_argument("--no-worker", dest="worker", action="store_false", help="do not run the Celery worker") + dev.add_argument("--no-reload", dest="reload", action="store_false", + help="do not restart the API and worker when a file changes") + dev.add_argument("-l", "--loglevel", default="INFO", help="worker log level (default: INFO)") + dev.set_defaults(func=_deploy("dev"), deploy=True) def build_parser() -> argparse.ArgumentParser: diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 0d95655f..8b1503a6 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -12,6 +12,7 @@ import socket import subprocess import sys import tempfile +import time import webbrowser from collections.abc import Callable, Mapping from dataclasses import dataclass @@ -623,12 +624,36 @@ def logs(args, context: Optional[Context] = None) -> int: else: print("(nothing logged yet)") if args.follow: - print("Following is not supported in native mode; use `tail -f` on the files above.", file=sys.stderr) + return _follow(logs_dir, wanted) return 0 options = (["--follow"] if args.follow else []) + (["--tail", str(args.tail)] if args.tail else []) return context.docker.compose(directory, "logs", *options, *args.services, check=False).returncode +def _follow(logs_dir: Path, services: list) -> int: + """Print new lines from each service's log until the terminal interrupts, prefixed by service.""" + handles: dict = {} + try: + while True: + for service in services: + if service not in handles: + path = logs_dir / f"{service}.log" + if not path.is_file(): + continue + handle = path.open("r", encoding="utf-8", errors="replace") + handle.seek(0, os.SEEK_END) + handles[service] = handle + for line in handles[service].readlines(): + print(f"{service:<6} | {line.rstrip()}") + sys.stdout.flush() + time.sleep(0.3) + except KeyboardInterrupt: + return 0 + finally: + for handle in handles.values(): + handle.close() + + def token(args, context: Optional[Context] = None) -> int: """Print the access token of a ``simple_jwt`` install.""" directory = stack.stack_dir(args.dir) @@ -681,10 +706,215 @@ def env(args, context: Optional[Context] = None) -> int: envfile.update(env_path, updates) except ValueError as exc: raise DeployError(str(exc)) from exc - print(f"Saved to {env_path}. Run `docsgpt up` to apply.") + print(f"Saved to {env_path}.") + directory = stack.stack_dir(args.dir) + if _mode(directory) == "native" and getattr(args, "restart", True): + context = context or Context.default(args) + services = context.service_manager() + names = _service_names(directory) + if any(services.is_running(name) for name in names): + for name in reversed(names): + services.stop(name) + for name in names: + services.start(name) + print("Restarted the services, so the change is live.") + return 0 + print("Run `docsgpt up` to apply.") return 0 +def dev(args, context: Optional[Context] = None) -> int: + """Run this checkout's API, worker and UI as children of this terminal.""" + from docsgpt.core import paths + from docsgpt.deploy import dev as dev_module + + checkout = paths.checkout_root() + if checkout is None: + raise DeployError( + "`docsgpt dev` runs the code in a source checkout, and this is an installed package. " + "Clone the repository and run it from there, or use `docsgpt up --native` to run this copy." + ) + if not _port_is_free(args.port): + raise DeployError( + f"port {args.port} is already in use, so the API cannot bind it. Stop what is on it " + f"(a previous `docsgpt dev`, or `docsgpt down` for an install), or pass --port." + ) + children = dev_module.plan(args, checkout) + print(f"DocsGPT from {checkout}") + for child in children: + print(f" {child.name:<6} {' '.join(child.command)}") + print(f"\nAPI http://{args.host}:{args.port}") + if getattr(args, "ui", False): + print(f"UI http://localhost:{dev_module.UI_PORT}") + print("Ctrl-C stops everything.\n") + return dev_module.run(children) + + +@dataclass +class Check: + """One line of ``docsgpt doctor``: what was looked at and what came back.""" + + name: str + level: str + detail: str + + +MARKS = {"ok": "ok ", "warn": "warn", "fail": "FAIL"} + + +def _migration_head() -> Optional[str]: + """The newest revision shipped with this package, or None when alembic cannot say.""" + try: + from alembic.config import Config + from alembic.script import ScriptDirectory + except ImportError: + return None + ini = Path(__file__).resolve().parents[1] / "alembic.ini" + if not ini.is_file(): + return None + config = Config(str(ini)) + config.set_main_option("script_location", str(ini.parent / "alembic")) + try: + return ScriptDirectory.from_config(config).get_current_head() + except Exception: # noqa: BLE001 - a broken script directory is a doctor finding, not a crash + return None + + +def _check_postgres(uri: Optional[str]) -> Check: + """Connect, and say whether the schema is the one this version expects.""" + if not uri: + return Check("postgres", "fail", "POSTGRES_URI is not set") + try: + import psycopg + except ImportError: + return Check("postgres", "fail", "the psycopg driver is not installed") + try: + with psycopg.connect(uri, connect_timeout=5) as connection, connection.cursor() as cursor: + cursor.execute("select current_setting('server_version')") + version = cursor.fetchone()[0] + cursor.execute("select to_regclass('public.alembic_version')") + applied = cursor.fetchone()[0] is not None + current = None + if applied: + cursor.execute("select version_num from alembic_version") + row = cursor.fetchone() + current = row[0] if row else None + except (psycopg.Error, OSError, ValueError) as exc: + return Check("postgres", "fail", f"cannot connect: {str(exc).strip()}") + head = _migration_head() + if not current: + return Check("postgres", "fail", f"PostgreSQL {version}, no schema yet; run `docsgpt migrate`") + if head and current != head: + return Check("postgres", "fail", f"PostgreSQL {version} at {current}, this version wants {head}; " + "run `docsgpt migrate`") + return Check("postgres", "ok", f"PostgreSQL {version}, schema at {current}") + + +def _check_redis(urls: Mapping[str, str]) -> Check: + """Ping every Redis the settings name; they are usually one server, three databases.""" + if not urls: + return Check("redis", "fail", "no Redis is configured (CELERY_BROKER_URL)") + try: + import redis + except ImportError: + return Check("redis", "fail", "the redis client is not installed") + for label, url in sorted(urls.items()): + try: + redis.Redis.from_url(url, socket_connect_timeout=3).ping() + except Exception as exc: # noqa: BLE001 - every client error here is the same finding + return Check("redis", "fail", f"{label} ({url}) does not answer: {str(exc).strip()}") + return Check("redis", "ok", f"answering on {len(urls)} database(s)") + + +def _check_provider(env: Mapping[str, str]) -> Check: + """Whether a model provider is set up well enough to answer a question.""" + provider = env.get("LLM_PROVIDER") or "docsgpt" + if provider == "docsgpt": + return Check("provider", "ok", "the DocsGPT public API (no key needed)") + if not (env.get("API_KEY") or env.get("OPENAI_API_KEY")): + return Check("provider", "fail", f"{provider} is configured but no API_KEY is set") + return Check("provider", "ok", f"{provider}{' at ' + env['OPENAI_BASE_URL'] if env.get('OPENAI_BASE_URL') else ''}") + + +def doctor(args, context: Optional[Context] = None) -> int: + """Check what DocsGPT needs on this machine, and say what is missing.""" + from docsgpt.core import paths + + if args.dir: + env_path = stack.stack_dir(args.dir) / ".env" + else: + try: + env_path = paths.env_file() + except FileNotFoundError as exc: + raise DeployError(str(exc)) from exc + env = envfile.read(env_path) + checks = [ + Check("settings", "ok" if env_path.is_file() else "warn", + f"{env_path}" if env_path.is_file() else f"{env_path} does not exist yet; defaults are in use"), + _check_postgres(args.postgres_uri or env.get("POSTGRES_URI")), + _check_redis({ + key: value for key, value in ( + ("broker", args.redis_url or env.get("CELERY_BROKER_URL")), + ("results", env.get("CELERY_RESULT_BACKEND")), + ("cache", env.get("CACHE_REDIS_URL")), + ) if value + }), + _check_provider(env), + ] + + port = int(env.get("DOCSGPT_PORT") or stack.DEFAULT_PORT) + if _port_is_free(port): + checks.append(Check("port", "ok", f"{port} is free")) + else: + checks.append(Check("port", "warn", f"{port} is in use, which is expected if DocsGPT is running")) + + directory = stack.stack_dir(args.dir) + if _mode(directory) == "native": + services = context.service_manager() if context else native.services_for_platform() + names = _service_names(directory) + running = [name for name in names if services.is_running(name)] + level = "ok" if len(running) == len(names) else "warn" + checks.append(Check("services", level, f"{len(running)} of {len(names)} running ({', '.join(names)})")) + + for check in checks: + print(f"[{MARKS[check.level]}] {check.name:<9} {check.detail}") + failed = [check for check in checks if check.level == "fail"] + if failed: + print(f"\n{len(failed)} problem(s) to fix before DocsGPT will work.", file=sys.stderr) + return 1 if failed else 0 + + +def _chosen_services(names: tuple[str, str], wanted: list) -> list: + """The services the user asked for, given either short names (api) or full ones.""" + if not wanted: + return list(names) + chosen = [] + for ask in wanted: + match = [name for name in names if name == ask or name.removeprefix("docsgpt-").startswith(ask)] + if not match: + raise DeployError(f"{ask!r} is not a service of this install; it has {', '.join(names)}.") + chosen.extend(match) + return chosen + + +def restart(args, context: Optional[Context] = None) -> int: + """Restart the services, changing nothing else.""" + context = context or Context.default(args) + directory = stack.stack_dir(args.dir) + if _installed(directory) is None: + return 1 + if _mode(directory) == "native": + services = context.service_manager() + chosen = _chosen_services(_service_names(directory), list(args.services)) + for name in reversed(chosen): + services.stop(name) + for name in chosen: + services.start(name) + print(f"Restarted {', '.join(chosen)}.") + return 0 + return context.docker.compose(directory, "restart", *args.services, check=False).returncode + + def _stack_image(env: Mapping[str, str]) -> str: """The image this install runs; the volume tars go through it, so nothing extra is pulled.""" tag = env.get("DOCSGPT_IMAGE_TAG") or "latest" diff --git a/docsgpt/deploy/dev.py b/docsgpt/deploy/dev.py new file mode 100644 index 00000000..a49dcf07 --- /dev/null +++ b/docsgpt/deploy/dev.py @@ -0,0 +1,214 @@ +"""``docsgpt dev``: this checkout's API, worker and UI as children of one terminal. + +``docsgpt up --native`` installs services meant to outlive the shell. Development wants the +opposite: processes rooted in the checkout, restarting when a file is saved, logging into one +terminal, and gone when Ctrl-C lands. This module decides which processes to run and supervises +them; nothing here imports the app itself. +""" + +from __future__ import annotations + +import os +import shlex +import shutil +import signal +import subprocess +import sys +import threading +import time +from dataclasses import dataclass, field +from pathlib import Path +from typing import Callable, Optional, TextIO + +from docsgpt.deploy.docker import DeployError + +MOCK_LLM_PORT = 8090 +UI_PORT = 5173 +STOP_GRACE = 10.0 + +# One colour per child so a glance at the terminal says who is talking. +COLOURS = {"api": "\033[36m", "worker": "\033[35m", "ui": "\033[32m", "llm": "\033[33m"} +RESET = "\033[0m" +WIDTH = 6 + + +@dataclass +class Child: + """One process ``docsgpt dev`` runs.""" + + name: str + command: list[str] + cwd: Path + env: dict[str, str] = field(default_factory=dict) + + +def watchfiles_available() -> bool: + """Whether the worker can be restarted on save; it arrives with uvicorn's standard extras.""" + try: + import watchfiles # noqa: F401 + except ImportError: + return False + return True + + +def _reloading_command(command: list[str], watched: Path) -> list[str]: + """``command`` under watchfiles, restarted when a Python file under ``watched`` changes.""" + return [sys.executable, "-m", "watchfiles", "--filter", "python", shlex.join(command), str(watched)] + + +def plan( + args, + checkout: Path, + *, + watching: Optional[bool] = None, + launcher: Optional[list[str]] = None, +) -> list[Child]: + """The children to run, in the order they should start.""" + from docsgpt.deploy import stack + + launcher = launcher or [sys.executable, "-m", "docsgpt"] + watching = watchfiles_available() if watching is None else watching + package = checkout / "docsgpt" + environment = {"DOCSGPT_HOME": str(checkout)} + children: list[Child] = [] + + if getattr(args, "mock_llm", False): + script = checkout / "scripts" / "mock_llm.py" + if not script.is_file(): + raise DeployError(f"{script} is missing, so there is no mock LLM to run.") + children.append( + Child( + name="llm", + command=[sys.executable, str(script), "--port", str(MOCK_LLM_PORT)], + cwd=checkout, + env=dict(environment), + ) + ) + # The children read these from the environment, so the checkout's .env is left alone. + chosen = stack.provider_settings( + "openai-compatible", model="mock", base_url=f"http://127.0.0.1:{MOCK_LLM_PORT}/v1" + ) + environment.update({key: value for key, value in chosen.items() if value is not None}) + + api = [*launcher, "api", "--host", args.host, "--port", str(args.port)] + if getattr(args, "reload", True): + api.append("--reload") + children.append(Child(name="api", command=api, cwd=checkout, env=dict(environment))) + + if getattr(args, "worker", True): + worker = [*launcher, "worker", "-l", getattr(args, "loglevel", "INFO")] + if getattr(args, "reload", True) and watching: + worker = _reloading_command(worker, package) + children.append(Child(name="worker", command=worker, cwd=checkout, env=dict(environment))) + + if getattr(args, "ui", False): + frontend = checkout / "frontend" + if not (frontend / "node_modules").is_dir(): + raise DeployError( + f"the frontend has no node_modules yet. Run `npm install --include=dev` in {frontend} " + "and try again, or leave --ui off." + ) + if not shutil.which("npm"): + raise DeployError("npm is not on PATH, so the frontend dev server cannot start.") + children.append(Child(name="ui", command=["npm", "run", "dev"], cwd=frontend, env=dict(environment))) + + return children + + +def _line(name: str, text: str, colour: bool) -> str: + """One output line, prefixed with the child that wrote it.""" + label = name.ljust(WIDTH) + if colour: + return f"{COLOURS.get(name, '')}{label}{RESET} | {text}" + return f"{label} | {text}" + + +def _pump(child: Child, process, out: TextIO, lock: threading.Lock, colour: bool) -> None: + """Copy one child's output to ``out``, a line at a time, prefixed.""" + stream = process.stdout + if stream is None: + return + for text in stream: + with lock: + out.write(_line(child.name, text.rstrip("\n"), colour)) + out.write("\n") + out.flush() + + +def _signal(process, number: int) -> None: + """Signal a child and, on POSIX, everything it started.""" + try: + if os.name == "nt": + process.terminate() + return + os.killpg(os.getpgid(process.pid), number) + except (ProcessLookupError, PermissionError, OSError): + pass + + +def _stop(running: list[tuple[Child, object]], grace: float, sleep: Callable[[float], None]) -> None: + """Interrupt the children, then insist if they are still there. + + A second Ctrl-C lands while this is waiting. It means "stop waiting", not "give up": the wait + ends and the children are killed, rather than the interrupt escaping and leaving them running. + """ + for _, process in running: + if process.poll() is None: + _signal(process, signal.SIGINT) + deadline = time.monotonic() + grace + while time.monotonic() < deadline and any(process.poll() is None for _, process in running): + try: + sleep(0.1) + except KeyboardInterrupt: + break + for _, process in running: + if process.poll() is None: + _signal(process, signal.SIGKILL) + + +def run( + children: list[Child], + *, + out: TextIO = sys.stdout, + spawn: Callable[..., object] = subprocess.Popen, + sleep: Callable[[float], None] = time.sleep, + grace: float = STOP_GRACE, + colour: Optional[bool] = None, +) -> int: + """Start the children and keep them running until one exits or the terminal interrupts.""" + colour = out.isatty() if colour is None else colour + lock = threading.Lock() + running: list[tuple[Child, object]] = [] + try: + for child in children: + process = spawn( + child.command, + cwd=str(child.cwd), + env={**os.environ, **child.env}, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + bufsize=1, + # Its own session, so Ctrl-C reaches this process and the children are stopped in order. + start_new_session=os.name != "nt", + ) + running.append((child, process)) + threading.Thread(target=_pump, args=(child, process, out, lock, colour), daemon=True).start() + + while True: + for child, process in running: + code = process.poll() + if code is not None: + with lock: + out.write(_line(child.name, f"exited with {code}", colour)) + out.write("\n") + out.flush() + return code or 1 + sleep(0.2) + except KeyboardInterrupt: + with lock: + out.write("\nStopping ...\n") + out.flush() + return 0 + finally: + _stop(running, grace, sleep) diff --git a/tests/deploy/test_dev.py b/tests/deploy/test_dev.py new file mode 100644 index 00000000..a8344954 --- /dev/null +++ b/tests/deploy/test_dev.py @@ -0,0 +1,180 @@ +"""`docsgpt dev`: the checkout's processes as children of one terminal.""" + +import argparse +import io +import signal +import sys +from pathlib import Path + +import pytest + +from docsgpt.deploy import dev +from docsgpt.deploy.docker import DeployError + + +def _args(**overrides): + values = {"host": "127.0.0.1", "port": 7091, "ui": False, "mock_llm": False, + "worker": True, "reload": True, "loglevel": "INFO"} + values.update(overrides) + return argparse.Namespace(**values) + + +def _named(children): + return [child.name for child in children] + + +class FakeProcess: + """A child that produces the given lines and is done.""" + + def __init__(self, lines=(), code=None): + self.stdout = iter(list(lines)) + self.code = code + self.pid = -1 + + def poll(self): + return self.code + + def terminate(self): + self.code = self.code if self.code is not None else -15 + + +class TestPlan: + def test_the_api_and_worker_run_from_the_checkout(self, tmp_path): + children = dev.plan(_args(), tmp_path, watching=False) + assert _named(children) == ["api", "worker"] + api, worker = children + assert api.command[-6:-1] == ["api", "--host", "127.0.0.1", "--port", "7091"] + assert api.command[-1] == "--reload" + assert worker.command[-2:] == ["-l", "INFO"] + for child in children: + assert child.cwd == tmp_path + assert child.env["DOCSGPT_HOME"] == str(tmp_path), "the checkout is the data home, not ~/.docsgpt" + + def test_the_worker_restarts_on_save_when_watchfiles_is_there(self, tmp_path): + """Celery has no reloader of its own, so it is wrapped in one.""" + worker = dev.plan(_args(), tmp_path, watching=True)[1] + assert "watchfiles" in worker.command + assert str(tmp_path / "docsgpt") in worker.command, "it watches the package, not the whole checkout" + assert "docsgpt worker" in " ".join(worker.command) + + def test_without_watchfiles_the_worker_still_runs(self, tmp_path): + worker = dev.plan(_args(), tmp_path, watching=False)[1] + assert "watchfiles" not in " ".join(worker.command) + + def test_no_reload_leaves_both_alone(self, tmp_path): + children = dev.plan(_args(reload=False), tmp_path, watching=True) + assert "--reload" not in children[0].command + assert "watchfiles" not in " ".join(children[1].command) + + def test_no_worker(self, tmp_path): + assert _named(dev.plan(_args(worker=False), tmp_path, watching=False)) == ["api"] + + def test_the_mock_llm_starts_first_and_the_others_are_pointed_at_it(self, tmp_path): + """A dev loop that needs no API key: the mock has to be up before the API asks it anything.""" + (tmp_path / "scripts").mkdir() + (tmp_path / "scripts" / "mock_llm.py").write_text("", encoding="utf-8") + children = dev.plan(_args(mock_llm=True), tmp_path, watching=False) + assert _named(children) == ["llm", "api", "worker"] + api = children[1] + assert api.env["LLM_PROVIDER"] == "openai" + assert api.env["OPENAI_BASE_URL"] == f"http://127.0.0.1:{dev.MOCK_LLM_PORT}/v1" + assert api.env["API_KEY"], "the client wants some key, even a placeholder" + assert "OPENAI_BASE_URL" not in children[0].env, "the mock itself does not need pointing at itself" + + def test_a_missing_mock_llm_script_is_reported(self, tmp_path): + with pytest.raises(DeployError, match="mock LLM"): + dev.plan(_args(mock_llm=True), tmp_path, watching=False) + + def test_the_ui_needs_its_dependencies(self, tmp_path): + (tmp_path / "frontend").mkdir() + with pytest.raises(DeployError, match="node_modules"): + dev.plan(_args(ui=True), tmp_path, watching=False) + + def test_the_ui_needs_npm_on_path(self, tmp_path, monkeypatch): + (tmp_path / "frontend" / "node_modules").mkdir(parents=True) + monkeypatch.setattr(dev.shutil, "which", lambda name: None) + with pytest.raises(DeployError, match="npm"): + dev.plan(_args(ui=True), tmp_path, watching=False) + + def test_the_ui_runs_in_the_frontend_directory(self, tmp_path, monkeypatch): + (tmp_path / "frontend" / "node_modules").mkdir(parents=True) + monkeypatch.setattr(dev.shutil, "which", lambda name: "/usr/local/bin/npm") + children = dev.plan(_args(ui=True), tmp_path, watching=False) + ui = children[-1] + assert ui.name == "ui" + assert ui.command == ["npm", "run", "dev"] + assert ui.cwd == tmp_path / "frontend" + + +class TestRun: + def _spawn(self, processes): + made = iter(processes) + + def spawn(command, **kwargs): + return next(made) + + return spawn + + def test_output_is_prefixed_with_the_child_that_wrote_it(self, tmp_path): + out = io.StringIO() + children = [dev.Child(name="api", command=["true"], cwd=tmp_path)] + code = dev.run(children, out=out, spawn=self._spawn([FakeProcess(["hello\n"], code=0)]), + sleep=lambda _: None, colour=False) + assert "api | hello" in out.getvalue() + assert code == 1, "a child that ends by itself ends the session, however it exited" + + def test_a_failing_child_returns_its_code(self, tmp_path): + out = io.StringIO() + children = [dev.Child(name="api", command=["false"], cwd=tmp_path)] + code = dev.run(children, out=out, spawn=self._spawn([FakeProcess(code=2)]), + sleep=lambda _: None, colour=False) + assert code == 2 + assert "exited with 2" in out.getvalue() + + def test_an_interrupt_stops_quietly(self, tmp_path): + out = io.StringIO() + + def sleep(_): + raise KeyboardInterrupt + + children = [dev.Child(name="api", command=["sleep"], cwd=tmp_path)] + code = dev.run(children, out=out, spawn=self._spawn([FakeProcess()]), sleep=sleep, colour=False) + assert code == 0 + assert "Stopping" in out.getvalue() + + def test_a_second_interrupt_during_shutdown_still_kills(self, tmp_path, monkeypatch): + """Ctrl-C twice is what you press when it did not die; it must not leave children behind.""" + sent = [] + monkeypatch.setattr(dev, "_signal", lambda process, number: sent.append(number)) + + def sleep(_): + raise KeyboardInterrupt + + children = [dev.Child(name="api", command=["sleep"], cwd=tmp_path)] + code = dev.run(children, out=io.StringIO(), spawn=self._spawn([FakeProcess()]), + sleep=sleep, colour=False, grace=5) + assert code == 0 + assert signal.SIGINT in sent + assert signal.SIGKILL in sent, "the second interrupt escalates instead of escaping" + + def test_children_get_the_checkout_environment(self, tmp_path, monkeypatch): + seen = {} + + def spawn(command, **kwargs): + seen.update(kwargs) + return FakeProcess(code=0) + + monkeypatch.setenv("SOMETHING_ELSE", "kept") + children = [dev.Child(name="api", command=["x"], cwd=tmp_path, env={"DOCSGPT_HOME": str(tmp_path)})] + dev.run(children, out=io.StringIO(), spawn=spawn, sleep=lambda _: None, colour=False) + assert seen["env"]["DOCSGPT_HOME"] == str(tmp_path) + assert seen["env"]["SOMETHING_ELSE"] == "kept", "the shell's environment is kept, not replaced" + assert seen["cwd"] == str(tmp_path) + assert seen["start_new_session"] is (sys.platform != "win32") + + +class TestReloadingCommand: + def test_an_interpreter_path_with_spaces_survives(self): + command = dev._reloading_command(["/opt/my venv/bin/python", "-m", "docsgpt", "worker"], Path("/srv/pkg")) + assert "'/opt/my venv/bin/python' -m docsgpt worker" in command + assert command[-1] == "/srv/pkg" diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py new file mode 100644 index 00000000..1309d512 --- /dev/null +++ b/tests/deploy/test_doctor.py @@ -0,0 +1,167 @@ +"""`docsgpt doctor`, `restart`, following native logs, and settings that apply themselves.""" + +import pytest + +from docsgpt.deploy import commands, envfile +from docsgpt.deploy.docker import DeployError + +from .test_commands import FakeDocker, _context, _run +from .test_native import FakeServices, _names, _native_context + + +def _installed_native(tmp_path, services): + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + assert _run(argv, _native_context(services)) == 0 + + +class TestRestart: + def test_it_stops_and_starts_both_without_touching_settings(self, tmp_path): + services = FakeServices() + _installed_native(tmp_path, services) + before = (tmp_path / ".env").read_text(encoding="utf-8") + services.started.clear() + services.stopped.clear() + + assert _run(["restart", "--dir", str(tmp_path)], _native_context(services)) == 0 + assert services.stopped == list(reversed(_names(tmp_path))), "the worker goes down first" + assert services.started == list(_names(tmp_path)) + assert (tmp_path / ".env").read_text(encoding="utf-8") == before + + def test_one_service_by_its_short_name(self, tmp_path): + services = FakeServices() + _installed_native(tmp_path, services) + services.started.clear() + assert _run(["restart", "api", "--dir", str(tmp_path)], _native_context(services)) == 0 + assert services.started == [_names(tmp_path)[0]] + + def test_a_name_this_install_does_not_have(self, tmp_path): + services = FakeServices() + _installed_native(tmp_path, services) + with pytest.raises(DeployError, match="not a service"): + _run(["restart", "frontend", "--dir", str(tmp_path)], _native_context(services)) + + def test_a_docker_install_restarts_its_containers(self, tmp_path): + assert _run(["up", "--yes", "--dir", str(tmp_path)], _context()) == 0 + docker = FakeDocker() + assert _run(["restart", "--dir", str(tmp_path)], _context(docker)) == 0 + assert ["restart"] in [args for _, args in docker.calls] + + +class TestEnvApplies: + def test_a_running_native_install_restarts_itself(self, tmp_path, capsys): + """The old advice was to run `docsgpt up` again, which reruns migrations to change one value.""" + services = FakeServices() + _installed_native(tmp_path, services) + services.started.clear() + + argv = ["env", "--dir", str(tmp_path), "set", "LLM_NAME=gpt-4o"] + assert _run(argv, _native_context(services)) == 0 + assert envfile.read(tmp_path / ".env")["LLM_NAME"] == "gpt-4o" + assert services.started == list(_names(tmp_path)) + assert "Restarted" in capsys.readouterr().out + + def test_no_restart_leaves_the_services_alone(self, tmp_path, capsys): + services = FakeServices() + _installed_native(tmp_path, services) + services.started.clear() + + argv = ["env", "--dir", str(tmp_path), "set", "LLM_NAME=gpt-4o", "--no-restart"] + assert _run(argv, _native_context(services)) == 0 + assert services.started == [] + assert "docsgpt up" in capsys.readouterr().out + + def test_a_stopped_install_is_not_started_by_a_settings_change(self, tmp_path): + services = FakeServices() + _installed_native(tmp_path, services) + for name in _names(tmp_path): + services.stop(name) + services.started.clear() + + argv = ["env", "--dir", str(tmp_path), "set", "LLM_NAME=gpt-4o"] + assert _run(argv, _native_context(services)) == 0 + assert services.started == [], "changing a setting does not start a stopped install" + + +class TestFollowLogs: + def test_it_prints_lines_written_after_it_started(self, tmp_path, capsys, monkeypatch): + logs = tmp_path / "logs" + logs.mkdir() + (logs / "api.log").write_text("old line\n", encoding="utf-8") + + rounds = {"n": 0} + + def sleep(_): + rounds["n"] += 1 + if rounds["n"] == 1: + (logs / "api.log").open("a", encoding="utf-8").write("new line\n") + return + raise KeyboardInterrupt + + monkeypatch.setattr(commands.time, "sleep", sleep) + assert commands._follow(logs, ["api"]) == 0 + printed = capsys.readouterr().out + assert "api | new line" in printed + assert "old line" not in printed, "it starts at the end, like tail -f" + + +class TestChecks: + def test_the_public_api_needs_no_key(self): + check = commands._check_provider({"LLM_PROVIDER": "docsgpt"}) + assert check.level == "ok" + + def test_a_provider_without_a_key_is_a_problem(self): + check = commands._check_provider({"LLM_PROVIDER": "openai"}) + assert check.level == "fail" + assert "API_KEY" in check.detail + + def test_a_provider_with_a_key_and_a_base_url(self): + check = commands._check_provider( + {"LLM_PROVIDER": "openai", "API_KEY": "x", "OPENAI_BASE_URL": "http://localhost:8090/v1"} + ) + assert check.level == "ok" + assert "8090" in check.detail + + def test_services_are_named_for_the_install(self, tmp_path): + names = _names(tmp_path) + assert commands._chosen_services(names, []) == list(names) + assert commands._chosen_services(names, ["worker"]) == [names[1]] + assert commands._chosen_services(names, [names[0]]) == [names[0]] + + +class TestDoctor: + def _only(self, monkeypatch, postgres, redis): + monkeypatch.setattr(commands, "_check_postgres", lambda uri: postgres) + monkeypatch.setattr(commands, "_check_redis", lambda urls: redis) + + def test_it_reports_every_check_and_succeeds_when_they_pass(self, tmp_path, capsys, monkeypatch): + self._only( + monkeypatch, + commands.Check("postgres", "ok", "PostgreSQL 16.2, schema at 0031"), + commands.Check("redis", "ok", "answering on 3 database(s)"), + ) + (tmp_path / ".env").write_text("LLM_PROVIDER=docsgpt\n", encoding="utf-8") + assert _run(["doctor", "--dir", str(tmp_path)], _context()) == 0 + out = capsys.readouterr().out + assert "postgres" in out and "redis" in out and "provider" in out + + def test_a_failing_check_makes_it_exit_one(self, tmp_path, capsys, monkeypatch): + self._only( + monkeypatch, + commands.Check("postgres", "fail", "cannot connect: refused"), + commands.Check("redis", "ok", "answering"), + ) + (tmp_path / ".env").write_text("LLM_PROVIDER=docsgpt\n", encoding="utf-8") + assert _run(["doctor", "--dir", str(tmp_path)], _context()) == 1 + captured = capsys.readouterr() + assert "FAIL" in captured.out + assert "1 problem" in captured.err + + def test_it_says_which_settings_file_it_read(self, tmp_path, capsys, monkeypatch): + self._only( + monkeypatch, + commands.Check("postgres", "ok", "fine"), + commands.Check("redis", "ok", "fine"), + ) + (tmp_path / ".env").write_text("LLM_PROVIDER=docsgpt\n", encoding="utf-8") + _run(["doctor", "--dir", str(tmp_path)], _context()) + assert str(tmp_path / ".env") in capsys.readouterr().out diff --git a/tests/test_cli.py b/tests/test_cli.py index dcad465e..bdb2977d 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -10,6 +10,7 @@ import click import pytest from docsgpt import cli +from docsgpt.core.paths import package_dir from docsgpt.version import __version__ @@ -126,7 +127,28 @@ class TestApi: uvicorn = types.SimpleNamespace(run=MagicMock()) monkeypatch.setitem(sys.modules, "uvicorn", uvicorn) assert cli.main(["api", "--reload", "--host", "127.0.0.1"]) == 0 - uvicorn.run.assert_called_once_with("docsgpt.asgi:asgi_app", host="127.0.0.1", port=7091, reload=True) + uvicorn.run.assert_called_once_with( + "docsgpt.asgi:asgi_app", host="127.0.0.1", port=7091, reload=True, + reload_dirs=[str(package_dir())], + ) + + def test_reload_watches_the_package_not_the_working_directory(self, monkeypatch, tmp_path): + """A checkout also holds .venv, node_modules and the indexes and inputs the app writes to, + so watching the working directory restarts the server mid-ingest.""" + uvicorn = types.SimpleNamespace(run=MagicMock()) + monkeypatch.setitem(sys.modules, "uvicorn", uvicorn) + monkeypatch.chdir(tmp_path) + assert cli.main(["api", "--reload"]) == 0 + watched = uvicorn.run.call_args.kwargs["reload_dirs"] + assert watched == [str(package_dir())] + assert str(tmp_path) not in watched + + def test_without_reload_nothing_is_watched(self, monkeypatch): + uvicorn = types.SimpleNamespace(run=MagicMock()) + monkeypatch.setitem(sys.modules, "uvicorn", uvicorn) + monkeypatch.setattr(sys, "platform", "win32") + assert cli.main(["api"]) == 0 + assert uvicorn.run.call_args.kwargs["reload_dirs"] is None class TestWorker: From 5578039c19c70d7cb9e8cf6e4dabffde7e481834 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:37:57 +0100 Subject: [PATCH 058/130] refactor(settings): treat unset spellings of every optional string as None Review follow-up. The per-group secret validators normalised a hand-picked list of API keys, which left other optional credentials and overrides (OPEN_ROUTER_API_KEY, S3 and Daytona keys, ELASTIC_PASSWORD, the OIDC trio, connector client ids, MICROSOFT_AUTHORITY, MCP_OAUTH_REDIRECT_URI) holding the literal "None" or "" a .env file spells "unset" with, so truthiness checks and fallbacks downstream saw a value. One rule on the group base replaces those lists: every Optional[str] field maps "", "None" and whitespace to None and strips real values. Plain str fields are left alone. The OIDC required-settings check therefore also rejects those spellings. EMBEDDINGS_POOLING is Literal["cls", "mean"] with case-insensitive parsing; its consumer silently ignored anything else. Bounds added where the consumer rejects or misbehaves on the value: SCHEDULE_RUN_OUTPUT_RETENTION_DAYS and MESSAGE_EVENTS_RETENTION_DAYS (the cleanup repositories raise on <= 0), EMBEDDINGS_DELEGATE_TIMEOUT, the remote-device idle/pairing/invocation TTLs and CELERY_VISIBILITY_TIMEOUT (> 0), REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS (> 605, the documented drain deadline), GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION (>= 0; negative would slice the pending list from the end). The generated reference now renders generic type arguments (dict[str, int] rather than dict). --- docs/content/Deploying/Settings-Reference.mdx | 32 ++++++------ docsgpt/core/settings/_shared.py | 29 ++++++++++- docsgpt/core/settings/auth.py | 5 -- docsgpt/core/settings/embeddings.py | 14 +++--- docsgpt/core/settings/events.py | 8 +-- docsgpt/core/settings/llm.py | 19 +------ docsgpt/core/settings/reference.py | 4 +- docsgpt/core/settings/retrieval.py | 2 +- docsgpt/core/settings/scheduler.py | 2 +- docsgpt/core/settings/speech.py | 7 +-- docsgpt/core/settings/vectorstores.py | 7 +-- docsgpt/core/settings/workers.py | 1 + tests/core/test_settings.py | 50 +++++++++++++++++-- 13 files changed, 113 insertions(+), 67 deletions(-) diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 02876714..30221972 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -271,7 +271,7 @@ Context window assumed when the model is not found in the registry. ### `RESERVED_TOKENS` -Type `dict`, default `{"system_prompt": 500, "current_query": 500, "safety_buffer": 1000}`. +Type `dict[str, int]`, default `{"system_prompt": 500, "current_query": 500, "safety_buffer": 1000}`. Tokens held back from the context window for the system prompt, the query and a safety buffer. @@ -378,7 +378,7 @@ Where embedding models and their tokenizers are cached. Persistent by default: F ### `EMBEDDINGS_POOLING` -Type `str`, default unset. +Type `"cls" | "mean"`, default unset. Pooling strategy ("cls" or "mean"). Read from the model's own repository; set only for a repository that declares none, or to override what it declares. @@ -402,7 +402,7 @@ Celery queue the embed task is routed to. ### `EMBEDDINGS_DELEGATE_TIMEOUT` -Type `int`, default `60`. +Type `int`, default `60`, must be > 0. Seconds the API waits for the worker to return an embedding. @@ -443,9 +443,9 @@ Model for ingest-time graph extraction; unset reuses LLM_PROVIDER/LLM_NAME. ### `GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION` -Type `int`, default `2000`. +Type `int`, default `2000`, must be >= 0. -Hard cap on chunks extracted per source (cost control). +Hard cap on chunks extracted per source (cost control); 0 extracts nothing. ## Vector stores @@ -668,7 +668,7 @@ Tasks prefetched per worker process; 1 caps SIGKILL loss to one task. ### `CELERY_VISIBILITY_TIMEOUT` -Type `int`, default `3600`. +Type `int`, default `3600`, must be > 0. Broker visibility timeout in seconds. Must exceed the longest legitimate task runtime but stay short enough that SIGKILLed tasks redeliver promptly. @@ -1220,13 +1220,13 @@ Length of the replay budget window. ### `MESSAGE_EVENTS_RETENTION_DAYS` -Type `int`, default `14`. +Type `int`, default `14`, must be > 0. Retention for the message_events journal, enforced by the cleanup_message_events beat task. Replay only needs streams a client could still be tailing. ### `REMOTE_DEVICE_SESSION_IDLE_SECONDS` -Type `int`, default `60`. +Type `int`, default `60`, must be > 0. Seconds without a heartbeat before a remote-device session is considered idle. @@ -1238,19 +1238,19 @@ Require signed commands from remote devices. ### `REMOTE_DEVICE_PAIRING_TTL_SECONDS` -Type `int`, default `600`. +Type `int`, default `600`, must be > 0. Lifetime of a pairing code. ### `REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS` -Type `int`, default `900`. +Type `int`, default `900`, must be > 605. Redis TTL of the per-device command queue, routing invocations cross-process so a scheduled run reaches the web-held device session. Must exceed the max drain deadline (605s) so a command for a briefly-offline device isn't evicted before its own drain gives up. ### `REMOTE_DEVICE_INVOCATION_TTL_SECONDS` -Type `int`, default `900`. +Type `int`, default `900`, must be > 0. Redis TTL of a pending remote-device invocation. @@ -1273,13 +1273,13 @@ Default agent type for agentless chats. ### `DEFAULT_AGENT_LIMITS` -Type `dict`, default `{"token_limit": 50000, "request_limit": 500}`. +Type `dict[str, int]`, default `{"token_limit": 50000, "request_limit": 500}`. Per-agent default quotas: tokens and requests. ### `DEFAULT_CHAT_TOOLS` -Type `list`, default `["memory", "read_webpage", "scheduler"]`. +Type `list[str]`, default `["memory", "read_webpage", "scheduler"]`. Config-free tools on by default in agentless chats. scheduler is dual-registered in BUILTIN_AGENT_TOOLS so one synthetic id resolves via defaults or the agent picker. Add code_executor and artifact_generator once a sandbox runner is configured; both execute through it and would fail on every call without one. @@ -1386,13 +1386,13 @@ Master switch; False disables every stage. ### `GUARDRAILS_CHECKS_ENABLED` -Type `list`, default `[]`. +Type `list[str]`, default `[]`. Allowlist of GuardrailCreator.checks keys; empty means every registered check. ### `GUARDRAILS_FLOOR` -Type `dict`, default `{}`. +Type `dict[str, Any]`, default `{}`. A GuardrailsConfig fragment every agent inherits and cannot weaken; agents may add controls or make an action stricter, never looser. "enabled" is required; without it the floor parses but applies to nothing. Example: \{"enabled": true, "mode": "scan_all", "controls": [\{"check": "secrets", "stage": "output", "action": "redact"\}]\} @@ -1463,7 +1463,7 @@ How far ahead a one-off run may be scheduled, in seconds (one year). ### `SCHEDULE_RUN_OUTPUT_RETENTION_DAYS` -Type `int`, default `90`. +Type `int`, default `90`, must be > 0. Days scheduled-run output is kept. diff --git a/docsgpt/core/settings/_shared.py b/docsgpt/core/settings/_shared.py index 235ee07b..59723bbf 100644 --- a/docsgpt/core/settings/_shared.py +++ b/docsgpt/core/settings/_shared.py @@ -9,16 +9,43 @@ domain's definitions live in their own module. from __future__ import annotations +import types +import typing from typing import Any, Optional +from pydantic import model_validator from pydantic_settings import BaseSettings, SettingsConfigDict +def _is_optional_str(annotation: Any) -> bool: + if typing.get_origin(annotation) not in (typing.Union, types.UnionType): + return False + return set(typing.get_args(annotation)) == {str, type(None)} + + class SettingsGroup(BaseSettings): - """Base for one domain's settings; groups are composed into ``Settings``.""" + """Base for one domain's settings; groups are composed into ``Settings``. + + Every ``Optional[str]`` field treats the spellings an unset value has in a + ``.env`` file (``KEY=``, ``KEY=None``, whitespace) as ``None``, so a check + like ``if settings.OIDC_ISSUER`` or a fallback like ``settings.X or default`` + sees "unset" rather than a truthy placeholder string. Real values are + stripped. Fields typed ``str`` keep whatever they are given. + """ model_config = SettingsConfigDict(extra="ignore") + @model_validator(mode="before") + @classmethod + def _unset_optional_strings(cls, data: Any) -> Any: + if not isinstance(data, dict): + return data + data = dict(data) + for name, field in cls.model_fields.items(): + if name in data and _is_optional_str(field.annotation): + data[name] = normalize_secret(data[name]) + return data + def normalize_choice(value: Any) -> Any: """Case-fold a closed-choice setting so ``PGVector`` and ``pgvector`` are the same choice.""" diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py index 04b039bc..c9309775 100644 --- a/docsgpt/core/settings/auth.py +++ b/docsgpt/core/settings/auth.py @@ -85,11 +85,6 @@ class AuthSettings(SettingsGroup): default=None, description="Bearer token for IdP SCIM clients (required when SCIM is enabled)." ) - @field_validator("INTERNAL_KEY", mode="before") - @classmethod - def _normalize_auth_secrets(cls, v): - return normalize_secret(v) - @field_validator("AUTH_TYPE", mode="before") @classmethod def _normalize_auth_type(cls, v): diff --git a/docsgpt/core/settings/embeddings.py b/docsgpt/core/settings/embeddings.py index 3ac749aa..f8c5c130 100644 --- a/docsgpt/core/settings/embeddings.py +++ b/docsgpt/core/settings/embeddings.py @@ -2,12 +2,12 @@ from __future__ import annotations -from typing import Optional +from typing import Literal, Optional from pydantic import Field, field_validator from docsgpt.core.paths import home_dir -from docsgpt.core.settings._shared import SettingsGroup, normalize_secret +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice class EmbeddingsSettings(SettingsGroup): @@ -56,7 +56,7 @@ class EmbeddingsSettings(SettingsGroup): "default is the temp dir." ), ) - EMBEDDINGS_POOLING: Optional[str] = Field( + EMBEDDINGS_POOLING: Optional[Literal["cls", "mean"]] = Field( default=None, description=( 'Pooling strategy ("cls" or "mean"). Read from the model\'s own repository; set only for a ' @@ -79,10 +79,10 @@ class EmbeddingsSettings(SettingsGroup): ) EMBEDDINGS_QUEUE: str = Field(default="embeddings", description="Celery queue the embed task is routed to.") EMBEDDINGS_DELEGATE_TIMEOUT: int = Field( - default=60, description="Seconds the API waits for the worker to return an embedding." + default=60, gt=0, description="Seconds the API waits for the worker to return an embedding." ) - @field_validator("EMBEDDINGS_KEY", mode="before") + @field_validator("EMBEDDINGS_POOLING", mode="before") @classmethod - def _normalize_embeddings_secrets(cls, v): - return normalize_secret(v) + def _normalize_pooling(cls, v): + return normalize_choice(v) diff --git a/docsgpt/core/settings/events.py b/docsgpt/core/settings/events.py index f8ca6742..3b170be9 100644 --- a/docsgpt/core/settings/events.py +++ b/docsgpt/core/settings/events.py @@ -58,6 +58,7 @@ class EventsSettings(SettingsGroup): EVENTS_REPLAY_BUDGET_WINDOW_SECONDS: int = Field(default=60, description="Length of the replay budget window.") MESSAGE_EVENTS_RETENTION_DAYS: int = Field( default=14, + gt=0, description=( "Retention for the message_events journal, enforced by the cleanup_message_events beat task. Replay " "only needs streams a client could still be tailing." @@ -66,14 +67,15 @@ class EventsSettings(SettingsGroup): # Remote Device feature. REMOTE_DEVICE_SESSION_IDLE_SECONDS: int = Field( - default=60, description="Seconds without a heartbeat before a remote-device session is considered idle." + default=60, gt=0, description="Seconds without a heartbeat before a remote-device session is considered idle." ) REMOTE_DEVICE_REQUIRE_SIGNATURE: bool = Field( default=False, description="Require signed commands from remote devices." ) - REMOTE_DEVICE_PAIRING_TTL_SECONDS: int = Field(default=600, description="Lifetime of a pairing code.") + REMOTE_DEVICE_PAIRING_TTL_SECONDS: int = Field(default=600, gt=0, description="Lifetime of a pairing code.") REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS: int = Field( default=900, + gt=605, description=( "Redis TTL of the per-device command queue, routing invocations cross-process so a scheduled run " "reaches the web-held device session. Must exceed the max drain deadline (605s) so a command for a " @@ -81,7 +83,7 @@ class EventsSettings(SettingsGroup): ), ) REMOTE_DEVICE_INVOCATION_TTL_SECONDS: int = Field( - default=900, description="Redis TTL of a pending remote-device invocation." + default=900, gt=0, description="Redis TTL of a pending remote-device invocation." ) REMOTE_DEVICE_OUTPUT_STREAM_MAXLEN: int = Field( default=10_000, description="Cap on buffered output entries per remote-device invocation stream." diff --git a/docsgpt/core/settings/llm.py b/docsgpt/core/settings/llm.py index 11c17cac..173c864a 100644 --- a/docsgpt/core/settings/llm.py +++ b/docsgpt/core/settings/llm.py @@ -5,10 +5,10 @@ from __future__ import annotations import os from typing import Optional -from pydantic import Field, field_validator +from pydantic import Field from docsgpt.core.paths import home_dir -from docsgpt.core.settings._shared import SettingsGroup, normalize_secret +from docsgpt.core.settings._shared import SettingsGroup class LLMSettings(SettingsGroup): @@ -104,18 +104,3 @@ class LLMSettings(SettingsGroup): OPENAI_REASONING_SUMMARY: str = Field( default="auto", description="Reasoning summary mode requested from the Responses API." ) - - @field_validator( - "API_KEY", - "OPENAI_API_KEY", - "ANTHROPIC_API_KEY", - "GOOGLE_API_KEY", - "GROQ_API_KEY", - "HUGGINGFACE_API_KEY", - "NOVITA_API_KEY", - "FALLBACK_LLM_API_KEY", - mode="before", - ) - @classmethod - def _normalize_llm_secrets(cls, v): - return normalize_secret(v) diff --git a/docsgpt/core/settings/reference.py b/docsgpt/core/settings/reference.py index 02d3721d..e5aa3392 100644 --- a/docsgpt/core/settings/reference.py +++ b/docsgpt/core/settings/reference.py @@ -62,7 +62,9 @@ def _type_name(annotation: Any) -> str: if origin is Literal: return " | ".join(json.dumps(v) for v in typing.get_args(annotation)) if origin is not None: - return getattr(origin, "__name__", str(origin)) + name = getattr(origin, "__name__", str(origin)) + args = typing.get_args(annotation) + return f"{name}[{', '.join(_type_name(arg) for arg in args)}]" if args else name return getattr(annotation, "__name__", str(annotation)) diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py index 8c716646..651d3f3d 100644 --- a/docsgpt/core/settings/retrieval.py +++ b/docsgpt/core/settings/retrieval.py @@ -29,7 +29,7 @@ class RetrievalSettings(SettingsGroup): default=None, description="Model for ingest-time graph extraction; unset reuses LLM_PROVIDER/LLM_NAME." ) GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field( - default=2000, description="Hard cap on chunks extracted per source (cost control)." + default=2000, ge=0, description="Hard cap on chunks extracted per source (cost control); 0 extracts nothing." ) @field_validator("VECTOR_STORE", mode="before") diff --git a/docsgpt/core/settings/scheduler.py b/docsgpt/core/settings/scheduler.py index 00a9f455..b86d9995 100644 --- a/docsgpt/core/settings/scheduler.py +++ b/docsgpt/core/settings/scheduler.py @@ -25,4 +25,4 @@ class SchedulerSettings(SettingsGroup): SCHEDULE_ONCE_MAX_HORIZON: int = Field( default=31_536_000, description="How far ahead a one-off run may be scheduled, in seconds (one year)." ) - SCHEDULE_RUN_OUTPUT_RETENTION_DAYS: int = Field(default=90, description="Days scheduled-run output is kept.") + SCHEDULE_RUN_OUTPUT_RETENTION_DAYS: int = Field(default=90, gt=0, description="Days scheduled-run output is kept.") diff --git a/docsgpt/core/settings/speech.py b/docsgpt/core/settings/speech.py index da5c245e..f25380b0 100644 --- a/docsgpt/core/settings/speech.py +++ b/docsgpt/core/settings/speech.py @@ -6,7 +6,7 @@ from typing import Literal, Optional from pydantic import Field, field_validator -from docsgpt.core.settings._shared import SettingsGroup, normalize_choice, normalize_secret +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice class SpeechSettings(SettingsGroup): @@ -25,11 +25,6 @@ class SpeechSettings(SettingsGroup): STT_ENABLE_TIMESTAMPS: bool = Field(default=False, description="Return word/segment timestamps.") STT_ENABLE_DIARIZATION: bool = Field(default=False, description="Label speakers in the transcript.") - @field_validator("ELEVENLABS_API_KEY", mode="before") - @classmethod - def _normalize_speech_secrets(cls, v): - return normalize_secret(v) - @field_validator("TTS_PROVIDER", "STT_PROVIDER", mode="before") @classmethod def _normalize_speech_providers(cls, v): diff --git a/docsgpt/core/settings/vectorstores.py b/docsgpt/core/settings/vectorstores.py index 067d1df0..010eefa3 100644 --- a/docsgpt/core/settings/vectorstores.py +++ b/docsgpt/core/settings/vectorstores.py @@ -8,7 +8,7 @@ from pydantic import Field, field_validator from docsgpt.core.db_uri import normalize_pgvector_connection_string from docsgpt.core.paths import home_dir -from docsgpt.core.settings._shared import SettingsGroup, normalize_secret +from docsgpt.core.settings._shared import SettingsGroup class VectorStoreSettings(SettingsGroup): @@ -82,8 +82,3 @@ class VectorStoreSettings(SettingsGroup): @classmethod def _normalize_pgvector_connection_string(cls, v): return normalize_pgvector_connection_string(v) - - @field_validator("QDRANT_API_KEY", mode="before") - @classmethod - def _normalize_vectorstore_secrets(cls, v): - return normalize_secret(v) diff --git a/docsgpt/core/settings/workers.py b/docsgpt/core/settings/workers.py index 86246cab..96efc8f2 100644 --- a/docsgpt/core/settings/workers.py +++ b/docsgpt/core/settings/workers.py @@ -17,6 +17,7 @@ class WorkerSettings(SettingsGroup): ) CELERY_VISIBILITY_TIMEOUT: int = Field( default=3600, + gt=0, description=( "Broker visibility timeout in seconds. Must exceed the longest legitimate task runtime but stay " "short enough that SIGKILLed tasks redeliver promptly." diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py index df0e0e07..db78a45b 100644 --- a/tests/core/test_settings.py +++ b/tests/core/test_settings.py @@ -6,6 +6,8 @@ validators from every group applied) and that the generated reference page tracks the definitions. """ +import types +import typing import warnings from pathlib import Path @@ -23,11 +25,22 @@ SECRET_FIELDS = ( "GROQ_API_KEY", "HUGGINGFACE_API_KEY", "NOVITA_API_KEY", + "OPEN_ROUTER_API_KEY", "EMBEDDINGS_KEY", "FALLBACK_LLM_API_KEY", "QDRANT_API_KEY", + "ELASTIC_PASSWORD", "ELEVENLABS_API_KEY", "INTERNAL_KEY", + "SCIM_TOKEN", + "OIDC_ISSUER", + "GITHUB_ACCESS_TOKEN", + "MICROSOFT_AUTHORITY", + "MCP_OAUTH_REDIRECT_URI", + "S3_ACCESS_KEY_ID", + "S3_SECRET_ACCESS_KEY", + "SANDBOX_GATEWAY_AUTH_TOKEN", + "DAYTONA_API_KEY", ) @@ -69,10 +82,27 @@ class TestComposition: class TestValidators: """Validators live on the group that owns the field; composition must keep all of them. - Two groups defining a validator under the same method name would silently - keep only one, so this checks every secret field, across every group. + Pydantic collects validators by method name across the MRO, so two groups + defining one under the same name would silently keep only one; the checks + below span every group. """ + def test_every_optional_string_treats_unset_spellings_as_none(self): + names = [ + name + for name, field in Settings.model_fields.items() + if typing.get_origin(field.annotation) in (typing.Union, types.UnionType) + and set(typing.get_args(field.annotation)) == {str, type(None)} + ] + assert len(names) > 60 + loaded = Settings.model_validate({name: " None " for name in names}) + with warnings.catch_warnings(): + warnings.simplefilter("ignore", DeprecationWarning) + assert [name for name in names if getattr(loaded, name) is not None] == [] + + def test_plain_strings_keep_empty_values(self): + assert Settings.model_validate({"MILVUS_TOKEN": "", "JWT_SECRET_KEY": ""}).MILVUS_TOKEN == "" + @pytest.mark.parametrize("name", SECRET_FIELDS) @pytest.mark.parametrize("raw", ["None", "none", "", " "]) def test_unset_secret_spellings_become_none(self, name, raw): @@ -125,6 +155,11 @@ class TestCrossFieldRules: with pytest.raises(ValidationError, match="AUTH_TYPE=oidc requires settings: OIDC_CLIENT_ID, OIDC_FRONTEND_URL"): Settings.model_validate({"AUTH_TYPE": "oidc", "OIDC_ISSUER": self.OIDC["OIDC_ISSUER"]}) + @pytest.mark.parametrize("raw", ["", "None", " "]) + def test_oidc_unset_spellings_do_not_satisfy_the_requirement(self, raw): + with pytest.raises(ValidationError, match="OIDC_CLIENT_ID"): + Settings.model_validate({"AUTH_TYPE": "oidc", **self.OIDC, "OIDC_CLIENT_ID": raw}) + def test_oidc_with_required_settings_loads(self): assert Settings.model_validate({"AUTH_TYPE": "OIDC", **self.OIDC}).AUTH_TYPE == "oidc" @@ -154,6 +189,7 @@ class TestClosedChoices: ("TTS_PROVIDER", "ElevenLabs", "elevenlabs"), ("STT_PROVIDER", "", "none"), ("TTS_PROVIDER", "NONE", "none"), + ("EMBEDDINGS_POOLING", "CLS", "cls"), ], ) def test_choices_are_case_insensitive(self, name, raw, expected): @@ -169,6 +205,7 @@ class TestClosedChoices: ("SANDBOX_BACKEND", "docker"), ("DOC_PARSER_ENGINE", "fast"), ("STT_PROVIDER", "whisper"), + ("EMBEDDINGS_POOLING", "max"), ], ) def test_unknown_choice_is_rejected(self, name, raw): @@ -177,7 +214,14 @@ class TestClosedChoices: @pytest.mark.parametrize( ("name", "raw"), - [("EMBEDDINGS_BATCH_SIZE", 0), ("COMPRESSION_THRESHOLD_PERCENTAGE", 1.5), ("UPLOAD_MAX_FILE_BYTES", 0)], + [ + ("EMBEDDINGS_BATCH_SIZE", 0), + ("COMPRESSION_THRESHOLD_PERCENTAGE", 1.5), + ("UPLOAD_MAX_FILE_BYTES", 0), + ("MESSAGE_EVENTS_RETENTION_DAYS", 0), + ("REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS", 605), + ("GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION", -1), + ], ) def test_out_of_range_numbers_are_rejected(self, name, raw): with pytest.raises(ValidationError): From 731baa7d31c145698c8ffd55876e21c02c9cb3d4 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 11:43:06 +0100 Subject: [PATCH 059/130] test: cover what doctor actually tells you The checks were mocked wholesale, so the branching that produces each diagnosis had never run: a database with no schema yet, one behind this version, one that refuses the connection, and which of the three Redis URLs failed. Each of those is the sentence a developer reads when something is wrong, so each is pinned. _migration_head is tested against the packaged alembic.ini itself: it needs no database, and it is the path resolution that breaks silently when files move. --- tests/deploy/test_doctor.py | 118 ++++++++++++++++++++++++++++++++++++ 1 file changed, 118 insertions(+) diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 1309d512..b92104ce 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -128,6 +128,124 @@ class TestChecks: assert commands._chosen_services(names, [names[0]]) == [names[0]] +class FakeCursor: + """Answers doctor's three queries in order: server version, whether the table is there, the revision.""" + + def __init__(self, version, table, revision): + self.answers = [(version,), (table,), (revision,) if revision else None] + self.given = 0 + + def execute(self, statement): + self.statement = statement + + def fetchone(self): + answer = self.answers[self.given] + self.given += 1 + return answer + + def __enter__(self): + return self + + def __exit__(self, *exception): + return False + + +class FakeConnection: + def __init__(self, cursor): + self._cursor = cursor + + def cursor(self): + return self._cursor + + def __enter__(self): + return self + + def __exit__(self, *exception): + return False + + +def _postgres_answering(monkeypatch, version="16.2", table="alembic_version", revision="0031_x", head="0031_x"): + import psycopg + + monkeypatch.setattr(psycopg, "connect", lambda *a, **k: FakeConnection(FakeCursor(version, table, revision))) + monkeypatch.setattr(commands, "_migration_head", lambda: head) + + +class TestPostgresCheck: + def test_at_head(self, monkeypatch): + _postgres_answering(monkeypatch) + check = commands._check_postgres("postgresql://localhost/d") + assert check.level == "ok" + assert "16.2" in check.detail and "0031_x" in check.detail + + def test_a_database_with_no_schema_yet(self, monkeypatch): + """The commonest first-run state: the database exists, nothing has been migrated into it.""" + _postgres_answering(monkeypatch, table=None, revision=None) + check = commands._check_postgres("postgresql://localhost/d") + assert check.level == "fail" + assert "docsgpt migrate" in check.detail + + def test_a_schema_behind_this_version(self, monkeypatch): + _postgres_answering(monkeypatch, revision="0029_old", head="0031_x") + check = commands._check_postgres("postgresql://localhost/d") + assert check.level == "fail" + assert "0029_old" in check.detail and "0031_x" in check.detail + assert "docsgpt migrate" in check.detail + + def test_a_database_that_does_not_answer(self, monkeypatch): + import psycopg + + def refuse(*args, **kwargs): + raise psycopg.OperationalError("connection refused") + + monkeypatch.setattr(psycopg, "connect", refuse) + check = commands._check_postgres("postgresql://localhost/d") + assert check.level == "fail" + assert "connection refused" in check.detail + + def test_without_a_uri_at_all(self): + check = commands._check_postgres(None) + assert check.level == "fail" + assert "POSTGRES_URI" in check.detail + + +class TestRedisCheck: + def test_every_database_answering(self, monkeypatch): + import redis + + monkeypatch.setattr(redis.Redis, "from_url", classmethod(lambda cls, url, **k: type("R", (), {"ping": lambda self: True})())) + check = commands._check_redis({"broker": "redis://localhost:6379/0", "cache": "redis://localhost:6379/2"}) + assert check.level == "ok" + assert "2" in check.detail + + def test_the_one_that_does_not_answer_is_named(self, monkeypatch): + """Three URLs usually differ only by database number, so the message has to say which.""" + import redis + + def from_url(cls, url, **kwargs): + class Client: + def ping(self): + raise ConnectionError(f"no route to {url}") + + return Client() + + monkeypatch.setattr(redis.Redis, "from_url", classmethod(from_url)) + check = commands._check_redis({"cache": "redis://localhost:6379/2"}) + assert check.level == "fail" + assert "cache" in check.detail and "6379/2" in check.detail + + def test_without_any_redis_configured(self): + assert commands._check_redis({}).level == "fail" + + +class TestMigrationHead: + def test_it_finds_the_revision_this_package_ships(self): + """No database needed: this is the alembic.ini path resolution, which breaks silently.""" + head = commands._migration_head() + assert head, "the packaged alembic.ini should resolve to a revision" + assert head[0].isdigit(), head + + class TestDoctor: def _only(self, monkeypatch, postgres, redis): monkeypatch.setattr(commands, "_check_postgres", lambda uri: postgres) From bcf2707efa3dcbba5be6e2ab3e516f5c0fecf68f Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 11:43:53 +0100 Subject: [PATCH 060/130] fix: doctor reports a bad DOCSGPT_PORT instead of raising on it _port_number was written so a hand-edited .env could not reach int() raw, and then doctor did exactly that: a nonnumeric port ended the command with a traceback rather than the message, in the one command whose job is to explain a broken setup. --- docsgpt/deploy/commands.py | 4 +++- tests/deploy/test_doctor.py | 11 +++++++++++ 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 8b1503a6..0811134f 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -862,7 +862,9 @@ def doctor(args, context: Optional[Context] = None) -> int: _check_provider(env), ] - port = int(env.get("DOCSGPT_PORT") or stack.DEFAULT_PORT) + # The one command that exists to explain a broken setup must not fall over on one. + port = _port_number(env["DOCSGPT_PORT"], f"DOCSGPT_PORT in {env_path}") if env.get("DOCSGPT_PORT") \ + else stack.DEFAULT_PORT if _port_is_free(port): checks.append(Check("port", "ok", f"{port} is free")) else: diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index b92104ce..3cc94dfd 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -274,6 +274,17 @@ class TestDoctor: assert "FAIL" in captured.out assert "1 problem" in captured.err + def test_a_hand_edited_port_is_reported_not_raised(self, tmp_path, monkeypatch): + """Doctor exists to explain a broken setup, so it must not traceback on one.""" + self._only( + monkeypatch, + commands.Check("postgres", "ok", "fine"), + commands.Check("redis", "ok", "fine"), + ) + (tmp_path / ".env").write_text("DOCSGPT_PORT=seven thousand\n", encoding="utf-8") + with pytest.raises(DeployError, match="not a port number"): + _run(["doctor", "--dir", str(tmp_path)], _context()) + def test_it_says_which_settings_file_it_read(self, tmp_path, capsys, monkeypatch): self._only( monkeypatch, From 41135fc894ac6d5d254c2446bd3fbde3f43d8257 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:51:28 +0100 Subject: [PATCH 061/130] refactor(settings): apply the unset rule to optional string Literals too Review follow-up. EMBEDDINGS_POOLING became Optional[Literal["cls", "mean"]], but the group-base rule that maps "", "None" and whitespace to None only matched Optional[str], so EMBEDDINGS_POOLING= in a .env file would have failed validation. The rule now also covers an Optional Literal whose choices are all strings, which lets the AUTH_TYPE validator drop its own copy of that handling. --- docsgpt/core/settings/_shared.py | 13 +++++++++++-- docsgpt/core/settings/auth.py | 6 +++--- tests/core/test_settings.py | 9 ++++++--- 3 files changed, 20 insertions(+), 8 deletions(-) diff --git a/docsgpt/core/settings/_shared.py b/docsgpt/core/settings/_shared.py index 59723bbf..5b8bfec4 100644 --- a/docsgpt/core/settings/_shared.py +++ b/docsgpt/core/settings/_shared.py @@ -18,15 +18,24 @@ from pydantic_settings import BaseSettings, SettingsConfigDict def _is_optional_str(annotation: Any) -> bool: + """``Optional[str]`` or ``Optional[Literal[...]]`` whose choices are all strings.""" if typing.get_origin(annotation) not in (typing.Union, types.UnionType): return False - return set(typing.get_args(annotation)) == {str, type(None)} + members = set(typing.get_args(annotation)) + if type(None) not in members or len(members) != 2: + return False + (member,) = members - {type(None)} + if member is str: + return True + return typing.get_origin(member) is typing.Literal and all( + isinstance(choice, str) for choice in typing.get_args(member) + ) class SettingsGroup(BaseSettings): """Base for one domain's settings; groups are composed into ``Settings``. - Every ``Optional[str]`` field treats the spellings an unset value has in a + Every ``Optional[str]`` field (and optional string ``Literal``) treats the spellings an unset value has in a ``.env`` file (``KEY=``, ``KEY=None``, whitespace) as ``None``, so a check like ``if settings.OIDC_ISSUER`` or a fallback like ``settings.X or default`` sees "unset" rather than a truthy placeholder string. Real values are diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py index c9309775..5069d6cf 100644 --- a/docsgpt/core/settings/auth.py +++ b/docsgpt/core/settings/auth.py @@ -6,7 +6,7 @@ from typing import Literal, Optional from pydantic import Field, field_validator, model_validator -from docsgpt.core.settings._shared import SettingsGroup, normalize_choice, normalize_secret +from docsgpt.core.settings._shared import SettingsGroup, normalize_choice #: Settings an OIDC deployment cannot run without; checked when AUTH_TYPE=oidc. @@ -88,8 +88,8 @@ class AuthSettings(SettingsGroup): @field_validator("AUTH_TYPE", mode="before") @classmethod def _normalize_auth_type(cls, v): - # ``AUTH_TYPE=None`` and ``AUTH_TYPE=`` in .env both mean "no authentication". - return normalize_choice(normalize_secret(v)) + # Unset spellings ("None", "") became None on the group base; this only case-folds a value. + return normalize_choice(v) @model_validator(mode="after") def _require_oidc_settings(self): diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py index db78a45b..ff53fba4 100644 --- a/tests/core/test_settings.py +++ b/tests/core/test_settings.py @@ -92,9 +92,11 @@ class TestValidators: name for name, field in Settings.model_fields.items() if typing.get_origin(field.annotation) in (typing.Union, types.UnionType) - and set(typing.get_args(field.annotation)) == {str, type(None)} + and type(None) in typing.get_args(field.annotation) + and all(a is str or typing.get_origin(a) is typing.Literal for a in typing.get_args(field.annotation) if a is not type(None)) ] assert len(names) > 60 + assert {"EMBEDDINGS_POOLING", "AUTH_TYPE", "OIDC_ISSUER"} <= set(names) loaded = Settings.model_validate({name: " None " for name in names}) with warnings.catch_warnings(): warnings.simplefilter("ignore", DeprecationWarning) @@ -171,9 +173,10 @@ class TestCrossFieldRules: class TestClosedChoices: """Enum-like settings are Literal types: a typo fails at startup instead of falling through.""" + @pytest.mark.parametrize("name", ["AUTH_TYPE", "EMBEDDINGS_POOLING"]) @pytest.mark.parametrize("raw", ["None", "none", "", " "]) - def test_auth_type_unset_spellings(self, raw): - assert Settings.model_validate({"AUTH_TYPE": raw}).AUTH_TYPE is None + def test_optional_choice_unset_spellings(self, name, raw): + assert getattr(Settings.model_validate({name: raw}), name) is None @pytest.mark.parametrize( ("name", "raw", "expected"), From 9326c6edb1620369c2b9603016ddede6dffa152f Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 11:52:28 +0100 Subject: [PATCH 062/130] fix: keep credentials out of doctor's output, and read the revision from the schema it checked The Redis check put the whole URL in its failure message. A managed Redis URL carries user:password@host, so a failed ping printed the password to the terminal and into any log or issue the output was pasted into. It names scheme://host:port/db now, and falls back to naming no URL when the value cannot be parsed at all. The Postgres check asked to_regclass about public.alembic_version and then read version_num through search_path, so another schema could answer with a different revision, or the query could fail, and doctor would send you to run migrations against a database that is already fine. --- docsgpt/deploy/commands.py | 17 +++++++++++++++-- tests/deploy/test_doctor.py | 35 +++++++++++++++++++++++++++++++++-- 2 files changed, 48 insertions(+), 4 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 0811134f..2da4929a 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -796,7 +796,8 @@ def _check_postgres(uri: Optional[str]) -> Check: applied = cursor.fetchone()[0] is not None current = None if applied: - cursor.execute("select version_num from alembic_version") + # public, like the to_regclass check above: search_path could resolve another one. + cursor.execute("select version_num from public.alembic_version") row = cursor.fetchone() current = row[0] if row else None except (psycopg.Error, OSError, ValueError) as exc: @@ -810,6 +811,18 @@ def _check_postgres(uri: Optional[str]) -> Check: return Check("postgres", "ok", f"PostgreSQL {version}, schema at {current}") +def _redis_endpoint(url: str) -> str: + """A Redis URL without its credentials: this ends up on a terminal, in CI logs and in issues.""" + try: + parts = urlsplit(url) + except ValueError: + return "the configured URL" + host = parts.hostname or "" + if parts.port: + host = f"{host}:{parts.port}" + return f"{parts.scheme}://{host}{parts.path}" if host else "the configured URL" + + def _check_redis(urls: Mapping[str, str]) -> Check: """Ping every Redis the settings name; they are usually one server, three databases.""" if not urls: @@ -822,7 +835,7 @@ def _check_redis(urls: Mapping[str, str]) -> Check: try: redis.Redis.from_url(url, socket_connect_timeout=3).ping() except Exception as exc: # noqa: BLE001 - every client error here is the same finding - return Check("redis", "fail", f"{label} ({url}) does not answer: {str(exc).strip()}") + return Check("redis", "fail", f"{label} ({_redis_endpoint(url)}) does not answer: {str(exc).strip()}") return Check("redis", "ok", f"answering on {len(urls)} database(s)") diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 3cc94dfd..9c73b77e 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -134,9 +134,10 @@ class FakeCursor: def __init__(self, version, table, revision): self.answers = [(version,), (table,), (revision,) if revision else None] self.given = 0 + self.statements = [] def execute(self, statement): - self.statement = statement + self.statements.append(statement) def fetchone(self): answer = self.answers[self.given] @@ -167,8 +168,10 @@ class FakeConnection: def _postgres_answering(monkeypatch, version="16.2", table="alembic_version", revision="0031_x", head="0031_x"): import psycopg - monkeypatch.setattr(psycopg, "connect", lambda *a, **k: FakeConnection(FakeCursor(version, table, revision))) + cursor = FakeCursor(version, table, revision) + monkeypatch.setattr(psycopg, "connect", lambda *a, **k: FakeConnection(cursor)) monkeypatch.setattr(commands, "_migration_head", lambda: head) + return cursor class TestPostgresCheck: @@ -178,6 +181,13 @@ class TestPostgresCheck: assert check.level == "ok" assert "16.2" in check.detail and "0031_x" in check.detail + def test_the_revision_comes_from_the_schema_that_was_checked(self, monkeypatch): + """to_regclass looks in public, so the second query must not resolve through search_path.""" + cursor = _postgres_answering(monkeypatch) + commands._check_postgres("postgresql://localhost/d") + revision_query = [statement for statement in cursor.statements if "version_num" in statement] + assert revision_query and all("public.alembic_version" in statement for statement in revision_query) + def test_a_database_with_no_schema_yet(self, monkeypatch): """The commonest first-run state: the database exists, nothing has been migrated into it.""" _postgres_answering(monkeypatch, table=None, revision=None) @@ -234,6 +244,27 @@ class TestRedisCheck: assert check.level == "fail" assert "cache" in check.detail and "6379/2" in check.detail + def test_a_password_in_the_url_is_never_printed(self, monkeypatch): + """doctor output goes into terminals, CI logs and pasted issue reports.""" + import redis + + def from_url(cls, url, **kwargs): + class Client: + def ping(self): + raise ConnectionError("connection refused") + + return Client() + + monkeypatch.setattr(redis.Redis, "from_url", classmethod(from_url)) + check = commands._check_redis({"broker": "rediss://default:sUpErSeCrEt@redis.example.com:6380/0"}) + assert check.level == "fail" + assert "sUpErSeCrEt" not in check.detail + assert "default" not in check.detail + assert "redis.example.com:6380/0" in check.detail, "the endpoint still has to be identifiable" + + def test_a_malformed_url_is_not_echoed_either(self): + assert commands._redis_endpoint("redis://[::1") == "the configured URL" + def test_without_any_redis_configured(self): assert commands._check_redis({}).level == "fail" From 18313c782371b9067635cb54d20e441fa5643330 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 12:01:28 +0100 Subject: [PATCH 063/130] fix: keep credentials out of the text doctor borrows from its clients Sanitising the URL in the message was not enough: the client's own error text went into the detail too, and both psycopg and redis-py quote the URL they were given. The reason is kept, the endpoint is kept, and the URL, username and password are taken out of it. The Redis test raised a generic error, so it asserted the password was absent without ever exercising the path that leaked. Both tests now raise what the clients actually raise. --- docsgpt/deploy/commands.py | 25 +++++++++++++++++++++---- tests/deploy/test_doctor.py | 21 +++++++++++++++++++-- 2 files changed, 40 insertions(+), 6 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 2da4929a..a62bee78 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -801,7 +801,7 @@ def _check_postgres(uri: Optional[str]) -> Check: row = cursor.fetchone() current = row[0] if row else None except (psycopg.Error, OSError, ValueError) as exc: - return Check("postgres", "fail", f"cannot connect: {str(exc).strip()}") + return Check("postgres", "fail", f"cannot connect to {_endpoint(uri)}: {_scrub(str(exc).strip(), uri)}") head = _migration_head() if not current: return Check("postgres", "fail", f"PostgreSQL {version}, no schema yet; run `docsgpt migrate`") @@ -811,8 +811,8 @@ def _check_postgres(uri: Optional[str]) -> Check: return Check("postgres", "ok", f"PostgreSQL {version}, schema at {current}") -def _redis_endpoint(url: str) -> str: - """A Redis URL without its credentials: this ends up on a terminal, in CI logs and in issues.""" +def _endpoint(url: str) -> str: + """A URL without its credentials: this ends up on a terminal, in CI logs and in issues.""" try: parts = urlsplit(url) except ValueError: @@ -823,6 +823,21 @@ def _redis_endpoint(url: str) -> str: return f"{parts.scheme}://{host}{parts.path}" if host else "the configured URL" +def _scrub(text: str, url: Optional[str]) -> str: + """Client errors quote the URL they were handed, credentials and all, so take them back out.""" + if not url: + return text + text = text.replace(url, _endpoint(url)) + try: + parts = urlsplit(url) + except ValueError: + return text + for secret in (parts.password, parts.username): + if secret: + text = text.replace(secret, "...") + return text + + def _check_redis(urls: Mapping[str, str]) -> Check: """Ping every Redis the settings name; they are usually one server, three databases.""" if not urls: @@ -835,7 +850,9 @@ def _check_redis(urls: Mapping[str, str]) -> Check: try: redis.Redis.from_url(url, socket_connect_timeout=3).ping() except Exception as exc: # noqa: BLE001 - every client error here is the same finding - return Check("redis", "fail", f"{label} ({_redis_endpoint(url)}) does not answer: {str(exc).strip()}") + return Check( + "redis", "fail", f"{label} ({_endpoint(url)}) does not answer: {_scrub(str(exc).strip(), url)}" + ) return Check("redis", "ok", f"answering on {len(urls)} database(s)") diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 9c73b77e..8e54efc3 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -213,6 +213,22 @@ class TestPostgresCheck: assert check.level == "fail" assert "connection refused" in check.detail + def test_a_database_password_is_never_printed(self, monkeypatch): + """psycopg quotes the connection string it was given, which carries the password.""" + import psycopg + + uri = "postgresql://docsgpt:hunter2@db.example.com:5432/docsgpt" + + def refuse(*args, **kwargs): + raise psycopg.OperationalError(f'connection to "{uri}" failed: timeout expired') + + monkeypatch.setattr(psycopg, "connect", refuse) + check = commands._check_postgres(uri) + assert check.level == "fail" + assert "hunter2" not in check.detail + assert "db.example.com:5432" in check.detail, "the endpoint still has to be identifiable" + assert "timeout expired" in check.detail, "and the reason has to survive the scrubbing" + def test_without_a_uri_at_all(self): check = commands._check_postgres(None) assert check.level == "fail" @@ -251,7 +267,8 @@ class TestRedisCheck: def from_url(cls, url, **kwargs): class Client: def ping(self): - raise ConnectionError("connection refused") + # redis-py quotes the URL it was handed, which is how the credentials got out. + raise ConnectionError(f"no route to {url}") return Client() @@ -263,7 +280,7 @@ class TestRedisCheck: assert "redis.example.com:6380/0" in check.detail, "the endpoint still has to be identifiable" def test_a_malformed_url_is_not_echoed_either(self): - assert commands._redis_endpoint("redis://[::1") == "the configured URL" + assert commands._endpoint("redis://[::1") == "the configured URL" def test_without_any_redis_configured(self): assert commands._check_redis({}).level == "fail" From 1cce6cfd673c5f00627d6fdbab81b5a4333e954d Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 12:02:40 +0100 Subject: [PATCH 064/130] test: cover the two things dev refuses to do Running outside a checkout and starting on a port something else holds are both refusals a developer will meet, and neither had a test. The second also asserts that nothing is spawned when the port is taken, which is the part that matters: the guard runs before any child process exists. --- tests/deploy/test_doctor.py | 49 +++++++++++++++++++++++++++++++++++++ 1 file changed, 49 insertions(+) diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 8e54efc3..185f8c4f 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -1,5 +1,7 @@ """`docsgpt doctor`, `restart`, following native logs, and settings that apply themselves.""" +import socket + import pytest from docsgpt.deploy import commands, envfile @@ -14,6 +16,53 @@ def _installed_native(tmp_path, services): assert _run(argv, _native_context(services)) == 0 +class TestDevCommand: + def _checkout(self, monkeypatch, tmp_path): + from docsgpt.core import paths + + monkeypatch.setattr(paths, "checkout_root", lambda: tmp_path) + + def _free_port(self): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: + probe.bind(("127.0.0.1", 0)) + return probe.getsockname()[1] + + def test_an_installed_package_is_pointed_at_native_mode(self, monkeypatch): + """`dev` runs the code you are editing; there is none to edit outside a checkout.""" + from docsgpt.core import paths + + monkeypatch.setattr(paths, "checkout_root", lambda: None) + with pytest.raises(DeployError, match="source checkout"): + _run(["dev"], _context()) + + def test_a_busy_port_is_refused_before_anything_starts(self, monkeypatch, tmp_path): + from docsgpt.deploy import dev as dev_module + + self._checkout(monkeypatch, tmp_path) + started = [] + monkeypatch.setattr(dev_module, "run", lambda children, **kwargs: started.append(children) or 0) + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as held: + held.bind(("127.0.0.1", 0)) + held.listen(1) + port = held.getsockname()[1] + with pytest.raises(DeployError, match="already in use"): + _run(["dev", "--port", str(port)], _context()) + assert started == [], "nothing is spawned when the port is taken" + + def test_it_runs_the_children_it_planned_and_says_where(self, monkeypatch, tmp_path, capsys): + from docsgpt.deploy import dev as dev_module + + self._checkout(monkeypatch, tmp_path) + started = [] + monkeypatch.setattr(dev_module, "run", lambda children, **kwargs: started.append(children) or 0) + port = self._free_port() + assert _run(["dev", "--port", str(port)], _context()) == 0 + assert [child.name for child in started[0]] == ["api", "worker"] + printed = capsys.readouterr().out + assert f"http://127.0.0.1:{port}" in printed + assert "Ctrl-C" in printed + + class TestRestart: def test_it_stops_and_starts_both_without_touching_settings(self, tmp_path): services = FakeServices() From 66fbb110492113b488a98b26dc62d59de1cc8f43 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 12:12:28 +0100 Subject: [PATCH 065/130] fix: doctor survives a Redis URL whose port will not parse urlsplit accepts redis://host:notaport/0; only parts.port raises, and it raises on access rather than at split time, so the ValueError fell outside the try. _check_redis catches the client error and then formats it through _endpoint, so doctor ended with a traceback from inside its own error path. --- docsgpt/deploy/commands.py | 9 +++++---- tests/deploy/test_doctor.py | 22 ++++++++++++++++++++-- 2 files changed, 25 insertions(+), 6 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index a62bee78..758fc715 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -815,12 +815,13 @@ def _endpoint(url: str) -> str: """A URL without its credentials: this ends up on a terminal, in CI logs and in issues.""" try: parts = urlsplit(url) + # .port is a property that parses on access, so it raises separately from the split itself. + host, port, scheme, path = parts.hostname or "", parts.port, parts.scheme, parts.path except ValueError: return "the configured URL" - host = parts.hostname or "" - if parts.port: - host = f"{host}:{parts.port}" - return f"{parts.scheme}://{host}{parts.path}" if host else "the configured URL" + if port: + host = f"{host}:{port}" + return f"{scheme}://{host}{path}" if host else "the configured URL" def _scrub(text: str, url: Optional[str]) -> str: diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 185f8c4f..ac046a85 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -328,8 +328,26 @@ class TestRedisCheck: assert "default" not in check.detail assert "redis.example.com:6380/0" in check.detail, "the endpoint still has to be identifiable" - def test_a_malformed_url_is_not_echoed_either(self): - assert commands._endpoint("redis://[::1") == "the configured URL" + @pytest.mark.parametrize("url", ["redis://[::1", "redis://localhost:not-a-port/0"]) + def test_a_malformed_url_is_not_echoed_either(self, url): + """urlsplit raises on the first; on the second it succeeds and .port raises on access.""" + assert commands._endpoint(url) == "the configured URL" + + def test_a_failure_on_a_url_with_an_unparsable_port_is_still_a_check(self, monkeypatch): + """The check catches the client error and then formats it, which is where this raised.""" + import redis + + def from_url(cls, url, **kwargs): + class Client: + def ping(self): + raise ConnectionError(f"no route to {url}") + + return Client() + + monkeypatch.setattr(redis.Redis, "from_url", classmethod(from_url)) + check = commands._check_redis({"broker": "redis://localhost:not-a-port/0"}) + assert check.level == "fail", "a bad port is a finding, not a traceback" + assert "the configured URL" in check.detail def test_without_any_redis_configured(self): assert commands._check_redis({}).level == "fail" From 9e86dd4e94f4c388a6ec61fffebda1cf032d5ceb Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 12:22:45 +0100 Subject: [PATCH 066/130] fix: refuse a busy mock LLM port before dev starts anything The mock starts first and the API and worker are pointed at it, but only the API port was preflighted. A busy 8090 therefore showed up as a child exiting once the rest were running, or as the API talking to whatever else was on that port. --ui is deliberately left alone: vite.config.ts sets no strictPort, so Vite moves to the next free port rather than failing. --- docsgpt/deploy/commands.py | 7 +++++++ tests/deploy/test_doctor.py | 17 +++++++++++++++++ 2 files changed, 24 insertions(+) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 758fc715..b8f063aa 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -739,6 +739,13 @@ def dev(args, context: Optional[Context] = None) -> int: f"port {args.port} is already in use, so the API cannot bind it. Stop what is on it " f"(a previous `docsgpt dev`, or `docsgpt down` for an install), or pass --port." ) + if getattr(args, "mock_llm", False) and not _port_is_free(dev_module.MOCK_LLM_PORT): + # It starts first and the others are pointed at it, so a busy port here would surface as the + # API talking to someone else's server, or as a child exiting once everything else is up. + raise DeployError( + f"port {dev_module.MOCK_LLM_PORT} is already in use, so the mock LLM cannot bind it. " + "Stop what is on it, or leave --mock-llm off and point DocsGPT at a real provider." + ) children = dev_module.plan(args, checkout) print(f"DocsGPT from {checkout}") for child in children: diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index ac046a85..c0a52b28 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -49,6 +49,23 @@ class TestDevCommand: _run(["dev", "--port", str(port)], _context()) assert started == [], "nothing is spawned when the port is taken" + def test_a_busy_mock_llm_port_is_refused_before_anything_starts(self, monkeypatch, tmp_path): + """It starts first and the rest are pointed at it, so a clash cannot wait until spawn time.""" + from docsgpt.deploy import dev as dev_module + + self._checkout(monkeypatch, tmp_path) + (tmp_path / "scripts").mkdir() + (tmp_path / "scripts" / "mock_llm.py").write_text("", encoding="utf-8") + started = [] + monkeypatch.setattr(dev_module, "run", lambda children, **kwargs: started.append(children) or 0) + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as held: + held.bind(("127.0.0.1", 0)) + held.listen(1) + monkeypatch.setattr(dev_module, "MOCK_LLM_PORT", held.getsockname()[1]) + with pytest.raises(DeployError, match="mock LLM cannot bind"): + _run(["dev", "--mock-llm", "--port", str(self._free_port())], _context()) + assert started == [], "nothing is spawned when the mock LLM has nowhere to listen" + def test_it_runs_the_children_it_planned_and_says_where(self, monkeypatch, tmp_path, capsys): from docsgpt.deploy import dev as dev_module From bb50d0a91d3e985183ffefff433bbb430329ccaf Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 12:39:41 +0100 Subject: [PATCH 067/130] fix: refuse to give the API the mock LLM's port Both preflights passed for `dev --mock-llm --port 8090`: the port was free, and then the mock and the API were each handed it. One child could not bind, and the API was pointed at that port as its model server while trying to listen on it. --- docsgpt/deploy/commands.py | 6 ++++++ tests/deploy/test_doctor.py | 14 ++++++++++++++ 2 files changed, 20 insertions(+) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index b8f063aa..94a144f6 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -739,6 +739,12 @@ def dev(args, context: Optional[Context] = None) -> int: f"port {args.port} is already in use, so the API cannot bind it. Stop what is on it " f"(a previous `docsgpt dev`, or `docsgpt down` for an install), or pass --port." ) + if getattr(args, "mock_llm", False) and args.port == dev_module.MOCK_LLM_PORT: + # Both checks below would pass: the port really is free, and then two children want it. + raise DeployError( + f"port {args.port} is where the mock LLM listens, so the API cannot have it too. " + "Give the API another port with --port." + ) if getattr(args, "mock_llm", False) and not _port_is_free(dev_module.MOCK_LLM_PORT): # It starts first and the others are pointed at it, so a busy port here would surface as the # API talking to someone else's server, or as a child exiting once everything else is up. diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index c0a52b28..4acca1d5 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -66,6 +66,20 @@ class TestDevCommand: _run(["dev", "--mock-llm", "--port", str(self._free_port())], _context()) assert started == [], "nothing is spawned when the mock LLM has nowhere to listen" + def test_the_api_cannot_be_given_the_mock_llm_port(self, monkeypatch, tmp_path): + """Both free-port checks pass here: the port is free, and then two children want it.""" + from docsgpt.deploy import dev as dev_module + + self._checkout(monkeypatch, tmp_path) + (tmp_path / "scripts").mkdir() + (tmp_path / "scripts" / "mock_llm.py").write_text("", encoding="utf-8") + started = [] + monkeypatch.setattr(dev_module, "run", lambda children, **kwargs: started.append(children) or 0) + monkeypatch.setattr(dev_module, "MOCK_LLM_PORT", self._free_port()) + with pytest.raises(DeployError, match="mock LLM listens"): + _run(["dev", "--mock-llm", "--port", str(dev_module.MOCK_LLM_PORT)], _context()) + assert started == [], "nothing is spawned when the two would collide" + def test_it_runs_the_children_it_planned_and_says_where(self, monkeypatch, tmp_path, capsys): from docsgpt.deploy import dev as dev_module From 48285f4405c3604b9bde70672a88a5b88aa98c23 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 12:51:50 +0100 Subject: [PATCH 068/130] fix: doctor looks for alembic_version where alembic puts it env.py sets no version_table_schema, so the table follows search_path. Pinning public. in doctor's queries made them agree with each other and disagree with alembic: on an install using another schema it reported no schema at all and sent the user to migrate an already-migrated database. Both queries resolve the table the same way alembic does now. --- docsgpt/deploy/commands.py | 8 +++++--- tests/deploy/test_doctor.py | 13 +++++++++---- 2 files changed, 14 insertions(+), 7 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 94a144f6..ca7be581 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -805,12 +805,14 @@ def _check_postgres(uri: Optional[str]) -> Check: with psycopg.connect(uri, connect_timeout=5) as connection, connection.cursor() as cursor: cursor.execute("select current_setting('server_version')") version = cursor.fetchone()[0] - cursor.execute("select to_regclass('public.alembic_version')") + # Unqualified, like alembic itself: env.py sets no version_table_schema, so the table + # lives wherever search_path puts it. Asserting public would call a migrated database empty. + cursor.execute("select to_regclass('alembic_version')") applied = cursor.fetchone()[0] is not None current = None if applied: - # public, like the to_regclass check above: search_path could resolve another one. - cursor.execute("select version_num from public.alembic_version") + # The same relation the check above resolved, by the same rules. + cursor.execute("select version_num from alembic_version") row = cursor.fetchone() current = row[0] if row else None except (psycopg.Error, OSError, ValueError) as exc: diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 4acca1d5..4cc4a343 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -261,12 +261,17 @@ class TestPostgresCheck: assert check.level == "ok" assert "16.2" in check.detail and "0031_x" in check.detail - def test_the_revision_comes_from_the_schema_that_was_checked(self, monkeypatch): - """to_regclass looks in public, so the second query must not resolve through search_path.""" + def test_both_queries_resolve_the_same_table(self, monkeypatch): + """Alembic sets no version_table_schema, so the table follows search_path; asserting a + schema in one query and not the other is how doctor called a migrated database empty.""" cursor = _postgres_answering(monkeypatch) commands._check_postgres("postgresql://localhost/d") - revision_query = [statement for statement in cursor.statements if "version_num" in statement] - assert revision_query and all("public.alembic_version" in statement for statement in revision_query) + looked_up = [statement for statement in cursor.statements if "to_regclass" in statement] + read = [statement for statement in cursor.statements if "version_num" in statement] + assert looked_up and read + assert all("public." not in statement for statement in looked_up + read), ( + "neither query may pin a schema alembic never promised" + ) def test_a_database_with_no_schema_yet(self, monkeypatch): """The commonest first-run state: the database exists, nothing has been migrated into it.""" From d0a0b352e0dd5003a380563190a8b9de0a22e263 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 13:03:59 +0100 Subject: [PATCH 069/130] fix: drop the host-substring assertions and explain the empty except CodeQL flags `"host:port" in message` as incomplete URL sanitisation. It is a message rather than a URL being authorised, so it is not a vulnerability, but the check was red and comparing against the sanitiser's own output asserts the real contract. The signal handler now says why it swallows: the child is already gone, and shutdown must not fail on what it is cleaning up. --- docsgpt/deploy/dev.py | 2 ++ tests/deploy/test_doctor.py | 9 ++++++--- 2 files changed, 8 insertions(+), 3 deletions(-) diff --git a/docsgpt/deploy/dev.py b/docsgpt/deploy/dev.py index a49dcf07..a552646e 100644 --- a/docsgpt/deploy/dev.py +++ b/docsgpt/deploy/dev.py @@ -143,6 +143,8 @@ def _signal(process, number: int) -> None: return os.killpg(os.getpgid(process.pid), number) except (ProcessLookupError, PermissionError, OSError): + # The child has already gone, or its group is no longer ours to signal. Either way there is + # nothing left to stop, and shutdown must not fail on the thing it is trying to clean up. pass diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 4cc4a343..1a55ce84 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -311,7 +311,9 @@ class TestPostgresCheck: check = commands._check_postgres(uri) assert check.level == "fail" assert "hunter2" not in check.detail - assert "db.example.com:5432" in check.detail, "the endpoint still has to be identifiable" + # Compared against what the sanitiser produced, not a host substring: asking whether a URL + # contains a host is the check CodeQL warns about, and it is not what this test means. + assert check.detail.startswith(f"cannot connect to {commands._endpoint(uri)}") assert "timeout expired" in check.detail, "and the reason has to survive the scrubbing" def test_without_a_uri_at_all(self): @@ -358,11 +360,12 @@ class TestRedisCheck: return Client() monkeypatch.setattr(redis.Redis, "from_url", classmethod(from_url)) - check = commands._check_redis({"broker": "rediss://default:sUpErSeCrEt@redis.example.com:6380/0"}) + url = "rediss://default:sUpErSeCrEt@redis.example.com:6380/0" + check = commands._check_redis({"broker": url}) assert check.level == "fail" assert "sUpErSeCrEt" not in check.detail assert "default" not in check.detail - assert "redis.example.com:6380/0" in check.detail, "the endpoint still has to be identifiable" + assert check.detail.startswith(f"broker ({commands._endpoint(url)}) does not answer") @pytest.mark.parametrize("url", ["redis://[::1", "redis://localhost:not-a-port/0"]) def test_a_malformed_url_is_not_echoed_either(self, url): From a47a2c34ec9d0bcbfbdd67256f84ea2f01a68184 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 13:04:33 +0100 Subject: [PATCH 070/130] docs(settings): render field constraints as code spans in the reference The docs site failed to build: a bare "<= 1" in the prose of the generated page is parsed by MDX as the start of a JSX tag ("Unexpected character '=' before name"). Constraints are rendered as code spans now, where MDX leaves them alone, and a test rejects any bare <, { or } outside a code span so a future description cannot reintroduce the failure. Verified with a local next build of the docs site. --- docs/content/Deploying/Settings-Reference.mdx | 80 +++++++++---------- docsgpt/core/settings/reference.py | 3 +- tests/core/test_settings.py | 9 +++ 3 files changed, 51 insertions(+), 41 deletions(-) diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 30221972..32181adb 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -86,7 +86,7 @@ Override for the callback URL; default is <request host>/api/auth/oidc/callba ### `OIDC_SESSION_LIFETIME_SECONDS` -Type `int`, default `28800`, must be > 0. +Type `int`, default `28800`, must be `> 0`. Lifetime of the minted session JWT in seconds (8h). @@ -354,13 +354,13 @@ Truncate each remote embed input to N tokens (overflow is lost). ### `EMBEDDINGS_BATCH_SIZE` -Type `int`, default `32`, must be >= 1. +Type `int`, default `32`, must be `>= 1`. Chunks per store transaction and per remote embed request. ### `EMBEDDINGS_MODEL_BATCH_SIZE` -Type `int`, default `1`, must be >= 1. +Type `int`, default `1`, must be `>= 1`. Documents per local ONNX forward pass. Each pass pads to its longest input, and that waste grows with the square of chunk length: at 1250 tokens, 32 peaked at 6.6 GB, 1 at 2.9 GB. @@ -402,7 +402,7 @@ Celery queue the embed task is routed to. ### `EMBEDDINGS_DELEGATE_TIMEOUT` -Type `int`, default `60`, must be > 0. +Type `int`, default `60`, must be `> 0`. Seconds the API waits for the worker to return an embedding. @@ -419,7 +419,7 @@ Vector store backend. ### `RETRIEVAL_MAX_PARALLEL_SOURCES` -Type `int`, default `4`, must be >= 1. +Type `int`, default `4`, must be `>= 1`. Concurrent per-source searches in one retrieval; the query is embedded once and shared. @@ -443,7 +443,7 @@ Model for ingest-time graph extraction; unset reuses LLM_PROVIDER/LLM_NAME. ### `GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION` -Type `int`, default `2000`, must be >= 0. +Type `int`, default `2000`, must be `>= 0`. Hard cap on chunks extracted per source (cost control); 0 extracts nothing. @@ -574,7 +574,7 @@ pgvector connection string. postgres://, postgresql:// and postgresql+psycopg:// ### `PGVECTOR_POOL_MAX_SIZE` -Type `int`, default `8`, must be >= 0. +Type `int`, default `8`, must be `>= 0`. Per-process connection pool size; 0 uses one direct connection per store. @@ -668,19 +668,19 @@ Tasks prefetched per worker process; 1 caps SIGKILL loss to one task. ### `CELERY_VISIBILITY_TIMEOUT` -Type `int`, default `3600`, must be > 0. +Type `int`, default `3600`, must be `> 0`. Broker visibility timeout in seconds. Must exceed the longest legitimate task runtime but stay short enough that SIGKILLed tasks redeliver promptly. ### `CELERY_WORKER_MAX_MEMORY_PER_CHILD` -Type `int`, default `4194304`, must be >= 0. +Type `int`, default `4194304`, must be `>= 0`. Recycle a prefork child past this resident size in KB; backstops docling/torch heap growth. Checked between tasks, so it does not bound the peak within one. 0 disables. ### `CELERY_WORKER_MAX_TASKS_PER_CHILD` -Type `int`, default `0`, must be >= 0. +Type `int`, default `0`, must be `>= 0`. Recycle a worker child after N tasks; 0 disables. @@ -703,43 +703,43 @@ Directory under the data home for uploaded sources. ### `UPLOAD_MAX_REQUEST_BYTES` -Type `int`, default `268435456`, must be > 0. +Type `int`, default `268435456`, must be `> 0`. Cap on an upload request body; applied by Flask before multipart parsing. ### `UPLOAD_MAX_FILE_BYTES` -Type `int`, default `104857600`, must be > 0. +Type `int`, default `104857600`, must be `> 0`. Cap on a single uploaded file; also enforced while copying. ### `PARSE_SPEC_MAX_BYTES` -Type `int`, default `10485760`, must be > 0. +Type `int`, default `10485760`, must be `> 0`. Cap on an OpenAPI/tool spec file accepted for parsing. ### `UPLOAD_MAX_ARCHIVE_BYTES` -Type `int`, default `262144000`, must be > 0. +Type `int`, default `262144000`, must be `> 0`. Cap on total bytes extracted from one uploaded archive. ### `UPLOAD_MAX_ARCHIVE_FILES` -Type `int`, default `10000`, must be > 0. +Type `int`, default `10000`, must be `> 0`. Cap on files extracted from one uploaded archive. ### `UPLOAD_MAX_ARCHIVE_RATIO` -Type `int`, default `1000`, must be > 0. +Type `int`, default `1000`, must be `> 0`. Maximum decompressed-to-compressed ratio before an archive is rejected. ### `UPLOAD_MAX_ARCHIVE_DEPTH` -Type `int`, default `3`, must be >= 0. +Type `int`, default `3`, must be `>= 0`. Maximum nesting depth of archives inside archives. @@ -787,7 +787,7 @@ Largest HTML/XML docling will parse, in bytes. ### `MARKUP_MAX_BYTES` -Type `int`, default `8000000`, must be >= 0. +Type `int`, default `8000000`, must be `>= 0`. HTML/XHTML larger than this (bytes) are head-truncated before the markdownify parser runs (the anydoc engine's HTML path). The tree that path builds costs ~50x the input (30 MB of HTML measured at 1.6 GB RSS) and the upload cap is 100 MB, so the gate is what keeps one upload from taking the ingest worker down. 0 disables it. @@ -835,13 +835,13 @@ Cap on the pixel count of an image passed to an agent. ### `GITHUB_INGEST_MAX_FILE_BYTES` -Type `int`, default `1048576`, must be >= 0. +Type `int`, default `1048576`, must be `>= 0`. Skip GitHub repo blobs larger than this (0 = no cap). ### `GITHUB_INGEST_MAX_WORKERS` -Type `int`, default `8`, must be >= 1. +Type `int`, default `8`, must be `>= 1`. Parallel file fetches per GitHub repo ingest. @@ -871,7 +871,7 @@ Absolute ceiling on the size-scaled parse window, in seconds. ### `DOCUMENT_PARSE_MAX_BYTES` -Type `int`, default `0`, must be >= 0. +Type `int`, default `0`, must be `>= 0`. Cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES). @@ -948,7 +948,7 @@ Native backend only: resolution at which pages without a text layer are rendered ### `OCR_MIN_CHARS_PER_PAGE` -Type `int`, default `20`, must be >= 0, also read from `DOCLING_OCR_MIN_CHARS_PER_PAGE`. +Type `int`, default `20`, must be `>= 0`, also read from `DOCLING_OCR_MIN_CHARS_PER_PAGE`. Chars-per-page floor below which an OCR'd PDF/image parse is treated as an OCR dropout rather than as content (long-running docling workers were observed returning zero characters for every scanned page after a long scanned PDF, with no error). docling retries once on a fresh full-page-OCR converter; both backends then fail loudly instead of indexing an empty document. 0 disables the guard. @@ -1149,7 +1149,7 @@ Bounds uvicorn's shutdown drain (uvicorn_worker doesn't forward --graceful-timeo ### `WSGI_THREADPOOL_WORKERS` -Type `int`, default `96`, must be >= 1. +Type `int`, default `96`, must be `>= 1`. Threads serving the WSGI (Flask) part of the app under the ASGI server. @@ -1172,31 +1172,31 @@ Internal SSE push channel (notifications and durable replay journal). False make ### `EVENTS_STREAM_MAXLEN` -Type `int`, default `1000`, must be >= 1. +Type `int`, default `1000`, must be `>= 1`. Per-user durable backlog cap in entries; ~24h of replay at typical rates. ### `SSE_KEEPALIVE_SECONDS` -Type `int`, default `15`, must be >= 1. +Type `int`, default `15`, must be `>= 1`. Interval between SSE keepalive comments. ### `SSE_MAX_CONCURRENT_PER_USER` -Type `int`, default `8`, must be >= 0. +Type `int`, default `8`, must be `>= 0`. Simultaneous SSE connections per user; each holds a pooled async Redis connection for its lifetime. 8 covers multi-tab use without one user starving the pool. 0 disables. ### `ASYNC_REDIS_MAX_CONNECTIONS` -Type `int`, default `2000`, must be >= 1. +Type `int`, default `2000`, must be `>= 1`. Pool size of the async Redis client behind the event-loop routes, per process. Every open notification tab, chat reconnect and device session holds one connection, so this caps concurrent streams per worker (redis-py's own default is 100). Keep the total across workers below the Redis server's maxclients (10000 by default). ### `EVENTS_REPLAY_MAX_PER_REQUEST` -Type `int`, default `200`, must be >= 1. +Type `int`, default `200`, must be `>= 1`. Backlog entries XRANGE returns per /api/events snapshot. Bounds what one replay moves from Redis to the wire: a client looping Last-Event-ID reconnects enumerates at most this many per round-trip. @@ -1220,13 +1220,13 @@ Length of the replay budget window. ### `MESSAGE_EVENTS_RETENTION_DAYS` -Type `int`, default `14`, must be > 0. +Type `int`, default `14`, must be `> 0`. Retention for the message_events journal, enforced by the cleanup_message_events beat task. Replay only needs streams a client could still be tailing. ### `REMOTE_DEVICE_SESSION_IDLE_SECONDS` -Type `int`, default `60`, must be > 0. +Type `int`, default `60`, must be `> 0`. Seconds without a heartbeat before a remote-device session is considered idle. @@ -1238,19 +1238,19 @@ Require signed commands from remote devices. ### `REMOTE_DEVICE_PAIRING_TTL_SECONDS` -Type `int`, default `600`, must be > 0. +Type `int`, default `600`, must be `> 0`. Lifetime of a pairing code. ### `REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS` -Type `int`, default `900`, must be > 605. +Type `int`, default `900`, must be `> 605`. Redis TTL of the per-device command queue, routing invocations cross-process so a scheduled run reaches the web-held device session. Must exceed the max drain deadline (605s) so a command for a briefly-offline device isn't evicted before its own drain gives up. ### `REMOTE_DEVICE_INVOCATION_TTL_SECONDS` -Type `int`, default `900`, must be > 0. +Type `int`, default `900`, must be `> 0`. Redis TTL of a pending remote-device invocation. @@ -1291,7 +1291,7 @@ Pre-fetch retrieval before the agent's first turn. ### `TOOL_RESULT_MAX_TOKENS` -Type `int`, default `20000`, must be >= 0. +Type `int`, default `20000`, must be `>= 0`. Cap on one tool result entering the LLM context (0 disables); journal and DB keep it whole. @@ -1303,7 +1303,7 @@ Compress long conversations once they approach the context window. ### `COMPRESSION_THRESHOLD_PERCENTAGE` -Type `float`, default `0.8`, must be > 0 and <= 1. +Type `float`, default `0.8`, must be `> 0` and `<= 1`. Fraction of the context window at which compression triggers. @@ -1327,7 +1327,7 @@ Keep only the last N compression points to prevent DB bloat. ### `COMPRESSION_RECENT_FIELD_MAX_TOKENS` -Type `int`, default `8000`, must be >= 0. +Type `int`, default `8000`, must be `>= 0`. Per-field cap on the verbatim tail kept after a compression point (0 disables). @@ -1410,7 +1410,7 @@ Persist scanned text alongside guardrail_events. Off by default: pre-redaction t ### `GUARDRAILS_EVENTS_RETENTION_DAYS` -Type `int`, default `30`, must be >= 1. +Type `int`, default `30`, must be `>= 1`. Days guardrail events are kept before the cleanup task removes them. @@ -1463,7 +1463,7 @@ How far ahead a one-off run may be scheduled, in seconds (one year). ### `SCHEDULE_RUN_OUTPUT_RETENTION_DAYS` -Type `int`, default `90`, must be > 0. +Type `int`, default `90`, must be `> 0`. Days scheduled-run output is kept. @@ -1582,13 +1582,13 @@ Default runtime language for created sandboxes. ### `DAYTONA_AUTO_STOP_INTERVAL` -Type `int`, default `15`, must be >= 0. +Type `int`, default `15`, must be `>= 0`. Minutes idle before Daytona auto-stops a sandbox (0 disables). ### `DAYTONA_AUTO_DELETE_INTERVAL` -Type `int`, default `60`, must be >= -1. +Type `int`, default `60`, must be `>= -1`. Minutes after stop before Daytona auto-deletes a sandbox (-1 disables). diff --git a/docsgpt/core/settings/reference.py b/docsgpt/core/settings/reference.py index e5aa3392..c471ee0f 100644 --- a/docsgpt/core/settings/reference.py +++ b/docsgpt/core/settings/reference.py @@ -104,7 +104,8 @@ def _render_field(name: str, field: FieldInfo) -> str: facts = [f"Type `{_type_name(field.annotation)}`", f"default {_default_text(field)}"] constraints = _constraints(field) if constraints: - facts.append("must be " + " and ".join(constraints)) + # Code spans: a bare ``<=`` in MDX prose is parsed as the start of a JSX tag. + facts.append("must be " + " and ".join(f"`{c}`" for c in constraints)) aliases = _aliases(name, field) if aliases: facts.append("also read from " + ", ".join(f"`{a}`" for a in aliases)) diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py index ff53fba4..36967fce 100644 --- a/tests/core/test_settings.py +++ b/tests/core/test_settings.py @@ -139,6 +139,15 @@ class TestReference: for name in Settings.model_fields: assert page.count(f"### `{name}`") == 1, name + def test_reference_prose_has_no_bare_angle_brackets_or_braces(self): + """MDX parses ``<`` and ``{`` in prose as JSX; only code spans may carry them raw.""" + for lineno, line in enumerate(render_reference().splitlines(), 1): + if line.startswith(("{/*", "---")): + continue + prose = "".join(line.split("`")[::2]) # drop the inside of every code span + prose = prose.replace("\\{", "").replace("\\}", "") # escaped braces are fine + assert "<" not in prose and "{" not in prose and "}" not in prose, f"line {lineno}: {line}" + def test_checked_in_reference_is_current(self): path: Path = reference_path() if not path.exists(): From da58c072a0ff69b334dc64366f3089fe8cc492ee Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 13:15:56 +0100 Subject: [PATCH 071/130] fix: five from the full review - doctor printed OPENAI_BASE_URL raw, the third place a credential-bearing URL reached the terminal; it goes through _endpoint like the others. - dev.run left its output readers unjoined, so a child's last lines could be lost on exit. The threads are kept and joined during teardown. - logs read each file and then reopened it to follow, so anything written in between appeared in neither. One handle now serves both. - dev took --port straight from argparse: 0 would have served on an ephemeral port while printing 0, and oversized values reach socket.bind. It goes through _port_number first. - doctor --redis-url overrode only the broker, so a stale result backend or cache was still pinged and the flag looked broken. All three endpoints now come from the URL given. --- docsgpt/deploy/commands.py | 58 +++++++++++++++++++++--------- docsgpt/deploy/dev.py | 9 ++++- tests/deploy/test_dev.py | 11 ++++++ tests/deploy/test_doctor.py | 71 +++++++++++++++++++++++++++++++++++++ 4 files changed, 131 insertions(+), 18 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index ca7be581..25a7d9a0 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -615,24 +615,35 @@ def logs(args, context: Optional[Context] = None) -> int: if _mode(directory) == "native": logs_dir = directory / "logs" wanted = args.services or ["api", "worker"] + # Opened before the first read and kept: a line written between printing what is there and + # starting to follow would otherwise appear in neither. + handles = {} for service in wanted: path = logs_dir / f"{service}.log" print(f"=== {path}") - if path.is_file(): - lines = path.read_text(encoding="utf-8", errors="replace").splitlines() - print("\n".join(lines[-args.tail:] if args.tail else lines)) - else: + if not path.is_file(): print("(nothing logged yet)") + continue + handle = path.open("r", encoding="utf-8", errors="replace") + lines = handle.read().splitlines() + print("\n".join(lines[-args.tail:] if args.tail else lines)) + handles[service] = handle if args.follow: - return _follow(logs_dir, wanted) + return _follow(logs_dir, wanted, handles) + for handle in handles.values(): + handle.close() return 0 options = (["--follow"] if args.follow else []) + (["--tail", str(args.tail)] if args.tail else []) return context.docker.compose(directory, "logs", *options, *args.services, check=False).returncode -def _follow(logs_dir: Path, services: list) -> int: - """Print new lines from each service's log until the terminal interrupts, prefixed by service.""" - handles: dict = {} +def _follow(logs_dir: Path, services: list, handles: Optional[dict] = None) -> int: + """Print new lines from each service's log until the terminal interrupts, prefixed by service. + + ``handles`` are the files already read, positioned where that read left off, so nothing written + in between is skipped. Files that did not exist yet are opened as they appear. + """ + handles = dict(handles or {}) try: while True: for service in services: @@ -734,7 +745,8 @@ def dev(args, context: Optional[Context] = None) -> int: "`docsgpt dev` runs the code in a source checkout, and this is an installed package. " "Clone the repository and run it from there, or use `docsgpt up --native` to run this copy." ) - if not _port_is_free(args.port): + port = _port_number(args.port, "--port") + if not _port_is_free(port): raise DeployError( f"port {args.port} is already in use, so the API cannot bind it. Stop what is on it " f"(a previous `docsgpt dev`, or `docsgpt down` for an install), or pass --port." @@ -879,7 +891,25 @@ def _check_provider(env: Mapping[str, str]) -> Check: return Check("provider", "ok", "the DocsGPT public API (no key needed)") if not (env.get("API_KEY") or env.get("OPENAI_API_KEY")): return Check("provider", "fail", f"{provider} is configured but no API_KEY is set") - return Check("provider", "ok", f"{provider}{' at ' + env['OPENAI_BASE_URL'] if env.get('OPENAI_BASE_URL') else ''}") + endpoint = _endpoint(env["OPENAI_BASE_URL"]) if env.get("OPENAI_BASE_URL") else "" + return Check("provider", "ok", f"{provider}{' at ' + endpoint if endpoint else ''}") + + +def _redis_to_check(args, env: Mapping[str, str]) -> dict: + """The Redis endpoints to ping: all three from --redis-url when given, else what .env holds. + + Overriding only the broker would still ping a stale result backend or cache, and the check + fails on the first endpoint that does not answer -- so --redis-url would appear not to work. + """ + if args.redis_url: + names = {"CELERY_BROKER_URL": "broker", "CELERY_RESULT_BACKEND": "results", "CACHE_REDIS_URL": "cache"} + return {names[key]: value for key, value in _redis_urls(args.redis_url).items()} + pairs = ( + ("broker", env.get("CELERY_BROKER_URL")), + ("results", env.get("CELERY_RESULT_BACKEND")), + ("cache", env.get("CACHE_REDIS_URL")), + ) + return {key: value for key, value in pairs if value} def doctor(args, context: Optional[Context] = None) -> int: @@ -898,13 +928,7 @@ def doctor(args, context: Optional[Context] = None) -> int: Check("settings", "ok" if env_path.is_file() else "warn", f"{env_path}" if env_path.is_file() else f"{env_path} does not exist yet; defaults are in use"), _check_postgres(args.postgres_uri or env.get("POSTGRES_URI")), - _check_redis({ - key: value for key, value in ( - ("broker", args.redis_url or env.get("CELERY_BROKER_URL")), - ("results", env.get("CELERY_RESULT_BACKEND")), - ("cache", env.get("CACHE_REDIS_URL")), - ) if value - }), + _check_redis(_redis_to_check(args, env)), _check_provider(env), ] diff --git a/docsgpt/deploy/dev.py b/docsgpt/deploy/dev.py index a552646e..18c4e485 100644 --- a/docsgpt/deploy/dev.py +++ b/docsgpt/deploy/dev.py @@ -181,6 +181,7 @@ def run( colour = out.isatty() if colour is None else colour lock = threading.Lock() running: list[tuple[Child, object]] = [] + pumps: list[threading.Thread] = [] try: for child in children: process = spawn( @@ -195,7 +196,9 @@ def run( start_new_session=os.name != "nt", ) running.append((child, process)) - threading.Thread(target=_pump, args=(child, process, out, lock, colour), daemon=True).start() + pump = threading.Thread(target=_pump, args=(child, process, out, lock, colour), daemon=True) + pump.start() + pumps.append(pump) while True: for child, process in running: @@ -214,3 +217,7 @@ def run( return 0 finally: _stop(running, grace, sleep) + # Join the readers: a child's last lines are still in flight when it exits, and dropping + # them loses exactly the output that says why it stopped. + for pump in pumps: + pump.join(timeout=grace) diff --git a/tests/deploy/test_dev.py b/tests/deploy/test_dev.py index a8344954..5426cbd6 100644 --- a/tests/deploy/test_dev.py +++ b/tests/deploy/test_dev.py @@ -123,6 +123,17 @@ class TestRun: assert "api | hello" in out.getvalue() assert code == 1, "a child that ends by itself ends the session, however it exited" + def test_output_in_flight_is_not_lost_when_run_returns(self, tmp_path): + """The reader is a thread: without joining it, a child's last lines can never be printed.""" + out = io.StringIO() + lines = [f"line {number}\n" for number in range(200)] + children = [dev.Child(name="api", command=["x"], cwd=tmp_path)] + code = dev.run(children, out=out, spawn=self._spawn([FakeProcess(lines, code=0)]), + sleep=lambda _: None, colour=False) + assert code == 1 + printed = out.getvalue() + assert "line 0" in printed and "line 199" in printed, "every line the child wrote is printed" + def test_a_failing_child_returns_its_code(self, tmp_path): out = io.StringIO() children = [dev.Child(name="api", command=["false"], cwd=tmp_path)] diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 1a55ce84..252af2fa 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -1,5 +1,6 @@ """`docsgpt doctor`, `restart`, following native logs, and settings that apply themselves.""" +import argparse import socket import pytest @@ -49,6 +50,18 @@ class TestDevCommand: _run(["dev", "--port", str(port)], _context()) assert started == [], "nothing is spawned when the port is taken" + @pytest.mark.parametrize("port", ["0", "-1", "70000"]) + def test_a_port_argparse_accepts_but_a_socket_cannot_use(self, monkeypatch, tmp_path, port): + """argparse takes any integer: 0 would serve on an ephemeral port while printing 0.""" + from docsgpt.deploy import dev as dev_module + + self._checkout(monkeypatch, tmp_path) + started = [] + monkeypatch.setattr(dev_module, "run", lambda children, **kwargs: started.append(children) or 0) + with pytest.raises(DeployError, match="not a port number"): + _run(["dev", "--port", port], _context()) + assert started == [] + def test_a_busy_mock_llm_port_is_refused_before_anything_starts(self, monkeypatch, tmp_path): """It starts first and the rest are pointed at it, so a clash cannot wait until spawn time.""" from docsgpt.deploy import dev as dev_module @@ -184,6 +197,33 @@ class TestFollowLogs: assert "old line" not in printed, "it starts at the end, like tail -f" +class TestLogsTransition: + def test_a_line_written_between_the_read_and_the_follow_is_not_lost(self, tmp_path, capsys, monkeypatch): + """The read and the follow used to open the file twice, and whatever landed between was gone.""" + services = FakeServices() + argv = ["up", "--native", "--dir", str(tmp_path), "--yes", "--postgres-uri", "postgresql://localhost/d"] + assert _run(argv, _native_context(services)) == 0 + logs = tmp_path / "logs" + (logs / "api.log").write_text("first\n", encoding="utf-8") + + def sleep(_): + raise KeyboardInterrupt + + original = commands._follow + + def follow(logs_dir, wanted, handles=None): + # Written after the read, before following starts. + (logs / "api.log").open("a", encoding="utf-8").write("during\n") + monkeypatch.setattr(commands.time, "sleep", sleep) + return original(logs_dir, wanted, handles) + + monkeypatch.setattr(commands, "_follow", follow) + assert _run(["logs", "-f", "api", "--dir", str(tmp_path)], _native_context(services)) == 0 + printed = capsys.readouterr().out + assert "first" in printed + assert "during" in printed, "the line written during the handover has to appear" + + class TestChecks: def test_the_public_api_needs_no_key(self): check = commands._check_provider({"LLM_PROVIDER": "docsgpt"}) @@ -201,6 +241,17 @@ class TestChecks: assert check.level == "ok" assert "8090" in check.detail + def test_a_provider_base_url_is_printed_without_its_credentials(self): + """An OpenAI-compatible endpoint can carry userinfo, and doctor prints its detail.""" + check = commands._check_provider( + {"LLM_PROVIDER": "openai", "API_KEY": "x", + "OPENAI_BASE_URL": "https://someone:sEcReTtOkEn@models.example.com/v1"} + ) + assert check.level == "ok" + assert "sEcReTtOkEn" not in check.detail + assert "someone" not in check.detail + assert "models.example.com" in check.detail + def test_services_are_named_for_the_install(self, tmp_path): names = _names(tmp_path) assert commands._chosen_services(names, []) == list(names) @@ -400,6 +451,26 @@ class TestMigrationHead: assert head[0].isdigit(), head +class TestRedisSelection: + def test_redis_url_replaces_every_endpoint(self): + """Overriding only the broker would still ping a stale backend, and one failure fails all.""" + args = argparse.Namespace(redis_url="redis://given:6379/5") + stale = { + "CELERY_BROKER_URL": "redis://old:6379/0", + "CELERY_RESULT_BACKEND": "redis://old:6379/1", + "CACHE_REDIS_URL": "redis://old:6379/2", + } + chosen = commands._redis_to_check(args, stale) + assert set(chosen) == {"broker", "results", "cache"} + assert all("old" not in url for url in chosen.values()), chosen + assert chosen["broker"].endswith("/5") and chosen["cache"].endswith("/7") + + def test_without_the_flag_the_settings_are_used(self): + args = argparse.Namespace(redis_url=None) + chosen = commands._redis_to_check(args, {"CELERY_BROKER_URL": "redis://localhost:6379/0"}) + assert chosen == {"broker": "redis://localhost:6379/0"} + + class TestDoctor: def _only(self, monkeypatch, postgres, redis): monkeypatch.setattr(commands, "_check_postgres", lambda uri: postgres) From 6f981a559814b7c556c942fc6bc26250240e6d04 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 13:19:34 +0100 Subject: [PATCH 072/130] test: compare against the sanitiser, not a host substring CodeQL flags "host" in message as incomplete URL sanitisation. I removed that shape from two assertions and introduced a third in the same commit; this takes it out of the provider test and the remaining Redis one, comparing with what _endpoint produced instead, which is the contract those tests actually mean. --- tests/deploy/test_doctor.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index 252af2fa..eb085038 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -243,14 +243,14 @@ class TestChecks: def test_a_provider_base_url_is_printed_without_its_credentials(self): """An OpenAI-compatible endpoint can carry userinfo, and doctor prints its detail.""" - check = commands._check_provider( - {"LLM_PROVIDER": "openai", "API_KEY": "x", - "OPENAI_BASE_URL": "https://someone:sEcReTtOkEn@models.example.com/v1"} - ) + url = "https://someone:sEcReTtOkEn@models.example.com/v1" + check = commands._check_provider({"LLM_PROVIDER": "openai", "API_KEY": "x", "OPENAI_BASE_URL": url}) assert check.level == "ok" assert "sEcReTtOkEn" not in check.detail assert "someone" not in check.detail - assert "models.example.com" in check.detail + # Compared with what the sanitiser produced rather than a host substring, which is both the + # real contract and the pattern CodeQL warns about. + assert check.detail == f"openai at {commands._endpoint(url)}" def test_services_are_named_for_the_install(self, tmp_path): names = _names(tmp_path) @@ -394,9 +394,10 @@ class TestRedisCheck: return Client() monkeypatch.setattr(redis.Redis, "from_url", classmethod(from_url)) - check = commands._check_redis({"cache": "redis://localhost:6379/2"}) + url = "redis://localhost:6379/2" + check = commands._check_redis({"cache": url}) assert check.level == "fail" - assert "cache" in check.detail and "6379/2" in check.detail + assert check.detail.startswith(f"cache ({commands._endpoint(url)}) does not answer") def test_a_password_in_the_url_is_never_printed(self, monkeypatch): """doctor output goes into terminals, CI logs and pasted issue reports.""" From 63bcdbe5a579278e6f8c963b5d97c383b7754803 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Thu, 17 Sep 2026 13:20:04 +0100 Subject: [PATCH 073/130] refactor(settings): fail at load when SCIM is enabled without a token Review follow-up. SCIM_ENABLED=true with no SCIM_TOKEN was only caught per request, as a 503 from the SCIM routes. The auth group's model validator now rejects that combination when Settings loads, the same way it rejects AUTH_TYPE=oidc without its required settings, so the misconfiguration is reported once at startup rather than on the first provisioning call. The route-level guard stays as defence in depth. --- docsgpt/core/settings/auth.py | 4 +++- tests/core/test_settings.py | 8 ++++++++ 2 files changed, 11 insertions(+), 1 deletion(-) diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py index 5069d6cf..85c57b84 100644 --- a/docsgpt/core/settings/auth.py +++ b/docsgpt/core/settings/auth.py @@ -92,9 +92,11 @@ class AuthSettings(SettingsGroup): return normalize_choice(v) @model_validator(mode="after") - def _require_oidc_settings(self): + def _require_dependent_settings(self): if self.AUTH_TYPE == "oidc": missing = [name for name in OIDC_REQUIRED if not getattr(self, name)] if missing: raise ValueError(f"AUTH_TYPE=oidc requires settings: {', '.join(missing)}") + if self.SCIM_ENABLED and not self.SCIM_TOKEN: + raise ValueError("SCIM_ENABLED requires settings: SCIM_TOKEN") return self diff --git a/tests/core/test_settings.py b/tests/core/test_settings.py index 36967fce..1074fe63 100644 --- a/tests/core/test_settings.py +++ b/tests/core/test_settings.py @@ -177,6 +177,14 @@ class TestCrossFieldRules: def test_oidc_settings_are_not_required_for_other_modes(self): assert Settings.model_validate({"AUTH_TYPE": "session_jwt"}).OIDC_ISSUER is None + @pytest.mark.parametrize("raw", [None, "", "None"]) + def test_scim_enabled_requires_a_token(self, raw): + with pytest.raises(ValidationError, match="SCIM_ENABLED requires settings: SCIM_TOKEN"): + Settings.model_validate({"SCIM_ENABLED": True, "SCIM_TOKEN": raw}) + + def test_scim_disabled_needs_no_token(self): + assert Settings.model_validate({"SCIM_ENABLED": False}).SCIM_TOKEN is None + @pytest.mark.unit class TestClosedChoices: From 8b064c96f25dcc4e9db08a4bde215089aaf92630 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 13:34:59 +0100 Subject: [PATCH 074/130] fix: keep credentials and 5000-digit paths out of the Redis URL errors MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _redis_urls raises before any check runs and cli.main prints what it raises, so every one of its five messages interpolated the URL it was given — password and all — straight to stderr. They render it through _endpoint now. _endpoint itself then echoed the whole path, so an over-long database number came back at full length: 5132 characters of error for one bad setting. It caps the path it renders, which protects every caller, since that string is written into terminals, CI logs and error messages rather than reused as a URL. --- docsgpt/deploy/commands.py | 18 ++++++++++++------ tests/deploy/test_doctor.py | 22 ++++++++++++++++++++++ 2 files changed, 34 insertions(+), 6 deletions(-) diff --git a/docsgpt/deploy/commands.py b/docsgpt/deploy/commands.py index 25a7d9a0..295ae0ac 100644 --- a/docsgpt/deploy/commands.py +++ b/docsgpt/deploy/commands.py @@ -215,30 +215,32 @@ def _redis_urls(base: str) -> dict[str, str]: parts = urlsplit(base) except ValueError as exc: # urlsplit raises on things like redis://[::1 ; that is a typo, not a crash. - raise DeployError(f"the Redis URL {base!r} could not be read: {exc}") from exc + raise DeployError(f"the Redis URL could not be read: {_scrub(str(exc), base)}") from exc try: port = parts.port # a non-numeric or out-of-range port raises here, not when the URL is split except ValueError as exc: raise DeployError( - f"the Redis URL {base!r} has an unusable port: {exc}. Pass a URL like " + f"the Redis URL ({_endpoint(base)}) has an unusable port: {_scrub(str(exc), base)}. " + "Pass a URL like " "redis://host:6379 or redis://host:6379/5." ) from exc if port == 0: # urlsplit is happy with it, since 0 is inside the range, but nothing can connect to it. raise DeployError( - f"the Redis URL {base!r} has an unusable port: 0. Pass a URL like " + f"the Redis URL ({_endpoint(base)}) has an unusable port: 0. Pass a URL like " "redis://host:6379 or redis://host:6379/5." ) if parts.scheme not in ("redis", "rediss"): raise DeployError( - f"the Redis URL {base!r} should start with redis:// or rediss://, with any options as " + f"the Redis URL ({_endpoint(base)}) should start with redis:// or rediss://, with any options as " "query parameters, so the broker, the result backend and the cache can be given a " "database each." ) path = parts.path.rstrip("/").lstrip("/") if path and not (path.isascii() and path.isdigit()): raise DeployError( - f"the Redis URL {base!r} has {path!r} where a database number would go. Pass a URL like " + f"the Redis URL ({_endpoint(base)}) has {path[:20]!r} where a database number would go. " + "Pass a URL like " "redis://host:6379 or redis://host:6379/5." ) try: @@ -247,7 +249,7 @@ def _redis_urls(base: str) -> dict[str, str]: # Python refuses to convert a digit string past its conversion limit, and that is a typo # rather than a crash. raise DeployError( - f"the Redis URL {base!r} has a database number too long to read. Pass a URL like " + f"the Redis URL ({_endpoint(base)}) has a database number too long to read. Pass a URL like " "redis://host:6379 or redis://host:6379/5." ) from exc return { @@ -848,6 +850,10 @@ def _endpoint(url: str) -> str: return "the configured URL" if port: host = f"{host}:{port}" + # The path is capped: this string is for a person to read, and it goes into terminals, CI logs + # and error messages. A 5000-digit database number would otherwise flood all three. + if len(path) > 40: + path = f"{path[:40]}..." return f"{scheme}://{host}{path}" if host else "the configured URL" diff --git a/tests/deploy/test_doctor.py b/tests/deploy/test_doctor.py index eb085038..e3072249 100644 --- a/tests/deploy/test_doctor.py +++ b/tests/deploy/test_doctor.py @@ -472,6 +472,28 @@ class TestRedisSelection: assert chosen == {"broker": "redis://localhost:6379/0"} +class TestRedisUrlErrors: + """The validator raises before any check runs, and cli.main prints what it raises.""" + + SECRET = "sUpErSeCrEt" + + @pytest.mark.parametrize("url", [ + "redis://user:sUpErSeCrEt@host:not-a-port/0", + "redis://user:sUpErSeCrEt@host:0/0", + "redis://user:sUpErSeCrEt@host:6379/queue", + "postgres://user:sUpErSeCrEt@host:6379/0", + "redis://user:sUpErSeCrEt@host:6379/" + "1" * 5000, + ]) + def test_a_rejected_url_never_carries_its_password_into_the_error(self, tmp_path, url): + (tmp_path / ".env").write_text("LLM_PROVIDER=docsgpt\n", encoding="utf-8") + with pytest.raises(DeployError) as raised: + _run(["doctor", "--dir", str(tmp_path), "--redis-url", url], _context()) + message = str(raised.value) + assert self.SECRET not in message + assert "user" not in message + assert len(message) < 400, "an over-long database number must not come back in the message" + + class TestDoctor: def _only(self, monkeypatch, postgres, redis): monkeypatch.setattr(commands, "_check_postgres", lambda uri: postgres) From 94a33c478122894ad3bf2e8d62ff8f50a34c264f Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:38:54 +0100 Subject: [PATCH 075/130] fix(graphrag): rank without scipy so graph retrieval stops falling back MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit networkx.pagerank delegates to a scipy implementation, and scipy is not a DocsGPT dependency — it only arrives transitively through the optional docling extra. In a default install every graph retrieval raised ModuleNotFoundError inside _ppr_scores, hit the per-source except, and degraded to ClassicRAG: the graph was built and paid for, then never used, with one ERROR line per source per query as the only signal. Rank with a local power iteration over the same row-normalized transition matrix: undirected edges normalized per endpoint, dangling nodes redistributed along the restart vector, and the restart vector normalized across the nodes the subgraph actually holds so seed mass cannot leak. Parity with networkx is asserted while scipy happens to be installed in the test env, and the retrieval path is exercised with the import blocked. --- docsgpt/retriever/graph_rag.py | 109 ++++++++++++++++++++++++-- tests/retriever/test_graph_rag.py | 122 ++++++++++++++++++++++++++++++ 2 files changed, 226 insertions(+), 5 deletions(-) diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index c7ac538e..5593a9b2 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -1,9 +1,13 @@ """GraphRAG local retriever — Personalized PageRank over a per-source graph. -Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> networkx -Personalized PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR -mass -> shared token budget. No LLM call at query time beyond the (optional, -reused) rephrase. +Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> Personalized +PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR mass -> shared +token budget. No LLM call at query time beyond the (optional, reused) rephrase. + +``networkx`` supplies the graph structure, but the ranking is the local power +iteration in :func:`_personalized_pagerank`: ``nx.pagerank`` delegates to scipy, +which DocsGPT does not depend on, so calling it turned every graph retrieval +into a silent ClassicRAG fallback. Composes :class:`ClassicRAG` rather than subclassing: PPR doesn't fit the ``_fetch_candidates`` hook, but the composed instance supplies the rephrase, the @@ -40,6 +44,99 @@ def _idf(doc_freq: Any) -> float: return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0) +def _restart_vector(nodes: List[Any], personalization: Dict[Any, float] | None) -> Dict[Any, float]: + """Normalized restart distribution over ``nodes``. + + Weights are clamped at zero (a cosine distance above 1 yields a negative + seed weight) and normalized across the nodes actually present in the graph, + so restart mass can never leak to a node the subgraph does not contain. An + absent or all-zero personalization collapses to a uniform restart. + """ + if personalization: + weights = { + node: max(float(personalization.get(node, 0.0) or 0.0), 0.0) + for node in nodes + } + total = sum(weights.values()) + if total > 0: + return {node: weight / total for node, weight in weights.items()} + uniform = 1.0 / len(nodes) + return {node: uniform for node in nodes} + + +def _personalized_pagerank( + graph: nx.Graph, + personalization: Dict[Any, float] | None = None, + *, + weight: str = "weight", + alpha: float = 0.85, + max_iter: int = 100, + tol: float = 1.0e-6, +) -> Dict[Any, float]: + """Personalized PageRank by power iteration — no scipy. + + ``networkx.pagerank`` delegates to a scipy implementation, and scipy is not + a DocsGPT dependency: in a default install the import raises and every graph + retrieval silently degrades to the ClassicRAG fallback. This is the same + algorithm over the same row-normalized transition matrix, so the ranking is + unchanged where scipy happens to be installed. + + Args: + graph: Undirected graph whose edges may carry a ``weight`` attribute. + personalization: Node -> restart weight; ``None`` means uniform. + weight: Edge attribute holding the weight. + alpha: Damping factor. + max_iter: Iteration cap. The last iterate is returned if it is hit — + retrieval degrades to a slightly less converged ranking rather than + raising, which is what the library does. + tol: Convergence tolerance; iteration stops below ``len(graph) * tol``. + + Returns: + Node -> PageRank mass, summing to ~1.0. Empty dict for an empty graph. + """ + nodes = list(graph.nodes) + node_count = len(nodes) + if node_count == 0: + return {} + + restart = _restart_vector(nodes, personalization) + + # Row-normalized transitions. An undirected edge is traversable from both + # endpoints, so each node normalizes over its own incident weights. + transitions: Dict[Any, List[Any]] = {} + for node in nodes: + neighbors = [] + total = 0.0 + for neighbor, data in graph[node].items(): + edge_weight = float(data.get(weight, 1.0) or 1.0) + if edge_weight <= 0: + continue + neighbors.append((neighbor, edge_weight)) + total += edge_weight + transitions[node] = ( + [(n, w / total) for n, w in neighbors] if total > 0 else [] + ) + + # A node with no usable edge is dangling: its mass would vanish each pass, + # so it is redistributed along the restart vector instead. + dangling = [node for node in nodes if not transitions[node]] + + ranks = {node: 1.0 / node_count for node in nodes} + for _ in range(max_iter): + previous = ranks + ranks = dict.fromkeys(nodes, 0.0) + leaked = alpha * sum(previous[node] for node in dangling) + for node in nodes: + share = alpha * previous[node] + for neighbor, transition in transitions[node]: + ranks[neighbor] += share * transition + for node in nodes: + ranks[node] += (leaked + 1.0 - alpha) * restart[node] + if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol: + break + return ranks + + class GraphRAGRetriever(BaseRetriever): """Per-source PPR retriever; falls back to ClassicRAG when a source has no graph.""" @@ -117,7 +214,9 @@ class GraphRAGRetriever(BaseRetriever): if not any(personalization.values()): personalization = None - ranks = nx.pagerank(graph, personalization=personalization, weight="weight") + ranks = _personalized_pagerank( + graph, personalization=personalization, weight="weight" + ) return { node: rank * _idf(graph.nodes[node].get("doc_freq", 0)) for node, rank in ranks.items() diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index e3fe2fae..5f8cdf80 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -800,3 +800,125 @@ class TestGraphRAGBatching: assert rag._get_data() == [] mock_store_cls.assert_not_called() rag._classic._get_data.assert_not_called() + + +# ── Personalized PageRank without scipy ────────────────────────────────────── + + +@pytest.fixture +def _no_scipy(monkeypatch): + """Make ``import scipy`` fail, as it does in a default install. + + ``scipy`` is not a DocsGPT dependency — it only reaches this test env + through the optional docling extra. ``networkx.pagerank`` delegates to its + scipy implementation, so ranking must not go through it. + """ + import sys + + for name in [m for m in list(sys.modules) if m == "scipy" or m.startswith("scipy.")]: + monkeypatch.delitem(sys.modules, name) + monkeypatch.setitem(sys.modules, "scipy", None) + + +def _chain_graph(): + """Weighted chain a-b-c-d plus a heavier shortcut a-d.""" + import networkx as nx + + graph = nx.Graph() + graph.add_weighted_edges_from( + [("a", "b", 1.0), ("b", "c", 2.0), ("c", "d", 1.0), ("a", "d", 0.5)] + ) + return graph + + +@pytest.mark.unit +class TestPersonalizedPageRankWithoutScipy: + def test_ranking_runs_when_scipy_is_missing(self, _no_scipy): + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = _chain_graph() + ranks = _personalized_pagerank( + graph, personalization={"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0} + ) + + assert set(ranks) == {"a", "b", "c", "d"} + assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6) + assert all(rank > 0 for rank in ranks.values()) + # Pinned from the parity test below, which runs the library + # implementation over the same graph while scipy is installed here. + assert sorted(ranks, key=ranks.get, reverse=True) == ["b", "a", "c", "d"] + # The seed outranks the node furthest from it along the heavy path. + assert ranks["a"] > ranks["d"] + + def test_matches_networkx_within_tolerance(self): + """Parity with the library implementation, while it is installed here.""" + import networkx as nx + + pytest.importorskip("scipy") + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = _chain_graph() + personalization = {"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0} + + ours = _personalized_pagerank(graph, personalization=personalization) + theirs = nx.pagerank(graph, personalization=personalization, weight="weight") + + for node in theirs: + assert ours[node] == pytest.approx(theirs[node], abs=1e-6) + + def test_uniform_personalization_when_none(self): + pytest.importorskip("scipy") + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = _chain_graph() + ours = _personalized_pagerank(graph, personalization=None) + theirs = nx.pagerank(graph, personalization=None, weight="weight") + + for node in theirs: + assert ours[node] == pytest.approx(theirs[node], abs=1e-6) + + def test_isolated_node_still_gets_mass(self): + """A node with no edges is dangling; its mass must not vanish.""" + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = nx.Graph() + graph.add_edge("a", "b", weight=1.0) + graph.add_node("lonely") + + ranks = _personalized_pagerank(graph, personalization=None) + + assert ranks["lonely"] > 0 + assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6) + + def test_empty_graph_returns_empty(self): + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + assert _personalized_pagerank(nx.Graph(), personalization=None) == {} + + @patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10) + @patch("docsgpt.retriever.graph_rag.GraphStore") + @patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True) + def test_graph_retrieval_does_not_fall_back_without_scipy( + self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed, _no_scipy + ): + """The whole PPR path runs with scipy absent — no ClassicRAG fallback.""" + nodes = [{"id": "n1", "doc_freq": 1}, {"id": "n2", "doc_freq": 1}] + edges = [{"src_node_id": "n1", "dst_node_id": "n2", "weight": 1.0}] + node_chunks = {"n1": ["c1"], "n2": ["c2"]} + chunk_texts = {"c1": "near", "c2": "far"} + seed_rows = [{"id": "n1", "distance": 0.0}] + store = _store_with_graph(nodes, edges, node_chunks, chunk_texts, seed_rows) + mock_store_cls.return_value = store + + rag = _make_retriever(chunks=2) + rag._classic_for_sources = Mock(side_effect=AssertionError("fell back")) + + docs = rag._get_data() + + assert [doc["text"] for doc in docs] == ["near", "far"] From 92e19ac177ef5bf170b8e63ecea5f6e0e4174f1d Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:40:20 +0100 Subject: [PATCH 076/130] fix(graphrag): dispatch extraction to the provider that serves the model MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _build_extraction_llm passed settings.LLM_PROVIDER with the resolved extraction model id. Those two disagree in any deployment that leaves the provider at its default: the model id comes from GRAPHRAG_EXTRACTION_MODEL or LLM_NAME, while the provider stays "docsgpt" — the hosted public endpoint, which does not serve it. The request is rejected, the shared fallback answers instead, and the graph gets built by a different model than the one configured, with nothing in the summary saying so. Resolve the provider from the model registry (owner-scoped, so per-user BYOM ids resolve too), fall back to settings.LLM_PROVIDER only when the model is unknown, and take the API key for the provider actually dispatched to rather than the generic settings.API_KEY. The effective provider is logged, since a silent swap was the whole failure mode. --- docsgpt/graphrag/extraction.py | 31 +++++++- tests/graphrag/test_extraction.py | 118 ++++++++++++++++++++++++++++++ 2 files changed, 147 insertions(+), 2 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index beaf3d66..aa76189f 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -23,6 +23,10 @@ import logging import re from typing import Any, Callable, Dict, List, Optional +from docsgpt.core.model_utils import ( + get_api_key_for_provider, + get_provider_from_model_id, +) from docsgpt.core.settings import settings from docsgpt.llm.llm_creator import LLMCreator from docsgpt.storage.db.source_config import SourceConfig @@ -63,14 +67,37 @@ def _resolve_max_chunks(config: SourceConfig) -> int: return config.graph.max_chunks or settings.GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION +def _resolve_extraction_provider( + model_id: Optional[str], user: Optional[str] +) -> str: + """The provider that serves ``model_id``, else the deployment default. + + ``settings.LLM_PROVIDER`` is only a default (``docsgpt``, the hosted public + endpoint, out of the box). Dispatching the resolved extraction model + through it sends the request to a provider that does not serve that model: + the call is rejected, the shared fallback answers instead, and the graph is + built by a different model than the one configured — with nothing in the + summary to say so. ``user`` scopes the lookup so a per-user (BYOM) model id + resolves as well. + """ + provider = ( + get_provider_from_model_id(model_id, user_id=user) if model_id else None + ) + return provider or settings.LLM_PROVIDER + + def _build_extraction_llm( model_id: Optional[str], user: Optional[str], request_id: Optional[str] ): """Build the extraction LLM tagged for token-usage attribution to the owner.""" decoded_token = {"sub": user} if user else None + provider = _resolve_extraction_provider(model_id, user) + logger.info( + "Graph extraction dispatching model=%s via provider=%s", model_id, provider + ) llm = LLMCreator.create_llm( - settings.LLM_PROVIDER, - api_key=settings.API_KEY, + provider, + api_key=get_api_key_for_provider(provider), user_api_key=None, decoded_token=decoded_token, model_id=model_id, diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index ad2d670f..eaaf881d 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -403,6 +403,124 @@ class TestModelResolution: assert extraction_module._resolve_max_chunks(config) == 5 +@pytest.mark.unit +class TestExtractionProviderResolution: + """The extraction model decides the provider, not ``LLM_PROVIDER``. + + ``settings.LLM_PROVIDER`` is the deployment default (``docsgpt`` out of the + box, i.e. the hosted public endpoint). Dispatching the resolved extraction + model through it sends the call to a provider that never serves that model: + the request is rejected, the shared fallback answers instead, and the graph + is quietly built by a different model than the one configured. + """ + + def _capture_create_llm(self, monkeypatch, llm=None): + captured = {} + + def _create(provider, *args, **kwargs): + captured["provider"] = provider + captured["args"] = args + captured["kwargs"] = kwargs + return llm or _StubLLM([]) + + monkeypatch.setattr( + extraction_module.LLMCreator, "create_llm", staticmethod(_create) + ) + return captured + + def test_provider_comes_from_the_model_registry(self, monkeypatch): + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt") + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", lambda *a, **k: "openai" + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "sk-openai" + ) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("gpt-4o-mini", "owner-1", "req-1") + + assert captured["provider"] == "openai" + assert captured["kwargs"]["api_key"] == "sk-openai" + assert captured["kwargs"]["model_id"] == "gpt-4o-mini" + + def test_owner_scopes_the_registry_lookup(self, monkeypatch): + """A per-user (BYOM) model only resolves when the owner is passed.""" + seen = {} + + def _resolve(model_id, user_id=None): + seen["model_id"] = model_id + seen["user_id"] = user_id + return "anthropic" + + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", _resolve + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "k" + ) + self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("byom-uuid", "owner-7", "req-1") + + assert seen == {"model_id": "byom-uuid", "user_id": "owner-7"} + + def test_unknown_model_falls_back_to_the_configured_provider(self, monkeypatch): + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt") + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", lambda *a, **k: None + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "fallback-key" + ) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("mystery-model", "owner-1", "req-1") + + assert captured["provider"] == "docsgpt" + assert captured["kwargs"]["api_key"] == "fallback-key" + + def test_no_model_id_skips_the_lookup(self, monkeypatch): + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "openai") + calls = [] + monkeypatch.setattr( + extraction_module, + "get_provider_from_model_id", + lambda *a, **k: calls.append(a) or "anthropic", + ) + monkeypatch.setattr( + extraction_module, "get_api_key_for_provider", lambda provider: "k" + ) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm(None, "owner-1", "req-1") + + assert calls == [] + assert captured["provider"] == "openai" + + def test_api_key_follows_the_resolved_provider(self, monkeypatch): + """The key must match the provider actually dispatched to.""" + monkeypatch.setattr(extraction_module.settings, "LLM_PROVIDER", "docsgpt") + monkeypatch.setattr(extraction_module.settings, "API_KEY", "generic-key") + monkeypatch.setattr( + extraction_module, "get_provider_from_model_id", lambda *a, **k: "anthropic" + ) + keyed_for = {} + + def _key(provider): + keyed_for["provider"] = provider + return "sk-anthropic" + + monkeypatch.setattr(extraction_module, "get_api_key_for_provider", _key) + captured = self._capture_create_llm(monkeypatch) + + extraction_module._build_extraction_llm("claude-x", "owner-1", "req-1") + + assert keyed_for["provider"] == "anthropic" + assert captured["kwargs"]["api_key"] == "sk-anthropic" + assert captured["kwargs"]["api_key"] != "generic-key" + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): From e3d819d9fdbbcd3f9eaffbec2b2a8cce5a822eac Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:44:23 +0100 Subject: [PATCH 077/130] fix(graphrag): stop losing chunks silently during a graph build MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three ways a chunk disappeared from a graph with no way to tell: A build checks out one connection and then spends minutes per chunk waiting on the model, so the connection idles long enough for the server or a pooler to drop it. The pool only validates a connection when it hands one out, and this one was handed out at the start of the build, so the next write raised "the connection is lost", the chunk was marked failed, and the build carried on a chunk short. apply_chunk and mark_chunk now reconnect and retry once; every statement they run is an idempotent upsert, so a replay cannot double-write. Only connection loss retries — a bad statement still surfaces. An unparseable model response marked the chunk failed and logged nothing at all, so failed_chunks was the only evidence and it named no chunk. Both failure modes now log the chunk id. The summary's node count summed per-chunk upserts, so an entity appearing in ten chunks counted ten times: it reported writes, not graph size. It now reports the distinct node count, falling back to the write count only if the count query fails. --- docsgpt/graphrag/extraction.py | 49 ++++++-- docsgpt/graphrag/store.py | 188 ++++++++++++++++++++---------- tests/graphrag/test_extraction.py | 103 ++++++++++++++++ tests/graphrag/test_store.py | 112 ++++++++++++++++++ 4 files changed, 379 insertions(+), 73 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index aa76189f..580b54cb 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -149,8 +149,15 @@ def _parse_extraction(raw: Any) -> Optional[Dict[str, List[Dict[str, Any]]]]: } -def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]: - """Run exactly one extraction call for a chunk (gleanings off).""" +def _extract_chunk( + llm, text: str, chunk_id: Optional[str] = None +) -> Optional[Dict[str, List[Dict[str, Any]]]]: + """Run exactly one extraction call for a chunk (gleanings off). + + Both failure modes name the chunk: an unparseable response used to return + ``None`` silently, so a graph could come back short with nothing in the + logs to say which chunk was dropped or why. + """ messages = [ {"role": "system", "content": _SYSTEM_PROMPT}, {"role": "user", "content": f"\n{text}\n"}, @@ -161,9 +168,17 @@ def _extract_chunk(llm, text: str) -> Optional[Dict[str, List[Dict[str, Any]]]]: messages=messages, ) except Exception as exc: - logger.warning("Graph extraction call failed, skipping chunk: %s", exc) + logger.warning( + "Graph extraction call failed for chunk %s, skipping: %s", chunk_id, exc + ) return None - return _parse_extraction(response) + parsed = _parse_extraction(response) + if parsed is None: + logger.warning( + "Graph extraction returned unparseable output for chunk %s; marking it failed.", + chunk_id, + ) + return parsed def _coerce_weight(value: Any) -> float: @@ -205,7 +220,9 @@ def extract_graph_for_source( Returns: A summary ``{nodes, edges, chunks_processed, skipped_over_cap, - failed_chunks}``. + failed_chunks}``, where ``nodes`` is how many distinct nodes the + source's graph holds after the run — not how many upserts ran, which + counts the same entity once per chunk it appears in. """ from docsgpt.graphrag.store import GraphStore @@ -228,7 +245,7 @@ def extract_graph_for_source( _resolve_extraction_model(config), user, request_id ) - nodes = 0 + node_upserts = 0 edges = 0 chunks_processed = 0 failed_chunks = 0 @@ -242,7 +259,7 @@ def extract_graph_for_source( { "current": chunks_processed + failed_chunks, "total": total, - "nodes": nodes, + "nodes": node_upserts, "edges": edges, } ) @@ -257,7 +274,7 @@ def extract_graph_for_source( _report() continue - extracted = _extract_chunk(llm, text) + extracted = _extract_chunk(llm, text, chunk_id) if extracted is None: store.mark_chunk(source_id, chunk_id, "failed") failed_chunks += 1 @@ -271,7 +288,7 @@ def extract_graph_for_source( chunk_nodes, chunk_edges = store.apply_chunk( source_id, chunk_id, entities, relationships, name_embeddings ) - nodes += chunk_nodes + node_upserts += chunk_nodes edges += chunk_edges store.mark_chunk(source_id, chunk_id, "done") chunks_processed += 1 @@ -290,6 +307,20 @@ def extract_graph_for_source( except Exception as exc: logger.warning("set_node_degrees failed for source %s: %s", source_id, exc) + # Upserts are writes, not nodes: one entity seen in ten chunks is ten + # upserts and a single node, so the old count overstated every graph whose + # entities recur. Report what the graph holds, falling back to the write + # count only if the count query itself fails. + nodes = node_upserts + try: + nodes = store.count_nodes(source_id) + except Exception as exc: + logger.warning( + "count_nodes failed for source %s; reporting upserts instead: %s", + source_id, + exc, + ) + return { "nodes": nodes, "edges": edges, diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 818004c0..ceb0fc48 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -17,6 +17,7 @@ import logging import uuid from typing import Any, Dict, List, Optional +import psycopg from psycopg.types.json import Jsonb from docsgpt.core.settings import settings @@ -73,6 +74,25 @@ def _pgvector_identifiers() -> tuple[str, str, str, str]: ) +def _is_connection_lost(exc: BaseException) -> bool: + """True when ``exc`` says the server connection went away, not that the SQL was bad. + + psycopg raises ``OperationalError`` ("the connection is lost") when the + socket dies under a statement and ``InterfaceError`` when the connection + object is already closed. Everything else — a bad statement, a constraint + violation — is a real failure that a retry would only repeat. + """ + return isinstance(exc, (psycopg.OperationalError, psycopg.InterfaceError)) + + +def _safe_rollback(conn) -> None: + """Roll back, tolerating a connection too broken to roll back.""" + try: + conn.rollback() + except Exception as exc: + logging.debug("Rollback on a broken connection failed: %s", exc) + + class GraphStore: """Stores and queries a per-source knowledge graph in the pgvector DB.""" @@ -138,6 +158,37 @@ class GraphStore: self._pooled = False return self._connection + def _write_with_reconnect(self, operation): + """Run ``operation(conn)``, once more on a fresh connection if it was dead. + + A graph build holds one checked-out connection for the length of the + whole extraction and spends minutes per chunk waiting on the model, so + the connection idles long enough for the server (or a pooler) to drop + it. The pool only validates a connection when it is handed out, and + this one was handed out at the start of the build, so the next write + raises and its chunk is lost from the graph. Every statement here is an + idempotent upsert, so replaying one on a new connection cannot + double-write. + + Args: + operation: Callable taking the connection and doing one write. + + Returns: + Whatever ``operation`` returns. + """ + for attempt in (1, 2): + conn = self._get_connection() + try: + return operation(conn) + except Exception as exc: + if attempt == 2 or not _is_connection_lost(exc): + raise + logging.warning( + "Graph write lost its connection (%s); reconnecting and retrying once.", + exc, + ) + self.close() + def _register_pgvector_types(self, conn) -> None: """Register pgvector's adapters, tolerating a not-yet-created extension. @@ -499,56 +550,61 @@ class GraphStore: (not linked to the chunk), mirroring the per-call path. ``name_embeddings`` maps ``normalized_name`` to its embedding. Degrees are not bumped here — the caller runs ``set_node_degrees`` once at the - end. Returns ``(nodes_upserted, edges_added)``. + end. Reconnects and retries once if the connection died while the + extraction was waiting on the model. Returns + ``(nodes_upserted, edges_added)``. """ self._ensure_tables_once() - conn = self._get_connection() - cursor = conn.cursor() - node_ids: Dict[str, str] = {} - edges_added = 0 - try: - for entity in entities: - normalized_name = entity["normalized_name"] - node_id = self._upsert_node( - cursor, - source_id, - entity["name"], - normalized_name, - entity.get("type"), - entity.get("description"), - name_embeddings.get(normalized_name), - ) - node_ids[normalized_name] = node_id - self._link_node_chunk(cursor, source_id, node_id, chunk_id) - for rel in relationships: - src_id = self._resolve_endpoint( - cursor, source_id, rel.get("source"), node_ids, name_embeddings - ) - dst_id = self._resolve_endpoint( - cursor, source_id, rel.get("target"), node_ids, name_embeddings - ) - if src_id is None or dst_id is None: - continue - self._add_edge( - cursor, - source_id, - src_id, - dst_id, - type=rel.get("type"), - description=rel.get("description"), - weight=float(rel.get("weight") or 1.0), - source_chunk_ids=[chunk_id], - ) - edges_added += 1 + def _write(conn): + cursor = conn.cursor() + node_ids: Dict[str, str] = {} + edges_added = 0 + try: + for entity in entities: + normalized_name = entity["normalized_name"] + node_id = self._upsert_node( + cursor, + source_id, + entity["name"], + normalized_name, + entity.get("type"), + entity.get("description"), + name_embeddings.get(normalized_name), + ) + node_ids[normalized_name] = node_id + self._link_node_chunk(cursor, source_id, node_id, chunk_id) - conn.commit() - return len(entities), edges_added - except Exception: - conn.rollback() - raise - finally: - cursor.close() + for rel in relationships: + src_id = self._resolve_endpoint( + cursor, source_id, rel.get("source"), node_ids, name_embeddings + ) + dst_id = self._resolve_endpoint( + cursor, source_id, rel.get("target"), node_ids, name_embeddings + ) + if src_id is None or dst_id is None: + continue + self._add_edge( + cursor, + source_id, + src_id, + dst_id, + type=rel.get("type"), + description=rel.get("description"), + weight=float(rel.get("weight") or 1.0), + source_chunk_ids=[chunk_id], + ) + edges_added += 1 + + conn.commit() + return len(entities), edges_added + except Exception: + _safe_rollback(conn) + raise + finally: + cursor.close() + + return self._write_with_reconnect(_write) def _resolve_endpoint( self, @@ -1026,25 +1082,29 @@ class GraphStore: cursor.close() def mark_chunk(self, source_id: str, chunk_id: str, status: str): + """Record a chunk's extraction status, reconnecting once if the connection died.""" self._ensure_tables_once() - conn = self._get_connection() - cursor = conn.cursor() - try: - cursor.execute( - """ - INSERT INTO graph_ingest_progress (source_id, chunk_id, status) - VALUES (%s, %s, %s) - ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status; - """, - (source_id, str(chunk_id), status), - ) - conn.commit() - except Exception as e: - conn.rollback() - logging.error(f"Error marking chunk: {e}") - raise - finally: - cursor.close() + + def _write(conn): + cursor = conn.cursor() + try: + cursor.execute( + """ + INSERT INTO graph_ingest_progress (source_id, chunk_id, status) + VALUES (%s, %s, %s) + ON CONFLICT (source_id, chunk_id) DO UPDATE SET status = EXCLUDED.status; + """, + (source_id, str(chunk_id), status), + ) + conn.commit() + except Exception as e: + _safe_rollback(conn) + logging.error(f"Error marking chunk: {e}") + raise + finally: + cursor.close() + + return self._write_with_reconnect(_write) def pending_chunks(self, source_id: str, all_chunk_ids: List[str]) -> List[str]: """Chunk ids from ``all_chunk_ids`` not yet marked ``done`` for the source.""" diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index eaaf881d..b6192450 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -521,6 +521,109 @@ class TestExtractionProviderResolution: assert captured["kwargs"]["api_key"] != "generic-key" +@pytest.mark.unit +class TestFailedChunksAreReported: + """Every dropped chunk has to leave a trace. + + A chunk whose extraction cannot be parsed is marked ``failed`` and skipped. + That path logged nothing at all, so a graph could come back short with the + summary's ``failed_chunks`` count as the only hint and no way to tell which + chunk, or why, from the logs. + """ + + def _fake_store(self, monkeypatch, chunk_ids): + from unittest.mock import MagicMock + + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = list(chunk_ids) + store.apply_chunk.return_value = (1, 0) + store.count_nodes.return_value = 1 + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda *a, **k: store + ) + return store + + def test_unparseable_output_is_logged_with_the_chunk_id( + self, monkeypatch, caplog, stub_embedding + ): + import logging + + store = self._fake_store(monkeypatch, ["c1"]) + _install_stub_llm(monkeypatch, _StubLLM(["not json at all"])) + + with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): + summary = extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk("c1", "some text")], + config=SourceConfig(), + request_id="req-1", + ) + + assert summary["failed_chunks"] == 1 + store.mark_chunk.assert_called_once() + assert store.mark_chunk.call_args.args[2] == "failed" + messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING] + assert any("c1" in message for message in messages), messages + + def test_llm_errors_still_name_the_chunk( + self, monkeypatch, caplog, stub_embedding + ): + import logging + + self._fake_store(monkeypatch, ["c7"]) + _install_stub_llm(monkeypatch, _StubLLM([RuntimeError("model exploded")])) + + with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): + extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk("c7", "some text")], + config=SourceConfig(), + request_id="req-1", + ) + + messages = [r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING] + assert any("c7" in message for message in messages), messages + + +@pytest.mark.integration +class TestSummaryNodeCount: + """``nodes`` must describe the graph, not the number of upserts.""" + + @pytest.fixture + def store(self, monkeypatch, postgresql): + store = _live_store(monkeypatch, postgresql.info) + yield store + store.close() + + def test_repeated_entity_counts_once( + self, store, monkeypatch, stub_embedding + ): + source_id = str(uuid.uuid4()) + try: + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + _install_stub_llm(monkeypatch, _StubLLM([payload, payload])) + + summary = extract_graph_for_source( + source_id, + user="owner-1", + chunks=[_chunk("c1", "Ada one."), _chunk("c2", "Ada two.")], + config=SourceConfig(), + request_id="req-1", + ) + + # Two chunks upserted the same entity: one node in the graph. + assert store.count_nodes(source_id) == 1 + assert summary["nodes"] == 1 + assert summary["chunks_processed"] == 2 + finally: + store.delete_by_source(source_id) + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index d712e7fe..f2ce6abe 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -890,3 +890,115 @@ class TestCountNodesMany: store, _, _ = self._store_with_mock_conn([(source_id.lower(), 3)]) assert store.count_nodes_many([source_id]) == {source_id: 3} + + +@pytest.mark.unit +class TestWritesSurviveALostConnection: + """A graph build holds one pooled connection across its LLM calls. + + Extraction spends minutes per chunk waiting on a model, so the connection + sits idle between writes and the server (or a pooler) can drop it. The pool + only validates a connection at checkout, and this one was checked out once + at the start of the build, so the next write raises and the chunk is marked + ``failed`` — silently losing it from the graph. The write reconnects and + retries once instead; the statements are idempotent upserts, so a retry + cannot double-write. + """ + + def _store_with_connections(self, conns): + """Store that hands out ``conns`` in order, one per (re)connect.""" + store = GraphStore.__new__(GraphStore) + store._tables_ensured = True + store._connection = None + handed = [] + closed = [] + + def _get_connection(): + if store._connection is None: + store._connection = conns[len(handed)] + handed.append(store._connection) + return store._connection + + def _close(): + if store._connection is not None: + closed.append(store._connection) + store._connection = None + + store._get_connection = _get_connection + store.close = _close + return store, handed, closed + + @staticmethod + def _conn(execute_error=None): + cursor = MagicMock() + cursor.fetchone.return_value = [str(uuid.uuid4())] + cursor.fetchall.return_value = [] + if execute_error is not None: + cursor.execute.side_effect = execute_error + conn = MagicMock() + conn.cursor.return_value = cursor + return conn + + def test_mark_chunk_retries_on_a_dropped_connection(self): + import psycopg + + dead = self._conn(psycopg.OperationalError("the connection is lost")) + alive = self._conn() + store, handed, closed = self._store_with_connections([dead, alive]) + + store.mark_chunk(str(uuid.uuid4()), "c1", "done") + + assert handed == [dead, alive] + assert closed == [dead] + alive.commit.assert_called_once() + + def test_apply_chunk_retries_on_a_dropped_connection(self): + import psycopg + + dead = self._conn(psycopg.OperationalError("the connection is lost")) + alive = self._conn() + store, handed, closed = self._store_with_connections([dead, alive]) + entities = [ + { + "name": "Ada", + "normalized_name": "ada", + "type": "person", + "description": "d", + } + ] + + nodes, edges = store.apply_chunk( + str(uuid.uuid4()), "c1", entities, [], {"ada": _embedding(0.5)} + ) + + assert (nodes, edges) == (1, 0) + assert handed == [dead, alive] + assert closed == [dead] + alive.commit.assert_called_once() + + def test_a_second_connection_failure_is_not_retried_again(self): + """One retry, not a loop: a genuinely unreachable DB still fails.""" + import psycopg + + dead = self._conn(psycopg.OperationalError("the connection is lost")) + also_dead = self._conn(psycopg.OperationalError("the connection is lost")) + store, handed, _ = self._store_with_connections([dead, also_dead]) + + with pytest.raises(psycopg.OperationalError): + store.mark_chunk(str(uuid.uuid4()), "c1", "done") + + assert handed == [dead, also_dead] + + def test_a_query_error_is_not_retried(self): + """Only connection loss is retryable; a bad statement must surface.""" + import psycopg + + broken = self._conn(psycopg.ProgrammingError("syntax error")) + spare = self._conn() + store, handed, _ = self._store_with_connections([broken, spare]) + + with pytest.raises(psycopg.ProgrammingError): + store.mark_chunk(str(uuid.uuid4()), "c1", "done") + + assert handed == [broken] + broken.rollback.assert_called_once() From 6ea737e57b3fd2bf70ed24225ccadca7c37b125d Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 15:49:56 +0100 Subject: [PATCH 078/130] fix(tasks): a deferred duplicate stands down instead of failing the task MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A task that runs longer than the broker's visibility timeout is redelivered while its first run is still going. The idempotency lease correctly stops the duplicate from doing the work, but the duplicate then re-queued itself once per LEASE_TTL until celery ran out of retries and raised MaxRetriesExceededError — so a perfectly healthy long task (a large graph extraction is the one that found this) reported a task failure, with a traceback, while the real run was still making progress next to it. Catch the exhaustion and return a "deferred" result instead. A normal deferral still re-queues: only the give-up path changes, and the lease holder's dedup row is left untouched so its own completion still records. --- docsgpt/api/user/idempotency.py | 28 ++++++-- tests/api/user/test_idempotency_decorator.py | 69 ++++++++++++++++++++ 2 files changed, 93 insertions(+), 4 deletions(-) diff --git a/docsgpt/api/user/idempotency.py b/docsgpt/api/user/idempotency.py index 1381f241..e38e8a73 100644 --- a/docsgpt/api/user/idempotency.py +++ b/docsgpt/api/user/idempotency.py @@ -9,6 +9,8 @@ import threading import uuid from typing import Any, Callable, Optional +from celery.exceptions import MaxRetriesExceededError + from docsgpt.storage.db.repositories.idempotency import IdempotencyRepository from docsgpt.storage.db.session import db_readonly, db_session @@ -81,10 +83,28 @@ def with_idempotency( "idempotency: live lease held; deferring task=%s key=%s", task_name, key, ) - raise self.retry( - countdown=LEASE_TTL_SECONDS, - max_retries=LEASE_RETRY_MAX, - ) + try: + raise self.retry( + countdown=LEASE_TTL_SECONDS, + max_retries=LEASE_RETRY_MAX, + ) + except MaxRetriesExceededError: + # The holder is simply slower than LEASE_RETRY_MAX + # deferrals — a task that outruns the broker's visibility + # timeout is redelivered while its first run is still + # going. Standing down is the correct end state for the + # duplicate; raising here would report a failure for a + # task that is running normally somewhere else. + logger.info( + "idempotency: lease still held after %s deferrals; " + "leaving task=%s key=%s to its holder", + LEASE_RETRY_MAX, task_name, key, + ) + return { + "status": "deferred", + "reason": "another worker holds the lease", + "idempotency_key": key, + } if attempt > MAX_TASK_ATTEMPTS: logger.error( diff --git a/tests/api/user/test_idempotency_decorator.py b/tests/api/user/test_idempotency_decorator.py index 35dfb9dd..b1d26f93 100644 --- a/tests/api/user/test_idempotency_decorator.py +++ b/tests/api/user/test_idempotency_decorator.py @@ -335,6 +335,75 @@ class TestLiveLeaseDefersConcurrentRun: assert row[2] == "completed" +@pytest.mark.unit +class TestLeaseDeferralGivesUpQuietly: + """Deferral is bookkeeping, not failure. + + A task that outruns the broker's visibility timeout is redelivered while + the first worker is still running it. The lease keeps the duplicate from + doing the work, but the duplicate kept re-queueing itself until celery + exhausted ``LEASE_RETRY_MAX`` and raised ``MaxRetriesExceededError``, so a + healthy long task logged a task failure. The duplicate should stand down + instead and leave the run to the worker that holds the lease. + """ + + def _hold_lease(self, pg_conn, key): + from docsgpt.storage.db.repositories.idempotency import ( + IdempotencyRepository, + ) + + IdempotencyRepository(pg_conn).try_claim_lease( + key=key, task_name="thing", + task_id="t-worker-1", owner_id="worker-1", + ) + + def test_exhausted_retries_return_deferred_instead_of_raising(self, pg_conn): + from celery.exceptions import MaxRetriesExceededError + + from docsgpt.api.user.idempotency import with_idempotency + + self._hold_lease(pg_conn, "k-long-run") + invocations = {"count": 0} + + @with_idempotency(task_name="thing") + def task(self, idempotency_key=None): + invocations["count"] += 1 + return {"ran": True} + + # Celery raises this from ``self.retry`` once max_retries is hit. + worker2 = _fake_celery_self("t-worker-2") + worker2.retry.side_effect = MaxRetriesExceededError("out of retries") + + with _patch_decorator_db(pg_conn): + result = task(worker2, idempotency_key="k-long-run") + + assert result["status"] == "deferred" + # The lease holder is still running it; the duplicate did not. + assert invocations["count"] == 0 + # The holder's row is untouched — not failed, not completed. + row = _row_for(pg_conn, "k-long-run") + assert row[2] == "pending" + + def test_a_normal_retry_still_propagates(self, pg_conn): + """Only exhaustion stands down; the first deferrals must re-queue.""" + from docsgpt.api.user.idempotency import with_idempotency + + self._hold_lease(pg_conn, "k-busy-once") + + @with_idempotency(task_name="thing") + def task(self, idempotency_key=None): + return {"ran": True} + + class _RetrySignal(Exception): + pass + + worker2 = _fake_celery_self("t-worker-2") + worker2.retry.side_effect = _RetrySignal("retry scheduled") + + with _patch_decorator_db(pg_conn), pytest.raises(_RetrySignal): + task(worker2, idempotency_key="k-busy-once") + + @pytest.mark.unit class TestExceptionPathReleasesLease: """When ``fn`` raises, the lease is dropped so the next attempt From 9e9f130ed028b6e65e7db88b9707dbf88f8ac084 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 16:03:39 +0100 Subject: [PATCH 079/130] refactor(graphrag): spell the write retry out, and log a capped PPR run MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Review feedback on _write_with_reconnect: an explicit return inside the loop plus an implicit fall-through past it reads as a path that returns None. The retry is exactly two attempts, so say so — first attempt, reconnect on connection loss, second and final attempt — and every path now returns or raises. Also narrows the transitions hint to the (node, weight) tuples it holds, and logs at debug when the power iteration hits its iteration cap instead of returning the last iterate with no trace. --- docsgpt/graphrag/store.py | 29 +++++++++++++++++------------ docsgpt/retriever/graph_rag.py | 9 ++++++++- 2 files changed, 25 insertions(+), 13 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index ceb0fc48..116d3611 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -175,19 +175,24 @@ class GraphStore: Returns: Whatever ``operation`` returns. + + Raises: + Exception: Anything ``operation`` raises that is not connection + loss, and anything the single retry raises. """ - for attempt in (1, 2): - conn = self._get_connection() - try: - return operation(conn) - except Exception as exc: - if attempt == 2 or not _is_connection_lost(exc): - raise - logging.warning( - "Graph write lost its connection (%s); reconnecting and retrying once.", - exc, - ) - self.close() + try: + return operation(self._get_connection()) + except Exception as exc: + if not _is_connection_lost(exc): + raise + logging.warning( + "Graph write lost its connection (%s); reconnecting and retrying once.", + exc, + ) + self.close() + # Second and final attempt, on a connection freshly checked out by + # ``_get_connection``. A failure here belongs to the caller. + return operation(self._get_connection()) def _register_pgvector_types(self, conn) -> None: """Register pgvector's adapters, tolerating a not-yet-created extension. diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index 5593a9b2..ef4b0ea5 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -103,7 +103,7 @@ def _personalized_pagerank( # Row-normalized transitions. An undirected edge is traversable from both # endpoints, so each node normalizes over its own incident weights. - transitions: Dict[Any, List[Any]] = {} + transitions: Dict[Any, List[tuple[Any, float]]] = {} for node in nodes: neighbors = [] total = 0.0 @@ -134,6 +134,13 @@ def _personalized_pagerank( ranks[node] += (leaked + 1.0 - alpha) * restart[node] if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol: break + else: + logging.debug( + "Personalized PageRank hit its %s-iteration cap on a %s-node " + "subgraph; ranking with the last iterate.", + max_iter, + node_count, + ) return ranks From 3f774d813c5dd22e6d4c0fc6602093570ee590e2 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 16:31:19 +0100 Subject: [PATCH 080/130] =?UTF-8?q?fix(graphrag):=20review=20fixes=20?= =?UTF-8?q?=E2=80=94=20replay-safe=20writes,=20strict=20count,=20zero=20we?= =?UTF-8?q?ights?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three findings from review, all on code this branch introduced. apply_chunk was not replay-safe. commit() can report connection loss *after* Postgres committed, and the reconnect retry then replays the write: _upsert_node bumps doc_freq a second time and _add_edge inserts another row, since graph_edges has no uniqueness constraint for a logical edge. The chunk's graph_ingest_progress row is now written in the same transaction as the rows it describes, and a replay that finds it already "done" returns (0, 0) without touching the graph. Extraction drops its separate mark_chunk("done"): the checkpoint and the graph can no longer disagree. count_nodes swallows every query failure and answers 0, so extraction's "fall back to the write count" handler could never run — a failed count after a successful build reported an empty graph. count_nodes grows a strict mode that re-raises; retrieval keeps the swallow, which is what routes a source to ClassicRAG. A zero edge weight was read as a full-strength link: `or 1.0` rewrote an explicit 0 before the <= 0 filter. Only missing and null weights default now. The same coercion sat in _ppr_scores, where it would have kept the ranker's rule unreachable from the product path, so it is fixed there too. --- docsgpt/graphrag/extraction.py | 9 ++- docsgpt/graphrag/store.py | 52 ++++++++++++-- docsgpt/retriever/graph_rag.py | 14 +++- tests/graphrag/test_extraction.py | 36 ++++++++++ tests/graphrag/test_store.py | 112 ++++++++++++++++++++++++++++++ tests/retriever/test_graph_rag.py | 60 ++++++++++++++++ 6 files changed, 274 insertions(+), 9 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 580b54cb..1bdfb054 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -290,7 +290,9 @@ def extract_graph_for_source( ) node_upserts += chunk_nodes edges += chunk_edges - store.mark_chunk(source_id, chunk_id, "done") + # ``apply_chunk`` marks the chunk done inside the transaction that + # writes its rows, so the checkpoint cannot disagree with the graph + # and a replayed write cannot apply the chunk twice. chunks_processed += 1 except Exception as exc: logger.warning( @@ -311,9 +313,12 @@ def extract_graph_for_source( # upserts and a single node, so the old count overstated every graph whose # entities recur. Report what the graph holds, falling back to the write # count only if the count query itself fails. + # ``strict`` is what makes the fallback below reachable: the default + # count swallows query failures and answers 0, which would report a + # successful build as an empty graph. nodes = node_upserts try: - nodes = store.count_nodes(source_id) + nodes = store.count_nodes(source_id, strict=True) except Exception as exc: logger.warning( "count_nodes failed for source %s; reporting upserts instead: %s", diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 116d3611..acb35911 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -556,8 +556,12 @@ class GraphStore: ``name_embeddings`` maps ``normalized_name`` to its embedding. Degrees are not bumped here — the caller runs ``set_node_degrees`` once at the end. Reconnects and retries once if the connection died while the - extraction was waiting on the model. Returns - ``(nodes_upserted, edges_added)``. + extraction was waiting on the model. + + The chunk's ``graph_ingest_progress`` row is written in this same + transaction, so the checkpoint and the rows it describes commit + together and a replay of an already-applied chunk returns ``(0, 0)`` + without touching the graph. Returns ``(nodes_upserted, edges_added)``. """ self._ensure_tables_once() @@ -566,6 +570,21 @@ class GraphStore: node_ids: Dict[str, str] = {} edges_added = 0 try: + # ``commit()`` can report connection loss *after* the server + # committed, and the retry then replays this write: doc_freq + # would be bumped twice and a second logical edge inserted + # (graph_edges has no uniqueness constraint). The progress row + # below is written in this transaction, so a replay sees it. + cursor.execute( + "SELECT status FROM graph_ingest_progress " + "WHERE source_id = %s AND chunk_id = %s;", + (source_id, str(chunk_id)), + ) + applied = cursor.fetchone() + if applied is not None and applied[0] == "done": + conn.rollback() + return 0, 0 + for entity in entities: normalized_name = entity["normalized_name"] node_id = self._upsert_node( @@ -601,6 +620,15 @@ class GraphStore: ) edges_added += 1 + cursor.execute( + """ + INSERT INTO graph_ingest_progress (source_id, chunk_id, status) + VALUES (%s, %s, 'done') + ON CONFLICT (source_id, chunk_id) + DO UPDATE SET status = EXCLUDED.status; + """, + (source_id, str(chunk_id)), + ) conn.commit() return len(entities), edges_added except Exception: @@ -671,8 +699,22 @@ class GraphStore: cursor.close() conn.rollback() - def count_nodes(self, source_id: str) -> int: - """Number of nodes for a source. Zero drives the ClassicRAG fallback.""" + def count_nodes(self, source_id: str, strict: bool = False) -> int: + """Number of nodes for a source. Zero drives the ClassicRAG fallback. + + Args: + source_id: Source whose nodes to count. + strict: Re-raise a query failure instead of reporting ``0``. + Retrieval wants the swallow — a broken count there just routes + the source to ClassicRAG — but a caller reporting how big a + graph is must not read a failed query as "the graph is empty". + + Returns: + int: The node count, or ``0`` when a query failure is swallowed. + + Raises: + Exception: The underlying query failure, when ``strict`` is set. + """ conn = self._get_connection() cursor = conn.cursor() try: @@ -683,6 +725,8 @@ class GraphStore: return int(cursor.fetchone()[0]) except Exception as e: logging.error(f"Error counting nodes: {e}") + if strict: + raise return 0 finally: cursor.close() diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index ef4b0ea5..81ed9bf0 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -108,7 +108,11 @@ def _personalized_pagerank( neighbors = [] total = 0.0 for neighbor, data in graph[node].items(): - edge_weight = float(data.get(weight, 1.0) or 1.0) + raw_weight = data.get(weight, 1.0) + # Default only a missing or null weight. ``or 1.0`` would also + # rewrite an explicit 0 — "these entities are not related" — into a + # full-strength transition, which changes the ranking. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) if edge_weight <= 0: continue neighbors.append((neighbor, edge_weight)) @@ -212,8 +216,12 @@ class GraphRAGRetriever(BaseRetriever): for edge in subgraph.get("edges", []): src, dst = edge["src_node_id"], edge["dst_node_id"] if src in graph and dst in graph: - weight = float(edge.get("weight") or 1.0) - graph.add_edge(src, dst, weight=weight) + raw_weight = edge.get("weight") + # Same rule the ranker applies: default only a missing or null + # weight. Coercing an explicit 0 to 1.0 here would make "these + # entities are not related" the strongest possible link. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) + graph.add_edge(src, dst, weight=edge_weight) if graph.number_of_nodes() == 0: return {} diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index b6192450..48e567d4 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -624,6 +624,42 @@ class TestSummaryNodeCount: store.delete_by_source(source_id) +@pytest.mark.unit +class TestSummaryCountFailure: + """A broken count query must not be reported as an empty graph.""" + + def test_a_failed_count_reports_the_write_count( + self, monkeypatch, stub_embedding + ): + from unittest.mock import MagicMock + + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = ["c1"] + store.apply_chunk.return_value = (2, 1) + store.count_nodes.side_effect = RuntimeError("count query failed") + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda *a, **k: store + ) + _install_stub_llm( + monkeypatch, + _StubLLM([_extraction_json([{"name": "Ada"}], [])]), + ) + + summary = extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk("c1", "Ada.")], + config=SourceConfig(), + request_id="req-1", + ) + + # Falls back to what was actually written, not to zero. + assert summary["nodes"] == 2 + # And it asked for a count that raises rather than one that returns 0, + # or the fallback above could never run. + assert store.count_nodes.call_args.kwargs.get("strict") is True + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index f2ce6abe..baa19604 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -1002,3 +1002,115 @@ class TestWritesSurviveALostConnection: assert handed == [broken] broken.rollback.assert_called_once() + + +@pytest.mark.unit +class TestCountNodesFailureModes: + """Retrieval wants a swallowed count; extraction wants to hear about it.""" + + def _store_with_failing_cursor(self): + store = GraphStore.__new__(GraphStore) + store._tables_ensured = True + cursor = MagicMock() + cursor.execute.side_effect = RuntimeError("relation does not exist") + conn = MagicMock() + conn.cursor.return_value = cursor + store._connection = conn + store._get_connection = lambda: conn + return store + + def test_default_reports_zero_to_drive_the_classic_fallback(self): + store = self._store_with_failing_cursor() + + assert store.count_nodes(str(uuid.uuid4())) == 0 + + def test_strict_surfaces_the_query_failure(self): + """A caller reporting graph size must not read a broken query as empty.""" + store = self._store_with_failing_cursor() + + with pytest.raises(RuntimeError): + store.count_nodes(str(uuid.uuid4()), strict=True) + + +@pytest.mark.integration +class TestApplyChunkIsReplaySafe: + """A retry after an ambiguous commit must not apply a chunk twice. + + ``_write_with_reconnect`` replays the write when the connection dies, and + ``commit()`` itself can raise connection loss *after* the server committed. + Replaying then bumps ``doc_freq`` a second time and inserts a second + logical edge (``graph_edges`` has no uniqueness constraint), so the chunk's + own progress row is written in the same transaction and short-circuits it. + """ + + @pytest.fixture + def store(self, postgresql): + store = GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) + try: + store._ensure_tables() + except Exception as exc: + pytest.skip(f"pgvector extension unavailable: {exc}") + yield store + store.close() + + def test_a_replayed_chunk_is_not_applied_twice(self, store): + source_id = str(uuid.uuid4()) + entities = [ + { + "name": "Ada", + "normalized_name": "ada", + "type": "person", + "description": "d", + } + ] + relationships = [ + { + "source": "Ada", + "target": "Engine", + "type": "worked_on", + "description": "x", + "weight": 2.0, + } + ] + embeddings = {"ada": _embedding(0.1), "engine": _embedding(0.2)} + try: + first = store.apply_chunk( + source_id, "c1", entities, relationships, embeddings + ) + replay = store.apply_chunk( + source_id, "c1", entities, relationships, embeddings + ) + + assert first == (1, 1) + assert replay == (0, 0) + node = store.get_node_by_normalized(source_id, "ada") + assert node["doc_freq"] == 1 + overview = store.get_graph_overview(source_id) + assert len(overview["edges"]) == 1 + # The write records its own progress, so the caller's checkpoint + # and the rows it describes commit together. + assert store.get_progress(source_id)["c1"] == "done" + finally: + store.delete_by_source(source_id) + + def test_a_different_chunk_still_applies(self, store): + """The guard is per chunk, not a blanket 'already saw this source'.""" + source_id = str(uuid.uuid4()) + entities = [ + { + "name": "Ada", + "normalized_name": "ada", + "type": "person", + "description": "d", + } + ] + embeddings = {"ada": _embedding(0.1)} + try: + store.apply_chunk(source_id, "c1", entities, [], embeddings) + second = store.apply_chunk(source_id, "c2", entities, [], embeddings) + + assert second == (1, 0) + node = store.get_node_by_normalized(source_id, "ada") + assert node["doc_freq"] == 2 + finally: + store.delete_by_source(source_id) diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index 5f8cdf80..8943174c 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -901,6 +901,66 @@ class TestPersonalizedPageRankWithoutScipy: assert _personalized_pagerank(nx.Graph(), personalization=None) == {} + def test_a_zero_weight_edge_is_not_traversable(self): + """Zero means "not related", not "use the default weight".""" + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + graph = nx.Graph() + graph.add_edge("seed", "zero", weight=0.0) + graph.add_edge("seed", "real", weight=1.0) + + ranks = _personalized_pagerank( + graph, personalization={"seed": 1.0, "zero": 0.0, "real": 0.0} + ) + + # ``zero`` is reachable only across the zero-weight edge, so no mass + # walks to it; ``real`` is on a live edge and must outrank it. + assert ranks["real"] > ranks["zero"] + assert ranks["zero"] == pytest.approx(0.0, abs=1e-9) + assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6) + + def test_stored_zero_weights_reach_the_ranker_intact(self): + """The subgraph builder must not coerce a stored 0 into a real edge. + + Without this the ranker's zero-weight rule is unreachable in + production: every 0 from ``graph_edges`` arrives as 1.0. + """ + subgraph = { + "nodes": [ + {"id": "seed", "doc_freq": 1}, + {"id": "zero", "doc_freq": 1}, + {"id": "real", "doc_freq": 1}, + ], + "edges": [ + {"src_node_id": "seed", "dst_node_id": "zero", "weight": 0}, + {"src_node_id": "seed", "dst_node_id": "real", "weight": 1.0}, + ], + } + # Called unbound with ``None`` for self: _ppr_scores reads no state. + scores = GraphRAGRetriever._ppr_scores(None, subgraph, {"seed": 1.0}) + + assert scores["real"] > scores["zero"] + assert scores["zero"] == pytest.approx(0.0, abs=1e-9) + + def test_missing_and_null_weights_default_to_one(self): + import networkx as nx + + from docsgpt.retriever.graph_rag import _personalized_pagerank + + absent = nx.Graph() + absent.add_edge("a", "b") # no weight attribute at all + null = nx.Graph() + null.add_edge("a", "b", weight=None) + + personalization = {"a": 1.0, "b": 0.0} + from_absent = _personalized_pagerank(absent, personalization=personalization) + from_null = _personalized_pagerank(null, personalization=personalization) + + assert from_absent["b"] == pytest.approx(from_null["b"], abs=1e-9) + assert from_absent["b"] > 0 + @patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10) @patch("docsgpt.retriever.graph_rag.GraphStore") @patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True) From b7e7872bf7722002aac7515f8af5960a579d7ddf Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 17 Sep 2026 16:51:06 +0100 Subject: [PATCH 081/130] fix(tasks): a deferred duplicate records no result, rather than success MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Returning a deferred marker traded one wrong signal for a worse one. A redelivery reuses the original task id — Context.as_execution_options carries task_id into the retry — so the duplicate's return marked the very id the client polls as SUCCESS. /api/task_status reports celery's state verbatim and the UI maps SUCCESS to "done", so the GraphRAG enable modal would announce a finished build, rendered from a payload with no counts, while the run holding the lease was still extracting. Raise Ignore instead: celery records no state for the duplicate, so the task id keeps whatever the holder sets and the poller keeps waiting. The autoretry wrapper re-raises Ignore ahead of autoretry_for, so the wider autoretry_for=(Exception,) on these tasks cannot turn it back into a retry. --- docsgpt/api/user/idempotency.py | 22 +++++++++++--------- tests/api/user/test_idempotency_decorator.py | 14 ++++++++----- 2 files changed, 21 insertions(+), 15 deletions(-) diff --git a/docsgpt/api/user/idempotency.py b/docsgpt/api/user/idempotency.py index e38e8a73..1cfc1b80 100644 --- a/docsgpt/api/user/idempotency.py +++ b/docsgpt/api/user/idempotency.py @@ -9,7 +9,7 @@ import threading import uuid from typing import Any, Callable, Optional -from celery.exceptions import MaxRetriesExceededError +from celery.exceptions import Ignore, MaxRetriesExceededError from docsgpt.storage.db.repositories.idempotency import IdempotencyRepository from docsgpt.storage.db.session import db_readonly, db_session @@ -90,21 +90,23 @@ def with_idempotency( ) except MaxRetriesExceededError: # The holder is simply slower than LEASE_RETRY_MAX - # deferrals — a task that outruns the broker's visibility + # deferrals: a task that outruns the broker's visibility # timeout is redelivered while its first run is still - # going. Standing down is the correct end state for the - # duplicate; raising here would report a failure for a - # task that is running normally somewhere else. + # going. Letting the exhaustion propagate would report a + # failure for a task that is running normally — but so + # would returning a value, only less visibly. A redelivery + # reuses the original task id (``Context`` carries + # ``task_id`` into the retry), so a return marks the very + # id the client polls SUCCESS, and ``/api/task_status`` + # hands that to the UI as a finished build. ``Ignore`` + # records no state at all, leaving the outcome to the run + # that actually holds the lease. logger.info( "idempotency: lease still held after %s deferrals; " "leaving task=%s key=%s to its holder", LEASE_RETRY_MAX, task_name, key, ) - return { - "status": "deferred", - "reason": "another worker holds the lease", - "idempotency_key": key, - } + raise Ignore() from None if attempt > MAX_TASK_ATTEMPTS: logger.error( diff --git a/tests/api/user/test_idempotency_decorator.py b/tests/api/user/test_idempotency_decorator.py index b1d26f93..cbed5f19 100644 --- a/tests/api/user/test_idempotency_decorator.py +++ b/tests/api/user/test_idempotency_decorator.py @@ -357,8 +357,8 @@ class TestLeaseDeferralGivesUpQuietly: task_id="t-worker-1", owner_id="worker-1", ) - def test_exhausted_retries_return_deferred_instead_of_raising(self, pg_conn): - from celery.exceptions import MaxRetriesExceededError + def test_exhausted_retries_stand_down_without_recording_a_result(self, pg_conn): + from celery.exceptions import Ignore, MaxRetriesExceededError from docsgpt.api.user.idempotency import with_idempotency @@ -374,10 +374,14 @@ class TestLeaseDeferralGivesUpQuietly: worker2 = _fake_celery_self("t-worker-2") worker2.retry.side_effect = MaxRetriesExceededError("out of retries") - with _patch_decorator_db(pg_conn): - result = task(worker2, idempotency_key="k-long-run") + # Ignore rather than a return value: a redelivery reuses the original + # task id, so returning would mark the id the client is polling + # SUCCESS — /api/task_status hands that straight to the UI, which + # would announce a finished (empty) build while the holder is still + # working. Ignore leaves the id's state to the holder. + with _patch_decorator_db(pg_conn), pytest.raises(Ignore): + task(worker2, idempotency_key="k-long-run") - assert result["status"] == "deferred" # The lease holder is still running it; the duplicate did not. assert invocations["count"] == 0 # The holder's row is untouched — not failed, not completed. From 762c7a19cd3af27e2e3a5708f960d8b9e50f5b5c Mon Sep 17 00:00:00 2001 From: Pavel Date: Thu, 17 Sep 2026 23:50:14 +0400 Subject: [PATCH 082/130] correct license for widget --- extensions/react-widget/package-lock.json | 2 +- extensions/react-widget/package.json | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/extensions/react-widget/package-lock.json b/extensions/react-widget/package-lock.json index 32fde950..d688b1a6 100644 --- a/extensions/react-widget/package-lock.json +++ b/extensions/react-widget/package-lock.json @@ -7,7 +7,7 @@ "": { "name": "docsgpt", "version": "0.7.1", - "license": "Apache-2.0", + "license": "MIT", "dependencies": { "@babel/plugin-transform-flow-strip-types": "^8.0.1", "@parcel/resolver-glob": "^2.16.4", diff --git a/extensions/react-widget/package.json b/extensions/react-widget/package.json index ddacae21..61d7f4be 100644 --- a/extensions/react-widget/package.json +++ b/extensions/react-widget/package.json @@ -99,7 +99,7 @@ "widget" ], "author": "Arc53", - "license": "Apache-2.0", + "license": "MIT", "bugs": { "url": "https://github.com/arc53/DocsGPT/issues" }, From a83e1dc0afb95c8b868b2fe83075c968fc5bdde7 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:41 +0100 Subject: [PATCH 083/130] feat(graphrag): seed the walk from what entities are, and rank with passages and vector hits Graph retrieval tied plain vector search at best and never beat it. Measured across five corpora, the bottleneck was seeding, not the graph: the walk started from nodes whose embeddings were computed from bare entity names, and a whole question shares almost nothing with a name like "Quill". Extraction now embeds each node from "name (type): description" and each relationship as the fact it asserts ("Alder streams_to Quill: ..."), stored on a new nullable graph_edges.fact_embedding column that ensure_vector_schema adds in place. Entity names are canonicalised (case, punctuation, word breaks and a cautious plural) so "VECTOR_STORE" and "vector stores" land on one node. Extraction calls run concurrently (GRAPHRAG_EXTRACTION_WORKERS, default 8) while embedding and graph writes stay serial on the task thread, so ordering and idempotency are unchanged; that measured 8.4x faster with identical output. Retrieval gains per-source options, stored under retrieval.graph and read live at query time: - seed_strategy: start from matching entities (default) or matching relationships, which can reach an entity the question never names; - passage_nodes (on): walk the source's passages alongside entities, with PageRank damping 0.5 instead of 0.85; - blend_vector (on): fuse the graph ranking with the source's vector ranking by reciprocal rank. The defaults are the measured-best configuration. Through GraphRAGRetriever, the new seeding moved recall@4 from 0.41 to 0.68 on a multi-hop corpus and from 0.50 to 1.00 on the docs corpus, and regressed none of the corpora measured. Existing graphs keep name-only embeddings until rebuilt. --- docs/content/Deploying/Settings-Reference.mdx | 6 + docsgpt/core/settings/retrieval.py | 9 + docsgpt/graphrag/extraction.py | 169 +++++++-- docsgpt/graphrag/naming.py | 94 +++++ docsgpt/graphrag/store.py | 331 +++++++++++++++++- docsgpt/retriever/graph_rag.py | 286 +++++++++++++-- docsgpt/storage/db/source_config.py | 24 +- tests/graphrag/test_extraction.py | 213 +++++++++++ tests/graphrag/test_retriever_default_path.py | 109 ++++++ tests/graphrag/test_retriever_passages.py | 132 +++++++ tests/graphrag/test_retriever_seeding.py | 129 +++++++ tests/graphrag/test_store.py | 131 ++++++- tests/retriever/test_graph_rag.py | 23 ++ 13 files changed, 1577 insertions(+), 79 deletions(-) create mode 100644 docsgpt/graphrag/naming.py create mode 100644 tests/graphrag/test_retriever_default_path.py create mode 100644 tests/graphrag/test_retriever_passages.py create mode 100644 tests/graphrag/test_retriever_seeding.py diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 32181adb..3221867f 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -447,6 +447,12 @@ Type `int`, default `2000`, must be `>= 0`. Hard cap on chunks extracted per source (cost control); 0 extracts nothing. +### `GRAPHRAG_EXTRACTION_WORKERS` + +Type `int`, default `8`, must be `>= 1` and `<= 32`. + +Concurrent extraction calls during ingest. Model calls run in parallel while graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial. + ## Vector stores diff --git a/docsgpt/core/settings/retrieval.py b/docsgpt/core/settings/retrieval.py index 651d3f3d..53dbf56b 100644 --- a/docsgpt/core/settings/retrieval.py +++ b/docsgpt/core/settings/retrieval.py @@ -31,6 +31,15 @@ class RetrievalSettings(SettingsGroup): GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = Field( default=2000, ge=0, description="Hard cap on chunks extracted per source (cost control); 0 extracts nothing." ) + GRAPHRAG_EXTRACTION_WORKERS: int = Field( + default=8, + ge=1, + le=32, + description=( + "Concurrent extraction calls during ingest. Model calls run in parallel while " + "graph writes stay serial, so ordering and idempotency are unchanged; 1 is fully serial." + ), + ) @field_validator("VECTOR_STORE", mode="before") @classmethod diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 1bdfb054..db8ea27a 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -27,6 +27,7 @@ from docsgpt.core.model_utils import ( get_api_key_for_provider, get_provider_from_model_id, ) +from docsgpt.graphrag.naming import normalize_entity_name from docsgpt.core.settings import settings from docsgpt.llm.llm_creator import LLMCreator from docsgpt.storage.db.source_config import SourceConfig @@ -224,6 +225,8 @@ def extract_graph_for_source( source's graph holds after the run — not how many upserts ran, which counts the same entity once per chunk it appears in. """ + from concurrent.futures import ThreadPoolExecutor + from docsgpt.graphrag.store import GraphStore store = GraphStore() @@ -266,43 +269,85 @@ def extract_graph_for_source( except Exception as exc: logger.debug("graph progress callback failed: %s", exc) - for chunk, chunk_id in to_process: + def _prepare(item): + """One chunk's LLM extraction — the only step run concurrently. + + A chunk spends almost all of its time waiting on the model, so that is + what runs in the pool. Everything else stays on the calling thread: + graph writes, so transactions and the progress checkpoint are exactly + what they were serially, and embedding. Inside a Celery worker the + embeddings client decides to embed locally from the task on the + *current thread's* stack; a pool thread has none, so it would instead + dispatch an embed task to the worker and wait on it, which Celery + refuses inside a task — failing every chunk of the build. + """ + chunk, chunk_id = item text = _chunk_text(chunk) if not text: - store.mark_chunk(source_id, chunk_id, "done") - chunks_processed += 1 - _report() - continue + return chunk_id, "empty", None extracted = _extract_chunk(llm, text, chunk_id) if extracted is None: - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() - continue - + return chunk_id, "failed", None try: entities = _build_entities(extracted["entities"]) relationships = _build_relationships(extracted["relationships"]) - name_embeddings = _embed_names(embedding, entities, relationships) - chunk_nodes, chunk_edges = store.apply_chunk( - source_id, chunk_id, entities, relationships, name_embeddings - ) - node_upserts += chunk_nodes - edges += chunk_edges - # ``apply_chunk`` marks the chunk done inside the transaction that - # writes its rows, so the checkpoint cannot disagree with the graph - # and a replayed write cannot apply the chunk twice. - chunks_processed += 1 except Exception as exc: logger.warning( - "Graph extraction write failed for chunk %s, skipping: %s", - chunk_id, - exc, + "Graph extraction failed for chunk %s, skipping: %s", chunk_id, exc ) - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() + return chunk_id, "failed", None + return chunk_id, "ok", (entities, relationships) + + workers = max(1, int(getattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 1) or 1)) + pool = None + if workers > 1 and len(to_process) > 1: + pool = ThreadPoolExecutor(max_workers=workers) + # ``map`` yields in submission order, so chunks are still applied in the + # order they were given and a run stays reproducible. + prepared = pool.map(_prepare, to_process) + else: + prepared = (_prepare(item) for item in to_process) + + try: + for chunk_id, status, payload in prepared: + if status == "empty": + store.mark_chunk(source_id, chunk_id, "done") + chunks_processed += 1 + _report() + continue + if status == "failed": + store.mark_chunk(source_id, chunk_id, "failed") + failed_chunks += 1 + _report() + continue + + entities, relationships = payload + try: + # On this thread, not in the pool — see ``_prepare``. + name_embeddings = _embed_names(embedding, entities, relationships) + _embed_facts(embedding, relationships) + chunk_nodes, chunk_edges = store.apply_chunk( + source_id, chunk_id, entities, relationships, name_embeddings + ) + node_upserts += chunk_nodes + edges += chunk_edges + # ``apply_chunk`` marks the chunk done inside the transaction that + # writes its rows, so the checkpoint cannot disagree with the graph + # and a replayed write cannot apply the chunk twice. + chunks_processed += 1 + except Exception as exc: + logger.warning( + "Graph extraction embed/write failed for chunk %s, skipping: %s", + chunk_id, + exc, + ) + store.mark_chunk(source_id, chunk_id, "failed") + failed_chunks += 1 + _report() + finally: + if pool is not None: + pool.shutdown(wait=True) try: store.set_node_degrees(source_id) @@ -347,7 +392,7 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: entities.append( { "name": name, - "normalized_name": name.lower(), + "normalized_name": normalize_entity_name(name), "type": str(e.get("type") or "") or None, "description": str(e.get("description") or "") or None, } @@ -373,6 +418,70 @@ def _build_relationships(raw_relationships: Any) -> List[Dict[str, Any]]: return relationships +def _fact_text(rel: Dict[str, Any]) -> str: + """A relationship rendered as the sentence it asserts. + + Embedded and stored on the edge so retrieval can match a question against + the *relation* rather than against entity names — the difference between + "which entity is this about" and "which fact answers this". + """ + source = str(rel.get("source") or "").strip() + target = str(rel.get("target") or "").strip() + if not source or not target: + return "" + relation = str(rel.get("type") or "related to").strip() or "related to" + text = f"{source} {relation} {target}" + description = str(rel.get("description") or "").strip() + return f"{text}: {description}" if description else text + + +def _embed_facts(embedding, relationships: List[Dict[str, Any]]) -> None: + """Attach a fact embedding to each relationship, in one batched call. + + Mutates the relationship dicts so the embedding travels with the edge into + ``apply_chunk`` without a second mapping to keep in step. Always on: it is + one extra batched call per chunk against an LLM call that already costs + far more, and it lets a source switch to relationship seeding at query time + without being rebuilt. + """ + pending = [(rel, _fact_text(rel)) for rel in relationships] + pending = [(rel, text) for rel, text in pending if text] + if not pending: + return + try: + vectors = embedding.embed_documents([text for _rel, text in pending]) + except Exception as exc: # noqa: BLE001 + # The graph is still correct without them; only fact seeding degrades. + logger.warning("Fact embedding failed, continuing without: %s", exc) + return + for (rel, _text), vector in zip(pending, vectors): + rel["fact_embedding"] = vector + + +def _seed_text(entity: Dict[str, Any]) -> str: + """The text a node's embedding is computed from. + + Retrieval seeds the graph walk by matching a whole question against these + embeddings, and a bare entity name is a poor thing to match a question + against — a question about what a service writes to shares almost no + surface with the name ``Quill``. Including the type and description gives + the match something to work with; measured across five corpora it moved + recall@4 by +0.07 to +0.50. + + Relationship endpoints keep their bare names: they arrive as strings with + no type or description attached. + """ + name = str(entity.get("name") or "").strip() + text = name + entity_type = str(entity.get("type") or "").strip() + if entity_type: + text += f" ({entity_type})" + description = str(entity.get("description") or "").strip() + if description: + text += f": {description}" + return text or name + + def _embed_names( embedding, entities: List[Dict[str, Any]], @@ -385,14 +494,16 @@ def _embed_names( """ name_by_norm: Dict[str, str] = {} for entity in entities: - name_by_norm.setdefault(entity["normalized_name"], entity["name"]) + name_by_norm.setdefault(entity["normalized_name"], _seed_text(entity)) for rel in relationships: for endpoint in (rel.get("source"), rel.get("target")): if endpoint is None: continue clean = str(endpoint).strip() if clean: - name_by_norm.setdefault(clean.lower(), clean) + # Same key the store resolves endpoints by, or the embedding + # computed here never reaches the node it was computed for. + name_by_norm.setdefault(normalize_entity_name(clean), clean) if not name_by_norm: return {} diff --git a/docsgpt/graphrag/naming.py b/docsgpt/graphrag/naming.py new file mode 100644 index 00000000..3dbfa1cb --- /dev/null +++ b/docsgpt/graphrag/naming.py @@ -0,0 +1,94 @@ +"""Canonical entity naming for the per-source knowledge graph. + +Nodes are merged on ``normalized_name``, which has been ``name.lower()``. That +splits entities a reader would call the same thing: measured on the DocsGPT docs +corpus, ``agent``/``agents``, ``VECTOR_STORE``/``Vector store``/``vector stores``, +``Celery worker``/``Celery workers`` and ``.env file``/``env_file`` all landed as +separate nodes — 58 such collisions across 1,704 entities, with 75% of entities +appearing in exactly one chunk as a result. + +:func:`canonical_name` folds the differences that are purely orthographic: +case, surrounding punctuation, underscore/hyphen word breaks, and a *cautious* +plural. Cautious matters: this corpus contains ``postgres``, ``kubernetes``, +``https`` and ``aws``, none of which are plurals, so a naive "strip trailing s" +would corrupt them into new entities rather than merge anything. + +Always on: every graph is built with canonical names. +""" + +from __future__ import annotations + +import re + +_PUNCT = re.compile(r"[^\w\s]+", re.UNICODE) +_UNDERSCORE = re.compile(r"[_\-]+") +_SPACE = re.compile(r"\s+") + +#: Words that end in "s" without being plural. Singularising these would invent +#: entities ("postgre", "kubernete") instead of merging existing ones. +_NOT_PLURAL = frozenset( + { + "postgres", "kubernetes", "https", "aws", "dns", "tls", "cors", "css", + "js", "sas", "gas", "ss", "class", "access", "process", "status", + "analysis", "basis", "axis", "https", "rss", "less", "express", + "redis", "nats", "kibana", "elasticsearch", "os", "ios", "macos", + "always", "sometimes", "series", "docs", "ops", "devops", "sse", + } +) + + +def _singular(word: str) -> str: + """Best-effort singular of one word, biased hard towards leaving it alone. + + Only the endings that are unambiguous in this domain are touched: + ``-ies`` -> ``-y`` (``policies``), ``-ses``/``-xes``/``-zes``/``-ches``/ + ``-shes`` -> drop ``es`` (``indexes``, ``batches``), and a bare trailing + ``s`` on a word long enough to be safe. Everything in :data:`_NOT_PLURAL`, + and anything ending in ``ss``/``us``/``is``, is returned unchanged. + """ + if len(word) < 4 or word in _NOT_PLURAL: + return word + if word.endswith(("ss", "us", "is")): + return word + if word.endswith("ies") and len(word) > 4: + return word[:-3] + "y" + if word.endswith(("ses", "xes", "zes", "ches", "shes")): + return word[:-2] + if word.endswith("s"): + return word[:-1] + return word + + +def canonical_name(name: str) -> str: + """Merge key for an entity name. + + Args: + name: The entity name as the model wrote it. + + Returns: + A lowercase, punctuation-free, singularised key. Returns ``""`` for an + empty or punctuation-only name, which callers treat as "no entity". + + Examples: + ``VECTOR_STORE`` and ``Vector stores`` -> ``vector store``; + ``.env file`` and ``env_file`` -> ``env file``; + ``postgres`` stays ``postgres``. + """ + if not name: + return "" + text = _UNDERSCORE.sub(" ", str(name)) + text = _PUNCT.sub(" ", text) + text = _SPACE.sub(" ", text).strip().lower() + if not text: + return "" + return " ".join(_singular(word) for word in text.split()) + + +def normalize_entity_name(name: str) -> str: + """The key an entity is merged on: its :func:`canonical_name`. + + Every graph the corpora were measured on was built this way, so it is the + only mode rather than a flag. A graph built before this used plain + ``lower()`` keys; re-extracting it merges onto these instead. + """ + return canonical_name(name) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index acb35911..134998ba 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -74,6 +74,16 @@ def _pgvector_identifiers() -> tuple[str, str, str, str]: ) +def _pgvector_vector_column() -> str: + """Resolve the embedding column name from the same ``PGVectorStore`` defaults.""" + import inspect + + from docsgpt.vectorstore.pgvector import PGVectorStore + + params = inspect.signature(PGVectorStore.__init__).parameters + return _safe_identifier(params["vector_column"].default) + + def _is_connection_lost(exc: BaseException) -> bool: """True when ``exc`` says the server connection went away, not that the SQL was bad. @@ -253,7 +263,7 @@ class GraphStore: ) cursor.execute( - """ + f""" CREATE TABLE IF NOT EXISTS graph_edges ( id UUID PRIMARY KEY, source_id UUID NOT NULL, @@ -262,10 +272,18 @@ class GraphStore: type TEXT, description TEXT, weight REAL DEFAULT 1.0, - source_chunk_ids JSONB + source_chunk_ids JSONB, + fact_embedding vector({dimension}) ); """ ) + # ``CREATE TABLE IF NOT EXISTS`` is a no-op on a database that + # already has the table, so a column added after the fact needs its + # own statement or every existing deployment silently lacks it. + cursor.execute( + f"ALTER TABLE graph_edges " + f"ADD COLUMN IF NOT EXISTS fact_embedding vector({dimension});" + ) cursor.execute( """ @@ -453,19 +471,78 @@ class GraphStore: description: Optional[str] = None, weight: float = 1.0, source_chunk_ids: Optional[List[str]] = None, - ) -> str: - """Insert an edge on an open cursor (no commit, no degree bump). + fact_embedding: Optional[List[float]] = None, + ) -> tuple[Optional[str], bool]: + """Write an edge on an open cursor (no commit, no degree bump). + + Returns ``(edge_id, created)``. Two shapes of noise are rejected here + rather than at read time, because once written neither is visible: + + * A self-loop feeds a node's PageRank mass straight back to itself. It + is dropped, reported as ``(None, False)``. + * A pair already related by the same type is *merged* rather than + inserted again. ``graph_edges`` carries no uniqueness constraint, so + re-extracting one relationship across many chunks otherwise writes a + row per chunk — a fifth of a real corpus's edges — inflating that + pair's traversal weight and spending the bounded subgraph fetch on + duplicates. The surviving row keeps the strongest weight seen and + every contributing chunk id. Callers that batch many edges run ``set_node_degrees`` once afterwards instead of bumping degree per edge. """ + if str(src_node_id) == str(dst_node_id): + return None, False + + cursor.execute( + """ + SELECT id + FROM graph_edges + WHERE source_id = %s AND src_node_id = %s AND dst_node_id = %s + AND type IS NOT DISTINCT FROM %s + LIMIT 1; + """, + (source_id, src_node_id, dst_node_id, type), + ) + existing = cursor.fetchone() + if existing: + edge_id = existing[0] + # The chunk ids are merged in SQL, against the row's own current + # value, rather than read here and written back: a read-modify-write + # would drop whatever a concurrent writer appended in between. + cursor.execute( + """ + UPDATE graph_edges + SET weight = GREATEST(COALESCE(weight, 0), %s), + description = COALESCE(description, %s), + -- Backfills the fact embedding for an edge first written + -- before fact embeddings were switched on. + fact_embedding = COALESCE(fact_embedding, %s::vector), + source_chunk_ids = COALESCE(source_chunk_ids, '[]'::jsonb) || ( + SELECT COALESCE(jsonb_agg(candidate), '[]'::jsonb) + FROM jsonb_array_elements(%s::jsonb) AS candidate + WHERE NOT COALESCE(source_chunk_ids, '[]'::jsonb) + @> jsonb_build_array(candidate) + ) + WHERE id = %s; + """, + ( + weight, + description, + fact_embedding, + Jsonb(list(source_chunk_ids or [])), + edge_id, + ), + ) + return str(edge_id), False + edge_id = str(uuid.uuid4()) cursor.execute( """ INSERT INTO graph_edges (id, source_id, src_node_id, dst_node_id, type, description, - weight, source_chunk_ids) - VALUES (%s, %s, %s, %s, %s, %s, %s, %s); + weight, source_chunk_ids, fact_embedding) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s); """, ( edge_id, @@ -476,9 +553,10 @@ class GraphStore: description, weight, Jsonb(source_chunk_ids or []), + fact_embedding, ), ) - return edge_id + return edge_id, True def add_edge( self, @@ -489,21 +567,28 @@ class GraphStore: description: Optional[str] = None, weight: float = 1.0, source_chunk_ids: Optional[List[str]] = None, - ) -> str: - """Insert an edge and bump the degree of both endpoints. Returns its id.""" + fact_embedding: Optional[List[float]] = None, + ) -> Optional[str]: + """Write an edge and bump the degree of both endpoints. Returns its id. + + Returns ``None`` for a self-loop, which is not written. A repeat of an + existing pair merges into that row and returns its id, leaving degree + alone — the endpoints gained no new neighbour. + """ self._ensure_tables_once() conn = self._get_connection() cursor = conn.cursor() try: - edge_id = self._add_edge( + edge_id, created = self._add_edge( cursor, source_id, src_node_id, dst_node_id, type, description, - weight, source_chunk_ids, - ) - cursor.execute( - "UPDATE graph_nodes SET degree = degree + 1 " - "WHERE source_id = %s AND id IN (%s, %s);", - (source_id, src_node_id, dst_node_id), + weight, source_chunk_ids, fact_embedding, ) + if created: + cursor.execute( + "UPDATE graph_nodes SET degree = degree + 1 " + "WHERE source_id = %s AND id IN (%s, %s);", + (source_id, src_node_id, dst_node_id), + ) conn.commit() return edge_id except Exception as e: @@ -608,7 +693,7 @@ class GraphStore: ) if src_id is None or dst_id is None: continue - self._add_edge( + _, created = self._add_edge( cursor, source_id, src_id, @@ -617,8 +702,10 @@ class GraphStore: description=rel.get("description"), weight=float(rel.get("weight") or 1.0), source_chunk_ids=[chunk_id], + fact_embedding=rel.get("fact_embedding"), ) - edges_added += 1 + if created: + edges_added += 1 cursor.execute( """ @@ -653,7 +740,11 @@ class GraphStore: clean = str(name).strip() if not clean: return None - normalized_name = clean.lower() + from docsgpt.graphrag.naming import normalize_entity_name + + normalized_name = normalize_entity_name(clean) + if not normalized_name: + return None if normalized_name in node_ids: return node_ids[normalized_name] node_id = self._upsert_node( @@ -839,6 +930,7 @@ class GraphStore: FROM graph_edges WHERE source_id = %s AND (src_node_id = ANY(%s) OR dst_node_id = ANY(%s)) + ORDER BY weight DESC NULLS LAST LIMIT %s; """, ( @@ -886,6 +978,7 @@ class GraphStore: FROM graph_edges WHERE source_id = %s AND src_node_id = ANY(%s) AND dst_node_id = ANY(%s) + ORDER BY weight DESC NULLS LAST LIMIT %s; """, (source_id, node_id_list, node_id_list, MAX_SUBGRAPH_EDGES), @@ -999,6 +1092,206 @@ class GraphStore: cursor.close() conn.rollback() + def seed_nodes_from_facts( + self, + source_id: str, + query_embedding: List[float], + fact_limit: int = 5, + limit: int = 10, + ) -> List[Dict[str, Any]]: + """Seed nodes drawn from the *relationships* nearest the question. + + Name matching asks "which entity is this question about", which a + multi-document question cannot answer: the entity holding the answer is + named in another document, not in the question. A fact string carries + the relation — "Alder streams_to Quill: ..." — so a question about what + a service writes to can match the edge itself and seed the walk on both + of its endpoints, including the one nothing in the question names. + + Endpoints are weighted by fact score divided by the entity's + ``doc_freq``: an entity appearing in every chunk is a poor seed even + when it sits on a well-matched fact, and dividing by how widely it + occurs prefers the specific endpoint over the hub. + + Rows match :meth:`search_nodes_by_embedding`'s shape, so the caller's + seed weighting is unchanged. Returns nothing when the source has no + fact embeddings, which is the signal to fall back to name matching. + """ + if not query_embedding: + return [] + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + """ + WITH top_facts AS ( + SELECT src_node_id, dst_node_id, + 1 - (fact_embedding <=> %s::vector) AS score + FROM graph_edges + WHERE source_id = %s AND fact_embedding IS NOT NULL + ORDER BY fact_embedding <=> %s::vector + LIMIT %s + ) + SELECT n.id::text, n.name, n.description, + MAX(f.score / GREATEST(COALESCE(n.doc_freq, 1), 1)) AS weight + FROM top_facts f + JOIN graph_nodes n + ON n.id = f.src_node_id OR n.id = f.dst_node_id + WHERE n.source_id = %s + GROUP BY n.id, n.name, n.description + ORDER BY weight DESC + LIMIT %s; + """, + ( + query_embedding, + source_id, + query_embedding, + max(1, int(fact_limit)), + source_id, + max(1, int(limit)), + ), + ) + return [ + { + "id": row[0], + "name": row[1], + "description": row[2], + # The caller reads weight back as ``1 - distance``. + "distance": 1.0 - float(row[3] or 0.0), + } + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error seeding nodes from facts: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def entity_relationships( + self, source_id: str, name: str, limit: int = 25 + ) -> List[Dict[str, Any]]: + """The relationships an entity takes part in, strongest first. + + This is the one thing a caller cannot get from vector search: which + *named* thing an entity is connected to. Matching is on the name rather + than a node id because the caller is an LLM holding a name it read in + the text, not an id. + """ + clean = (name or "").strip() + if not clean: + return [] + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + """ + SELECT s.name, e.type, d.name, e.description + FROM graph_edges e + JOIN graph_nodes s ON s.id = e.src_node_id + JOIN graph_nodes d ON d.id = e.dst_node_id + WHERE e.source_id = %s AND (s.name ILIKE %s OR d.name ILIKE %s) + ORDER BY e.weight DESC NULLS LAST + LIMIT %s; + """, + (source_id, f"%{clean}%", f"%{clean}%", max(1, int(limit))), + ) + return [ + {"source": row[0], "type": row[1], "target": row[2], "description": row[3]} + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error reading relationships for {name!r}: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def entity_pages( + self, source_id: str, name: str, limit: int = 4 + ) -> List[Dict[str, Any]]: + """Chunks an entity appears in, with the chunk it is *about* first. + + A plain substring match answers "Halvard" with pages that merely mention + Halvard, and an unordered ``LIMIT`` then decides which of those the + caller sees. Nodes whose name is the entity (or the entity plus a + qualifier the extractor appended, "Quill" -> "Quill Store") are + preferred, and among those the chunk whose text opens with the name + comes first; a substring match is the fallback so an unusual name still + resolves. + """ + clean = (name or "").strip() + if not clean: + return [] + table, text_col, metadata_col, source_col = _pgvector_identifiers() + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + f""" + SELECT d.{metadata_col}, d.{text_col}, + (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject + FROM graph_node_chunks gc + JOIN graph_nodes n ON n.id = gc.node_id + JOIN {table} d ON d.id::text = gc.chunk_id + WHERE gc.source_id = %s AND d.{source_col} = %s + AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) + GROUP BY d.{metadata_col}, d.{text_col}, is_subject + ORDER BY is_subject DESC, (d.{text_col} ILIKE %s) DESC + LIMIT %s; + """, + ( + clean.lower(), f"{clean.lower()} %", + source_id, source_id, + clean.lower(), f"{clean.lower()} %", f"%{clean}%", + f"{clean}%", + max(1, int(limit)), + ), + ) + return [ + {"metadata": row[0] or {}, "text": row[1] or ""} + for row in cursor.fetchall() + ] + except Exception as e: + logging.error(f"Error reading pages for {name!r}: {e}") + return [] + finally: + cursor.close() + conn.rollback() + + def chunk_similarities( + self, source_id: str, chunk_ids: List[str], query_embedding: List[float] + ) -> Dict[str, float]: + """Cosine similarity between the query and specific chunks of a source. + + Passage nodes need their own relevance to claim a share of the walk's + restart mass, and that number lives in the co-located pgvector table — + the same one :meth:`get_chunk_texts` reads. Restricted to the chunk ids + the subgraph actually reached, so this never scans the whole source. + """ + if not chunk_ids or not query_embedding: + return {} + table, _text_col, _metadata_col, source_col = _pgvector_identifiers() + vector_col = _pgvector_vector_column() + conn = self._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + f""" + SELECT id::text, 1 - ({vector_col} <=> %s::vector) + FROM {table} + WHERE {source_col} = %s AND id::text = ANY(%s); + """, + (query_embedding, source_id, [str(c) for c in chunk_ids]), + ) + return {row[0]: float(row[1]) for row in cursor.fetchall()} + except Exception as e: + logging.error(f"Error scoring chunks against the query: {e}") + return {} + finally: + cursor.close() + conn.rollback() + def get_chunk_texts( self, source_id: str, diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index 81ed9bf0..818f4f07 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -32,6 +32,7 @@ from docsgpt.graphrag.store import GraphStore from docsgpt.retriever.base import BaseRetriever from docsgpt.retriever.classic_rag import ClassicRAG from docsgpt.retriever.labels import labels_from_metadata +from docsgpt.storage.db.source_config import GraphRetrievalConfig from docsgpt.utils import num_tokens_from_string from docsgpt.vectorstore.base import get_embeddings @@ -39,11 +40,27 @@ SEED_NODES = 10 SUBGRAPH_HOPS = 1 +PASSAGE_NODE_WEIGHT = 0.05 +FACT_SEED_FACTS = 5 +RRF_K = 60 + +# PageRank damping per ranking mode — each is the value that mode was measured +# at. Lower keeps mass nearer the seeds; with passages in the walk 0.5 measured +# better, while entity-only ranking was measured at the conventional 0.85. +DAMPING_WITH_PASSAGES = 0.5 +DAMPING_ENTITIES_ONLY = 0.85 + + def _idf(doc_freq: Any) -> float: """Node-specificity weight: rarer entities (low ``doc_freq``) score higher.""" return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0) +def _damping(passage_nodes: bool) -> float: + """PageRank damping for a ranking mode: the value that mode was measured at.""" + return DAMPING_WITH_PASSAGES if passage_nodes else DAMPING_ENTITIES_ONLY + + def _restart_vector(nodes: List[Any], personalization: Dict[Any, float] | None) -> Dict[Any, float]: """Normalized restart distribution over ``nodes``. @@ -210,18 +227,9 @@ class GraphRAGRetriever(BaseRetriever): After PPR, each node's mass is scaled by ``1/log(2 + doc_freq)`` so a high-degree hub contributes less than a specific entity at equal mass. """ - graph = nx.Graph() - for node in subgraph.get("nodes", []): - graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0)) - for edge in subgraph.get("edges", []): - src, dst = edge["src_node_id"], edge["dst_node_id"] - if src in graph and dst in graph: - raw_weight = edge.get("weight") - # Same rule the ranker applies: default only a missing or null - # weight. Coercing an explicit 0 to 1.0 here would make "these - # entities are not related" the strongest possible link. - edge_weight = 1.0 if raw_weight is None else float(raw_weight) - graph.add_edge(src, dst, weight=edge_weight) + # Through the class, not ``self``: this method reads no instance state, + # and callers (and tests) rely on being able to invoke it unbound. + graph = GraphRAGRetriever._subgraph_graph(subgraph) if graph.number_of_nodes() == 0: return {} @@ -230,13 +238,33 @@ class GraphRAGRetriever(BaseRetriever): personalization = None ranks = _personalized_pagerank( - graph, personalization=personalization, weight="weight" + graph, + personalization=personalization, + weight="weight", + alpha=_damping(passage_nodes=False), ) return { node: rank * _idf(graph.nodes[node].get("doc_freq", 0)) for node, rank in ranks.items() } + @staticmethod + def _subgraph_graph(subgraph) -> "nx.Graph": + """The fetched subgraph as a weighted undirected graph.""" + graph = nx.Graph() + for node in subgraph.get("nodes", []): + graph.add_node(node["id"], doc_freq=node.get("doc_freq", 0)) + for edge in subgraph.get("edges", []): + src, dst = edge["src_node_id"], edge["dst_node_id"] + if src in graph and dst in graph: + raw_weight = edge.get("weight") + # Default only a missing or null weight. Coercing an explicit 0 + # to 1.0 would make "these entities are not related" the + # strongest possible link. + edge_weight = 1.0 if raw_weight is None else float(raw_weight) + graph.add_edge(src, dst, weight=edge_weight) + return graph + def _rank_chunks(self, store, source_id, node_scores) -> List[str]: """Score chunks by summed (PPR mass x IDF) of their linked nodes; top candidates. @@ -254,6 +282,74 @@ class GraphRAGRetriever(BaseRetriever): candidates = max(self.chunks * 2, self.chunks + 5) return ranked[: max(1, candidates)] + def _rank_chunks_with_passages( + self, store, source_id, subgraph, seeds, query_embedding + ) -> List[str]: + """Rank chunks by walking a graph that contains the chunks themselves. + + :meth:`_rank_chunks` reads a chunk's score *off* its entities, summing + their PPR mass — so a chunk touching many mid-scoring generic entities + outranks one touching the few entities the question is about. Putting + the chunks in the walk instead, each joined to its own entities and + carrying a small share of the restart mass proportional to its own + vector similarity, makes a chunk reachable both ways: by being about the + question, and by being connected to what is. Graph retrieval then + contains vector retrieval rather than competing with it. + + Only an improvement when the seeds are good: measured across five + corpora it helped alongside richer seed embeddings and *hurt* with + bare-name seeds (0.73 -> 0.57 on one corpus). Graphs are now always + built with the richer seed text; one built before that change should be + rebuilt before this is relied on. + """ + node_ids = [node["id"] for node in subgraph.get("nodes", [])] + chunk_links = store.get_chunk_ids_for_nodes(source_id, node_ids) + candidate_ids = sorted({c for chunks in chunk_links.values() for c in chunks}) + if not candidate_ids: + return [] + + graph = self._subgraph_graph(subgraph) + similarities = store.chunk_similarities( + source_id, candidate_ids, query_embedding + ) + # Normalised so the passage share is a fixed fraction of the restart + # mass rather than whatever absolute cosine this embedding model emits. + scores = [similarities.get(c, 0.0) for c in candidate_ids] + low, high = (min(scores), max(scores)) if scores else (0.0, 0.0) + spread = high - low + + personalization = dict(seeds) + passage_of: Dict[str, str] = {} + for chunk_id in candidate_ids: + linked = [n for n, chunks in chunk_links.items() if chunk_id in chunks] + linked = [n for n in linked if n in graph] + if not linked: + continue + passage_node = f"chunk::{chunk_id}" + passage_of[passage_node] = chunk_id + for node in linked: + graph.add_edge(passage_node, node, weight=1.0) + similarity = similarities.get(chunk_id, 0.0) + normalized = (similarity - low) / spread if spread > 0 else 0.0 + personalization[passage_node] = normalized * PASSAGE_NODE_WEIGHT + + if graph.number_of_nodes() == 0 or not any(personalization.values()): + return [] + + ranks = _personalized_pagerank( + graph, + personalization=personalization, + weight="weight", + alpha=_damping(passage_nodes=True), + ) + chunk_scores = { + chunk_id: ranks.get(passage_node, 0.0) + for passage_node, chunk_id in passage_of.items() + } + ranked = sorted(chunk_scores, key=lambda c: chunk_scores[c], reverse=True) + candidates = max(self.chunks * 2, self.chunks + 5) + return ranked[: max(1, candidates)] + def _source_top_k(self, source_id) -> int: """How many chunks this source may contribute — its own top-k. @@ -273,6 +369,111 @@ class GraphRAGRetriever(BaseRetriever): base = self.base_chunks if self.base_chunks is not None else self.chunks return max(1, base // max(1, len(self.vectorstores))) + def _vector_ranking(self, source_id, query_embedding: List[float]) -> List[tuple]: + """The source's own vector ranking, as ``(text, metadata)`` in score order. + + Used only by the hybrid path. Vector hits carry no row id, so the fused + ranking is keyed on the chunk text itself — the one identifier both + rankings share — and the metadata travels with it so a hit the graph + never surfaced can still be emitted as a document. + """ + from docsgpt.vectorstore.vector_creator import VectorCreator + + store = None + try: + store = VectorCreator.create_vectorstore( + settings.VECTOR_STORE, source_id, settings.EMBEDDINGS_KEY + ) + hits = store.search( + self._classic._get_rephrased_question(), + k=max(self.chunks * 4, 20), + query_vector=query_embedding, + ) + except Exception as e: + logging.error( + "GraphRAG hybrid: vector ranking failed for %s: %s", source_id, e + ) + return [] + finally: + close = getattr(store, "close", None) + if close is not None: + try: + close() + except Exception as e: + logging.debug("Error closing hybrid vector store: %s", e) + ranked = [] + for hit in hits: + text = getattr(hit, "page_content", None) + metadata = getattr(hit, "metadata", None) + if text is None and isinstance(hit, dict): + text = hit.get("text") or hit.get("page_content") + metadata = hit.get("metadata") + if text: + ranked.append((text, metadata or {})) + return ranked + + @staticmethod + def _rrf_order(rankings: List[List[str]], k: int) -> Dict[str, float]: + """Reciprocal rank fusion over ranked lists of the same key type. + + Rank-based on purpose: PPR mass and cosine similarity are not on + comparable scales, and normalising either one invents a calibration + that does not exist. + """ + scores: Dict[str, float] = {} + for ranking in rankings: + for position, key in enumerate(ranking): + scores[key] = scores.get(key, 0.0) + 1.0 / (k + position + 1) + return scores + + def _graph_options(self, source_id) -> GraphRetrievalConfig: + """This source's graph retrieval options, or the recommended defaults. + + Options travel on the per-source retrieval config the Dispatcher hands + over. A request that carries no per-source detail gets the defaults, + which are the measured-best configuration rather than a neutral one. + """ + cfg = (getattr(self, "per_source_retrieval", None) or {}).get(source_id) + options = cfg.get("graph") if isinstance(cfg, dict) else getattr(cfg, "graph", None) + if isinstance(options, GraphRetrievalConfig): + return options + try: + return GraphRetrievalConfig.model_validate(options or {}) + except Exception: + return GraphRetrievalConfig() + + def _seed_rows( + self, store, source_id, query_embedding: List[float] + ) -> List[Dict[str, Any]]: + """The nodes the walk restarts from, per the source's ``seed_strategy``. + + Seeding decides more than ranking does: a walk that starts on the wrong + nodes cannot be rescued downstream. + + ``entities`` + Cosine NN over entity embeddings, built from each entity's name, + type and description so a whole question has something to match. + The default: best or tied-best on every corpus measured. + ``relationships`` + Cosine NN over relationship sentences ("A streams_to B: ..."), + seeding both endpoints of the best-matching facts. The only way to + start on an entity the question never names; strongest on + chain-structured content, weaker on ordinary prose. + + Relationship seeding falls back to entity matching for a source with no + fact embeddings (one built before they were recorded), so it still + retrieves rather than returning nothing. + """ + if self._graph_options(source_id).seed_strategy == "relationships": + by_fact = store.seed_nodes_from_facts( + source_id, query_embedding, fact_limit=FACT_SEED_FACTS, limit=SEED_NODES + ) + if by_fact: + return by_fact + return store.search_nodes_by_embedding( + source_id, query_embedding, k=SEED_NODES + ) + def _graph_docs_for_source( self, store, source_id, query_embedding: List[float] ) -> List[Dict[str, Any]]: @@ -284,9 +485,7 @@ class GraphRAGRetriever(BaseRetriever): query_embedding: Embedding of the rephrased question, computed once by the caller for the whole retrieval. """ - seed_rows = store.search_nodes_by_embedding( - source_id, query_embedding, k=SEED_NODES - ) + seed_rows = self._seed_rows(store, source_id, query_embedding) if not seed_rows: return [] @@ -300,26 +499,63 @@ class GraphRAGRetriever(BaseRetriever): for row in seed_rows } + options = self._graph_options(source_id) subgraph = store.get_subgraph(source_id, seed_ids, hops=SUBGRAPH_HOPS) - node_scores = self._ppr_scores(subgraph, seeds) - if not node_scores: + if options.passage_nodes: + chunk_ids = self._rank_chunks_with_passages( + store, source_id, subgraph, seeds, query_embedding + ) + else: + node_scores = self._ppr_scores(subgraph, seeds) + if not node_scores: + return [] + chunk_ids = self._rank_chunks(store, source_id, node_scores) + if not chunk_ids: return [] - chunk_ids = self._rank_chunks(store, source_id, node_scores) chunk_data = store.get_chunk_texts(source_id, chunk_ids) + # ``(text, metadata)`` in rank order. Chunk ids stop being the currency + # here: a hit contributed by the vector ranking has no graph chunk id, + # and keying on ids is what made an earlier version of this fusion able + # only to reorder the graph's own candidates. + candidates: List[tuple] = [] + for chunk_id in chunk_ids: + chunk = chunk_data.get(chunk_id) + text = chunk.get("text") if chunk else None + if text: + candidates.append((text, chunk.get("metadata"))) + + if options.blend_vector: + # The graph ranks by how much PPR mass landed on a chunk's + # entities, which says nothing about whether the chunk is about the + # question. Fusing with the source's own vector ranking keeps the + # graph's reach while letting plain relevance back in — including + # chunks the graph never surfaced, which is where most of the value + # is: no reordering can rescue a question whose answer the graph + # missed entirely. + vector_hits = self._vector_ranking(source_id, query_embedding) + if vector_hits: + metadata_by_text = {text: meta for text, meta in candidates} + for text, meta in vector_hits: + metadata_by_text.setdefault(text, meta) + fused = self._rrf_order( + [[t for t, _ in candidates], [t for t, _ in vector_hits]], + RRF_K, + ) + candidates = [ + (text, metadata_by_text.get(text)) + for text in sorted(fused, key=lambda t: fused[t], reverse=True) + ] + docs: List[Dict[str, Any]] = [] token_budget = max(int(self.doc_token_limit * 0.9), 100) cumulative_tokens = 0 source_top_k = self._source_top_k(source_id) - for chunk_id in chunk_ids: + for text, metadata in candidates: if len(docs) >= source_top_k: break - chunk = chunk_data.get(chunk_id) - text = chunk.get("text") if chunk else None - if not text: - continue - labels = labels_from_metadata(chunk.get("metadata"), text, source_id) + labels = labels_from_metadata(metadata, text, source_id) doc_tokens = num_tokens_from_string(f"{labels['filename']}\n{text}") if cumulative_tokens + doc_tokens >= token_budget: break diff --git a/docsgpt/storage/db/source_config.py b/docsgpt/storage/db/source_config.py index 20873b46..8f8d6209 100644 --- a/docsgpt/storage/db/source_config.py +++ b/docsgpt/storage/db/source_config.py @@ -12,7 +12,7 @@ reproduces today's chunking byte-for-byte. from __future__ import annotations -from typing import Optional +from typing import Literal, Optional from pydantic import BaseModel, ConfigDict, field_validator, model_validator @@ -90,6 +90,27 @@ class ChunkingConfig(BaseModel): duplicate_headers: bool = False +class GraphRetrievalConfig(BaseModel): + """How the graph retriever walks a graphrag source (live; no re-ingest). + + The defaults are the configuration that measured best across the corpora + tested rather than a neutral starting point: seed from entity matches, put + the passages in the walk, and blend with the source's own vector ranking. + """ + + model_config = ConfigDict(extra="forbid") + + # Where the walk starts: entities whose descriptions match the question, or + # relationships ("A streams_to B") that do. Relationships can start the + # walk on an entity the question never names. + seed_strategy: Literal["entities", "relationships"] = "entities" + # Chunks join the walk as nodes, so a passage is reachable both by being + # about the question and by being connected to what is. + passage_nodes: bool = True + # Fuse the graph ranking with plain vector search by reciprocal rank. + blend_vector: bool = True + + class RetrievalConfig(BaseModel): """Query-time retrieval knobs (live; no re-ingest needed).""" @@ -102,6 +123,7 @@ class RetrievalConfig(BaseModel): rephrase_query: bool = True # toggle ClassicRAG._rephrase_query side-call reranker: Optional[dict] = None # reserved: future cross-encoder/LLM reorder prescreen: Optional[dict] = None # None = off; else PreScreenConfig dict (D12) + graph: GraphRetrievalConfig = GraphRetrievalConfig() # graphrag retriever only @field_validator("chunks") @classmethod diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index 48e567d4..eda8dd7c 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -138,6 +138,129 @@ def _extraction_json(entities, relationships): return json.dumps({"entities": entities, "relationships": relationships}) +class TestFactText: + """A relationship rendered as the sentence it asserts. + + This is what fact seeding matches a question against, so it has to read as + a claim rather than as three fields concatenated. + """ + + def test_renders_the_relationship_as_a_sentence(self): + text = extraction_module._fact_text( + { + "source": "Alder", + "target": "Quill", + "type": "streams_to", + "description": "Alder streams audit events to Quill.", + } + ) + + assert text == "Alder streams_to Quill: Alder streams audit events to Quill." + + def test_omits_an_absent_description(self): + text = extraction_module._fact_text( + {"source": "Alder", "target": "Quill", "type": "streams_to"} + ) + + assert text == "Alder streams_to Quill" + + def test_defaults_a_missing_relation(self): + text = extraction_module._fact_text({"source": "Alder", "target": "Quill"}) + + assert text == "Alder related to Quill" + + @pytest.mark.parametrize( + "rel", + [ + {"source": "Alder", "target": ""}, + {"source": "", "target": "Quill"}, + {}, + ], + ) + def test_an_edge_without_both_endpoints_has_no_fact(self, rel): + assert extraction_module._fact_text(rel) == "" + + +class TestEmbedFacts: + """Fact embeddings are always recorded, so a source can switch to + relationship seeding at query time without being rebuilt.""" + + def _relationships(self): + return [{"source": "Alder", "target": "Quill", "type": "streams_to"}] + + def test_attaches_one_embedding_per_fact_in_a_single_call(self): + relationships = self._relationships() + [{"source": "", "target": "Nowhere"}] + calls = [] + + class _Embedding: + def embed_documents(self, texts): + calls.append(texts) + return [[0.5] * 4 for _ in texts] + + extraction_module._embed_facts(_Embedding(), relationships) + + # One batched call, and the endpoint-less relationship is skipped + # rather than embedded as an empty string. + assert calls == [["Alder streams_to Quill"]] + assert relationships[0]["fact_embedding"] == [0.5] * 4 + assert "fact_embedding" not in relationships[1] + + def test_survives_an_embedding_failure(self): + """The graph is still correct without fact embeddings — only + relationship seeding degrades, and it falls back to entities — so a + failure here must not fail the chunk.""" + relationships = self._relationships() + + class _Embedding: + def embed_documents(self, texts): + raise RuntimeError("embeddings down") + + extraction_module._embed_facts(_Embedding(), relationships) + + assert "fact_embedding" not in relationships[0] + + +class TestSeedText: + """What a node's embedding is computed from. + + Retrieval matches a whole question against these embeddings, so what goes + into them decides what the graph walk can start from. + """ + + def _entity(self): + return { + "name": "Quill", + "normalized_name": "quill", + "type": "store", + "description": "A write-ahead store.", + } + + def test_includes_type_and_description(self): + assert ( + extraction_module._seed_text(self._entity()) + == "Quill (store): A write-ahead store." + ) + + def test_falls_back_to_the_name_when_fields_are_missing(self): + assert extraction_module._seed_text({"name": "Quill"}) == "Quill" + + def test_embedded_text_is_keyed_by_the_normalized_name(self): + """The richer text must reach ``embed_documents``, keyed by the same + normalized name the store resolves nodes by — otherwise the embedding + is computed for a node it never reaches.""" + captured = {} + + class _Embedding: + def embed_documents(self, texts): + captured["texts"] = texts + return [[0.0] * 4 for _ in texts] + + result = extraction_module._embed_names(_Embedding(), [self._entity()], []) + + assert captured["texts"] == ["Quill (store): A write-ahead store."] + assert set(result) == {"quill"} + + @pytest.mark.integration class TestExtractionLive: @pytest.fixture @@ -193,6 +316,96 @@ class TestExtractionLive: finally: store.delete_by_source(source_id) + def test_parallel_workers_process_every_chunk_once( + self, store, source_id, monkeypatch, stub_embedding + ): + """Running the model calls concurrently must not change what gets written. + + Extraction spends nearly all of a chunk's time waiting on the model, so + the calls run in a pool while every graph write stays on the calling + thread. Six chunks share one entity here: whatever order the pool + finishes in, that entity is upserted once, each chunk is linked, and all + six are marked processed. + """ + from docsgpt.core.settings import settings + + try: + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + llm = _StubLLM([payload] * 6) + _install_stub_llm(monkeypatch, llm) + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + + summary = extract_graph_for_source( + source_id, + user="owner-1", + chunks=[ + _chunk(f"c{i}", f"Ada appears here, take {i}.") for i in range(6) + ], + config=SourceConfig(), + request_id="req-parallel", + ) + + assert summary["chunks_processed"] == 6 + assert summary["failed_chunks"] == 0 + assert summary["nodes"] == 1 + assert len(llm.gen_calls) == 6 + + node = store.get_node_by_normalized(source_id, "ada") + assert node is not None + mapping = store.get_chunk_ids_for_nodes(source_id, [node["id"]]) + assert sorted(mapping[node["id"]]) == [f"c{i}" for i in range(6)] + finally: + store.delete_by_source(source_id) + + def test_embedding_runs_on_the_calling_thread( + self, store, source_id, monkeypatch, stub_embedding + ): + """Only the LLM call may run in the extraction pool, never embedding. + + Inside a Celery worker the embeddings client decides to embed locally + from the task on the *current thread's* stack. A pool thread has none, + so from there it dispatches an embed task to the worker and waits on + it — which Celery refuses inside a task, so every chunk of a graph + build failed. + """ + import threading + + from docsgpt.core.settings import settings + + caller = threading.current_thread() + seen = [] + real_embed_names = extraction_module._embed_names + + def _recording_embed_names(*args, **kwargs): + seen.append(threading.current_thread()) + return real_embed_names(*args, **kwargs) + + monkeypatch.setattr(extraction_module, "_embed_names", _recording_embed_names) + try: + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + _install_stub_llm(monkeypatch, _StubLLM([payload] * 4)) + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + + summary = extract_graph_for_source( + source_id, + user="owner-1", + chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(4)], + config=SourceConfig(), + request_id="req-thread", + ) + + assert summary["failed_chunks"] == 0 + assert len(seen) == 4 + assert all(thread is caller for thread in seen) + finally: + store.delete_by_source(source_id) + def test_same_entity_across_chunks_merges( self, store, source_id, monkeypatch, stub_embedding ): diff --git a/tests/graphrag/test_retriever_default_path.py b/tests/graphrag/test_retriever_default_path.py new file mode 100644 index 00000000..c11a0f0b --- /dev/null +++ b/tests/graphrag/test_retriever_default_path.py @@ -0,0 +1,109 @@ +"""The graph retriever's default path, end to end through ``_graph_docs_for_source``. + +The shipped defaults — seed from entities, walk the passages, blend with vector +search — are the configuration that measured best, so they are what most graph +sources run. This drives that whole path with a store that returns real values, +and checks each per-source option actually switches its stage off. +""" + +from __future__ import annotations + +from docsgpt.retriever.graph_rag import GraphRAGRetriever +from docsgpt.storage.db.source_config import RetrievalConfig + +TEXTS = { + "c-alder": "Alder streams audit events to Quill.", + "c-quill": "Quill is compacted every six hours.", +} +VECTOR_ONLY = "A passage only plain vector search found." + + +class _Store: + """A two-entity chain: the question matches Alder, the answer is on Quill.""" + + def __init__(self): + self.calls: list[str] = [] + + def search_nodes_by_embedding(self, source_id, query_embedding, k=10): + return [{"id": "alder", "name": "Alder", "distance": 0.1}] + + def get_subgraph(self, source_id, node_ids, hops=1): + return { + "nodes": [{"id": "alder", "doc_freq": 1}, {"id": "quill", "doc_freq": 1}], + "edges": [{"src_node_id": "alder", "dst_node_id": "quill", "weight": 1.0}], + } + + def get_chunk_ids_for_nodes(self, source_id, node_ids): + return {"alder": ["c-alder"], "quill": ["c-quill"]} + + def chunk_similarities(self, source_id, chunk_ids, query_embedding): + self.calls.append("chunk_similarities") + return {"c-alder": 0.9, "c-quill": 0.2} + + def get_chunk_texts(self, source_id, chunk_ids): + return { + c: {"text": TEXTS[c], "metadata": {"title": c}} + for c in chunk_ids + if c in TEXTS + } + + +def _retriever(per_source=None): + """A retriever without its constructor (which builds a ClassicRAG).""" + retriever = object.__new__(GraphRAGRetriever) + retriever.chunks = 3 + retriever.base_chunks = None + retriever.doc_token_limit = 50000 + retriever.vectorstores = ["src"] + retriever.per_source_retrieval = per_source or {} + retriever.vector_calls = 0 + + def _vector_ranking(source_id, query_embedding): + retriever.vector_calls += 1 + return [(VECTOR_ONLY, {"title": "vector"})] + + retriever._vector_ranking = _vector_ranking + return retriever + + +def _texts(docs): + return [doc["text"] for doc in docs] + + +class TestDefaultPath: + def test_walks_passages_and_blends_in_vector_hits(self): + store = _Store() + retriever = _retriever() + + docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2]) + + # The answer sits one edge away from the seed: the walk reached it. + assert TEXTS["c-quill"] in _texts(docs) + # A hit only vector search found is blended in, not lost. + assert VECTOR_ONLY in _texts(docs) + assert store.calls == ["chunk_similarities"] + assert retriever.vector_calls == 1 + + +class TestPerSourceOptions: + def test_passage_walk_can_be_switched_off(self): + store = _Store() + retriever = _retriever( + {"src": RetrievalConfig(chunks=3, graph={"passage_nodes": False})} + ) + + docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2]) + + assert "chunk_similarities" not in store.calls + assert TEXTS["c-quill"] in _texts(docs) + + def test_vector_blending_can_be_switched_off(self): + store = _Store() + retriever = _retriever( + {"src": RetrievalConfig(chunks=3, graph={"blend_vector": False})} + ) + + docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2]) + + assert retriever.vector_calls == 0 + assert VECTOR_ONLY not in _texts(docs) diff --git a/tests/graphrag/test_retriever_passages.py b/tests/graphrag/test_retriever_passages.py new file mode 100644 index 00000000..fcc561ff --- /dev/null +++ b/tests/graphrag/test_retriever_passages.py @@ -0,0 +1,132 @@ +"""Chunks as nodes in the walk, and the damping that decides how far mass spreads. + +``_rank_chunks`` reads a chunk's score off its entities by summing their PPR +mass, which rewards a chunk for touching *many* entities rather than the right +ones. The passage-node path puts the chunks in the graph instead, so a chunk is +reachable both by being about the question and by being connected to what is. + +These tests use a stub store: the ranking is graph arithmetic, and pinning it +against a real database would measure Postgres rather than the ranking. +""" + +from __future__ import annotations + +import pytest + +from docsgpt.retriever.graph_rag import GraphRAGRetriever, _damping + + +class _StubStore: + """The two reads the passage path makes, and nothing else.""" + + def __init__(self, chunk_links, similarities): + self._chunk_links = chunk_links + self._similarities = similarities + + def get_chunk_ids_for_nodes(self, source_id, node_ids): + return {n: c for n, c in self._chunk_links.items() if n in set(node_ids)} + + def chunk_similarities(self, source_id, chunk_ids, query_embedding): + return {c: self._similarities.get(c, 0.0) for c in chunk_ids} + + +def _retriever(chunks=2): + """A retriever without its constructor — which builds a ClassicRAG, opens + settings-driven collaborators, and has nothing to do with ranking.""" + retriever = object.__new__(GraphRAGRetriever) + retriever.chunks = chunks + return retriever + + +def _subgraph(): + return { + "nodes": [ + {"id": "a", "doc_freq": 1}, + {"id": "b", "doc_freq": 1}, + {"id": "hub", "doc_freq": 40}, + ], + "edges": [ + {"src_node_id": "a", "dst_node_id": "hub", "weight": 1.0}, + {"src_node_id": "b", "dst_node_id": "hub", "weight": 1.0}, + ], + } + + +class TestDamping: + """Each ranking mode runs at the damping it was measured at.""" + + def test_passage_walk_keeps_mass_near_the_seeds(self): + assert _damping(passage_nodes=True) == 0.5 + + def test_entity_only_ranking_keeps_the_conventional_value(self): + assert _damping(passage_nodes=False) == 0.85 + + +class TestPassageNodes: + def test_ranks_the_chunk_the_question_matches(self, monkeypatch): + """Two chunks are equally connected; only their own relevance differs, + so the more relevant one must win.""" + store = _StubStore( + chunk_links={"a": ["c1"], "b": ["c2"]}, + similarities={"c1": 0.1, "c2": 0.9}, + ) + + ranked = _retriever()._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0, "b": 1.0}, [0.0] * 4 + ) + + assert ranked[0] == "c2" + + def test_a_chunk_reached_only_through_the_graph_still_ranks(self, monkeypatch): + """The point of the walk: a chunk with no similarity of its own is + still reachable through the entity the seeds point at.""" + store = _StubStore( + chunk_links={"a": ["c1"], "b": ["c2"]}, + similarities={"c1": 0.0, "c2": 0.0}, + ) + + ranked = _retriever()._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4 + ) + + assert set(ranked) == {"c1", "c2"} + + def test_no_linked_chunks_returns_nothing(self): + store = _StubStore(chunk_links={}, similarities={}) + + assert ( + _retriever()._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4 + ) + == [] + ) + + def test_over_fetches_past_the_chunk_budget(self, monkeypatch): + """Same contract as ``_rank_chunks``: candidates exceed the budget so + chunks with missing text cannot drop the final count below it.""" + links = {"a": [f"c{i}" for i in range(10)]} + store = _StubStore( + chunk_links=links, + similarities={f"c{i}": i / 10 for i in range(10)}, + ) + + ranked = _retriever(chunks=2)._rank_chunks_with_passages( + store, "src", _subgraph(), {"a": 1.0}, [0.0] * 4 + ) + + assert len(ranked) == max(2 * 2, 2 + 5) + + +class TestChunkSimilaritiesGuard: + """The store call the passage path depends on short-circuits before it + touches a connection, so an empty subgraph costs no query.""" + + @pytest.mark.parametrize( + "chunk_ids,embedding", [([], [0.1]), (["c1"], []), ([], [])] + ) + def test_empty_inputs_return_empty(self, chunk_ids, embedding): + from docsgpt.graphrag.store import GraphStore + + store = object.__new__(GraphStore) + + assert store.chunk_similarities("src", chunk_ids, embedding) == {} diff --git a/tests/graphrag/test_retriever_seeding.py b/tests/graphrag/test_retriever_seeding.py new file mode 100644 index 00000000..0918aa4c --- /dev/null +++ b/tests/graphrag/test_retriever_seeding.py @@ -0,0 +1,129 @@ +"""Where the graph walk starts, per the source's graph retrieval options. + +Seeding decides more than ranking does — a walk that starts on the wrong nodes +cannot be rescued downstream. The options are per source and live (no +re-ingest), carried on the per-source retrieval config the Dispatcher hands the +retriever, so both the dispatch and how the options are resolved are pinned. + +The fallback matters most: relationship seeding reads fact embeddings written +at ingest, and a source built before they were recorded has none. It must keep +retrieving through entity matching rather than returning nothing. +""" + +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from docsgpt.retriever.graph_rag import GraphRAGRetriever +from docsgpt.storage.db.source_config import GraphRetrievalConfig, RetrievalConfig + +ENTITY_ROWS = [{"id": "n1", "name": "Quill", "distance": 0.2}] +FACT_ROWS = [ + {"id": "f1", "name": "Alder", "distance": 0.1}, + {"id": "f2", "name": "Quill", "distance": 0.1}, +] + + +class _StubStore: + def __init__(self, fact_rows=None): + self.fact_rows = list(FACT_ROWS) if fact_rows is None else fact_rows + self.calls: list[str] = [] + + def seed_nodes_from_facts(self, source_id, query_embedding, fact_limit=5, limit=10): + self.calls.append("facts") + return self.fact_rows + + def search_nodes_by_embedding(self, source_id, query_embedding, k=10): + self.calls.append("entities") + return list(ENTITY_ROWS) + + +def _retriever(per_source=None): + """A retriever without its constructor, which builds a ClassicRAG.""" + retriever = object.__new__(GraphRAGRetriever) + if per_source is not None: + retriever.per_source_retrieval = per_source + return retriever + + +def _relationships_config(): + return RetrievalConfig(graph={"seed_strategy": "relationships"}) + + +class TestDefaults: + def test_measured_best_configuration_is_the_default(self): + options = GraphRetrievalConfig() + + assert options.seed_strategy == "entities" + assert options.passage_nodes is True + assert options.blend_vector is True + + def test_a_source_with_no_per_source_config_gets_the_defaults(self): + assert _retriever()._graph_options("src") == GraphRetrievalConfig() + + def test_existing_retrieval_configs_validate_without_the_new_block(self): + """Source configs saved before this existed carry no ``graph`` key.""" + config = RetrievalConfig.model_validate({"retriever": "graphrag"}) + + assert config.graph == GraphRetrievalConfig() + + @pytest.mark.parametrize( + "bad", [{"seed_strategy": "vector"}, {"seed_strategy": "union"}, {"damping": 0.5}] + ) + def test_retired_and_unknown_options_are_rejected(self, bad): + with pytest.raises(ValidationError): + GraphRetrievalConfig.model_validate(bad) + + +class TestEntitySeeding: + def test_seeds_from_entities_and_never_reads_facts(self): + store = _StubStore() + + rows = _retriever()._seed_rows(store, "src", [0.0, 0.1]) + + assert [r["id"] for r in rows] == ["n1"] + assert store.calls == ["entities"] + + +class TestRelationshipSeeding: + def test_seeds_from_facts_when_the_source_asks_for_it(self): + store = _StubStore() + retriever = _retriever({"src": _relationships_config()}) + + rows = retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert [r["id"] for r in rows] == ["f1", "f2"] + assert store.calls == ["facts"] + + def test_reads_the_option_from_a_plain_dict_config_too(self): + store = _StubStore() + retriever = _retriever({"src": {"graph": {"seed_strategy": "relationships"}}}) + + retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert store.calls == ["facts"] + + def test_falls_back_to_entities_without_fact_embeddings(self): + """A source built before fact embeddings were recorded still retrieves.""" + store = _StubStore(fact_rows=[]) + retriever = _retriever({"src": _relationships_config()}) + + rows = retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert [r["id"] for r in rows] == ["n1"] + assert store.calls == ["facts", "entities"] + + def test_options_are_per_source(self): + store = _StubStore() + retriever = _retriever({"other": _relationships_config()}) + + retriever._seed_rows(store, "src", [0.0, 0.1]) + + assert store.calls == ["entities"] + + def test_a_malformed_stored_option_falls_back_to_the_defaults(self): + """A bad value must not take graph retrieval down with it.""" + retriever = _retriever({"src": {"graph": {"seed_strategy": "nonsense"}}}) + + assert retriever._graph_options("src") == GraphRetrievalConfig() diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index baa19604..02a315db 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -157,6 +157,122 @@ class TestGraphStoreLive: finally: store.delete_by_source(source_id) + def test_seed_nodes_from_facts_returns_both_endpoints_of_the_match( + self, store, source_id + ): + """Fact seeding's whole point: the question matches the *relationship*, + and both of its endpoints become seeds — including the one the question + never names.""" + try: + alder = store.upsert_node(source_id, "Alder", "alder", "service", "d") + quill = store.upsert_node(source_id, "Quill", "quill", "store", "d") + birch = store.upsert_node(source_id, "Birch", "birch", "service", "d") + ridge = store.upsert_node(source_id, "Ridge", "ridge", "store", "d") + store.add_edge( + source_id, alder, quill, "streams_to", "Alder streams to Quill", + 1.0, ["c1"], fact_embedding=_embedding(1.0), + ) + store.add_edge( + source_id, birch, ridge, "streams_to", "Birch streams to Ridge", + 1.0, ["c2"], fact_embedding=_embedding(-1.0), + ) + + rows = store.seed_nodes_from_facts( + source_id, _embedding(1.0), fact_limit=1, limit=10 + ) + + assert {row["name"] for row in rows} == {"Alder", "Quill"} + assert all(row["distance"] <= 1.0 for row in rows) + finally: + store.delete_by_source(source_id) + + def test_seed_nodes_from_facts_is_empty_without_fact_embeddings( + self, store, source_id + ): + """A source ingested before fact embeddings existed returns nothing, + which is the signal the retriever falls back to name matching on.""" + try: + a = store.upsert_node(source_id, "A", "a") + b = store.upsert_node(source_id, "B", "b") + store.add_edge(source_id, a, b, "rel") + + assert store.seed_nodes_from_facts(source_id, _embedding(1.0)) == [] + finally: + store.delete_by_source(source_id) + + def test_add_edge_skips_self_loops(self, store, source_id): + """A relationship whose endpoints resolve to one node is noise. + + A self-loop feeds a node's PageRank mass straight back to itself, and a + real extraction produced 121 of them on a 98-page corpus. + """ + try: + a = store.upsert_node(source_id, "A", "a", "thing", "desc a") + assert ( + store.add_edge(source_id, a, a, "related", "a relates to a", 1.0, ["c1"]) + is None + ) + assert store.get_subgraph(source_id, [a], hops=1)["edges"] == [] + finally: + store.delete_by_source(source_id) + + def test_add_edge_merges_a_repeated_pair(self, store, source_id): + """The same relationship seen in many chunks is one edge, not many rows. + + ``graph_edges`` carries no uniqueness constraint, so re-extracting a + relationship used to insert a row per chunk — 19.9% of a real corpus's + edges — inflating traversal weight and wasting the subgraph fetch + budget. The surviving row keeps the strongest weight and both chunk ids. + """ + try: + a = store.upsert_node(source_id, "A", "a", "thing", "desc a") + b = store.upsert_node(source_id, "B", "b", "thing", "desc b") + first = store.add_edge(source_id, a, b, "related", "d", 2.0, ["chunk-1"]) + second = store.add_edge(source_id, a, b, "related", "d", 5.0, ["chunk-2"]) + + assert second == first + edges = store.get_subgraph(source_id, [a, b], hops=1)["edges"] + assert len(edges) == 1 + assert float(edges[0]["weight"]) == 5.0 + + # Both chunks are still recorded as evidence for the merged edge. + conn = store._get_connection() + cursor = conn.cursor() + try: + cursor.execute( + "SELECT source_chunk_ids FROM graph_edges WHERE id = %s;", (first,) + ) + chunk_ids = cursor.fetchone()[0] + finally: + cursor.close() + conn.rollback() + assert sorted(chunk_ids) == ["chunk-1", "chunk-2"] + finally: + store.delete_by_source(source_id) + + def test_get_subgraph_keeps_the_heaviest_edges_when_capped( + self, store, source_id, monkeypatch + ): + """A capped fetch must drop the weakest edges, not an arbitrary subset. + + The cap is applied with ``LIMIT``; without an ordering Postgres is free + to return any rows at all, so a dense graph silently retrieves a random + neighbourhood. + """ + try: + a = store.upsert_node(source_id, "A", "a", "thing", "d") + b = store.upsert_node(source_id, "B", "b", "thing", "d") + c = store.upsert_node(source_id, "C", "c", "thing", "d") + store.add_edge(source_id, a, b, "light", "d", 1.0, ["c1"]) + store.add_edge(source_id, a, c, "heavy", "d", 9.0, ["c1"]) + + monkeypatch.setattr(store_module, "MAX_SUBGRAPH_EDGES", 1) + edges = store.get_subgraph(source_id, [a], hops=1)["edges"] + + assert [e["type"] for e in edges] == ["heavy"] + finally: + store.delete_by_source(source_id) + def test_apply_chunk_writes_nodes_links_and_edges(self, store, source_id): """One transactional write: entities linked to the chunk, edges added, and a bare relationship endpoint upserted but not chunk-linked.""" @@ -203,18 +319,23 @@ class TestGraphStoreLive: store.delete_by_source(source_id) def test_self_loop_degree_agrees_across_paths(self, store, source_id): - """``add_edge``'s incremental +1 and ``set_node_degrees`` recompute must - agree on a self-loop (count it once).""" + """``add_edge``'s incremental bump and ``set_node_degrees`` recompute must + agree on a self-loop. + + They now agree on zero rather than one: the self-loop is rejected at + write time, so neither path has an edge to count. The property under + test is that the two paths agree, not the number they agree on. + """ try: node = store.upsert_node(source_id, "Solo", "solo") - store.add_edge(source_id, node, node, "self") + assert store.add_edge(source_id, node, node, "self") is None incremental = store.get_node_by_normalized(source_id, "solo")["degree"] - assert incremental == 1 + assert incremental == 0 store.set_node_degrees(source_id) recomputed = store.get_node_by_normalized(source_id, "solo")["degree"] - assert recomputed == 1 + assert recomputed == incremental == 0 finally: store.delete_by_source(source_id) diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index 8943174c..e4d1bf46 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -45,6 +45,29 @@ def _patch_embed(monkeypatch): ) +@pytest.fixture(autouse=True) +def _entity_only_ranking(monkeypatch): + """Pin the ranking path these tests were written for. + + Everything here exercises entity-only PPR ranking without vector blending, + driven through ``MagicMock`` stores. The shipped default now walks the + passages and blends with vector search — covered end to end in + ``tests/graphrag/test_retriever_default_path.py`` with a store that returns + real values. Pinning keeps each test here asserting what it was written to + assert, rather than whatever a mock happens to return on a path it never set + up. + """ + from docsgpt.storage.db.source_config import GraphRetrievalConfig + + monkeypatch.setattr( + GraphRAGRetriever, + "_graph_options", + lambda self, source_id: GraphRetrievalConfig( + passage_nodes=False, blend_vector=False + ), + ) + + # ── Fallback to ClassicRAG ──────────────────────────────────────────────────── From 6b9b1933371b1e9bcc8e32ac3b3b037d918ec901 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:41 +0100 Subject: [PATCH 084/130] feat(agents): let agents walk a source's graph when they can search it A one-shot graph ranking diffuses over the whole neighbourhood; a question whose answer sits two hops away is better served by following the edges. The graph_search tool gives an agent search_entities, get_relationships and read_entity_pages over its graph sources. On a multi-hop corpus where the bridging entity is never named, answers went from 0/10 with classic vector retrieval to 10/10 with the tool, end to end through /stream. The tool has no setting of its own. It is offered exactly where the agent can already search: agentic and research agents always, a classic agent only for sources exposed as a search tool. A graph source left at prefetch is used for ranking only. Tests now pin GRAPHRAG_ENABLED to its shipped default, as CI has: with a dev .env enabling it, every agent test's graph check read the developer's real database and left a pool to it behind. --- docsgpt/agents/agentic_agent.py | 2 + docsgpt/agents/classic_agent.py | 2 + docsgpt/agents/research_agent.py | 2 + docsgpt/agents/tools/graph_search.py | 277 +++++++++++++++++++++++ tests/conftest.py | 15 ++ tests/graphrag/test_graph_search_tool.py | 172 ++++++++++++++ 6 files changed, 470 insertions(+) create mode 100644 docsgpt/agents/tools/graph_search.py create mode 100644 tests/graphrag/test_graph_search_tool.py diff --git a/docsgpt/agents/agentic_agent.py b/docsgpt/agents/agentic_agent.py index b83c485d..0b87485c 100644 --- a/docsgpt/agents/agentic_agent.py +++ b/docsgpt/agents/agentic_agent.py @@ -2,6 +2,7 @@ import logging from typing import Dict, Generator, Optional from docsgpt.agents.base import BaseAgent +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.logging import LogContext @@ -33,6 +34,7 @@ class AgenticAgent(BaseAgent): ) -> Generator[Dict, None, None]: tools_dict = self.tool_executor.get_tools() add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) self._prepare_tools(tools_dict) diff --git a/docsgpt/agents/classic_agent.py b/docsgpt/agents/classic_agent.py index 2bc25130..5956ed3c 100644 --- a/docsgpt/agents/classic_agent.py +++ b/docsgpt/agents/classic_agent.py @@ -2,6 +2,7 @@ import logging from typing import Dict, Generator, Optional from docsgpt.agents.base import BaseAgent +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.logging import LogContext @@ -38,6 +39,7 @@ class ClassicAgent(BaseAgent): tools_dict = self.tool_executor.get_tools() if self.retriever_config: add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) self._prepare_tools(tools_dict) diff --git a/docsgpt/agents/research_agent.py b/docsgpt/agents/research_agent.py index 1d19a7a9..ee92507a 100644 --- a/docsgpt/agents/research_agent.py +++ b/docsgpt/agents/research_agent.py @@ -6,6 +6,7 @@ from typing import Dict, Generator, List, Optional from docsgpt.agents.base import BaseAgent from docsgpt.agents.tool_executor import ToolExecutor +from docsgpt.agents.tools.graph_search import add_graph_search_tool from docsgpt.agents.tools.internal_search import ( INTERNAL_TOOL_ID, add_internal_search_tool, @@ -277,6 +278,7 @@ class ResearchAgent(BaseAgent): tools_dict = self.tool_executor.get_tools() add_internal_search_tool(tools_dict, self.retriever_config) + add_graph_search_tool(tools_dict, self.retriever_config) if self.wiki_config: add_wiki_tool(tools_dict, self.wiki_config) diff --git a/docsgpt/agents/tools/graph_search.py b/docsgpt/agents/tools/graph_search.py new file mode 100644 index 00000000..c2c15a7b --- /dev/null +++ b/docsgpt/agents/tools/graph_search.py @@ -0,0 +1,277 @@ +"""Let the model search the knowledge graph itself, one edge at a time. + +Graph retrieval normally runs as a ranker: seed a walk from the question, +diffuse mass over a subgraph, hand back the highest-scoring chunks. Measured +across five corpora that never beat plain vector search, because a question +whose answer lives two documents away has nothing in it for the seeding step to +match — the bridging entity is named in the *first* document, not the question. + +Exposing the graph as tools removes the guess. The model can look up the +service, read which store it names, then fetch that store's page: the chain +followed deliberately rather than approximated by a diffusion. On a corpus built +so that vector search cannot shortcut the chain, this took two-hop answers from +1/8 to 8/8, against 0.40 for vector and 0.47 for one-shot graph retrieval. + +It is not a general win, and is deliberately not a default. On ordinary prose +documentation it *lost* to plain vector search (0.50 against 0.90): it answers +well when a question names an entity and wanders when the question is a task +description. It also costs several model round-trips per answer instead of one. +So it is offered only where a source owner has already chosen search over +prefetch — the per-source exposure setting, or an agentic/research agent — and +suits content that is genuinely chain-structured: runbooks, service catalogues, +infrastructure inventories. +""" + +from __future__ import annotations + +import logging +from typing import Any, Dict, List, Optional + +from docsgpt.agents.tools.base import Tool +from docsgpt.core.settings import settings + +logger = logging.getLogger(__name__) + +GRAPH_TOOL_ID = "graph_search" +MAX_PAGE_CHARS = 1500 + + +class GraphSearchTool(Tool): + """Entity lookup, relationship traversal and page reads over a source's graph.""" + + internal = True + + def __init__(self, config: Dict): + self.config = config or {} + self._store = None + self.retrieved_docs: List[Dict] = [] + + # -- plumbing ------------------------------------------------------------ + def _sources(self) -> List[str]: + source = self.config.get("source") or {} + active = source.get("active_docs") or [] + if isinstance(active, str): + active = [active] + return [str(s) for s in active if s] + + def _get_store(self): + if self._store is None: + from docsgpt.graphrag.store import GraphStore + + self._store = GraphStore() + return self._store + + def _embed(self, text: str) -> Optional[List[float]]: + try: + from docsgpt.vectorstore.base import get_embeddings + + return get_embeddings().embed_query(text) + except Exception as e: # noqa: BLE001 + logger.error(f"Graph tool could not embed the query: {e}") + return None + + # -- actions ------------------------------------------------------------- + def execute_action(self, action_name: str, **kwargs): + if not settings.GRAPHRAG_ENABLED: + return "The knowledge graph is not enabled for this deployment." + if not self._sources(): + return "No graph-backed sources are configured." + try: + if action_name == "search_entities": + return self._search_entities(**kwargs) + if action_name == "get_relationships": + return self._get_relationships(**kwargs) + if action_name == "read_entity_pages": + return self._read_entity_pages(**kwargs) + except Exception as e: # noqa: BLE001 + logger.error(f"Graph tool action {action_name} failed: {e}", exc_info=True) + return "The graph lookup failed." + return f"Unknown action: {action_name}" + + def _search_entities(self, **kwargs) -> str: + query = str(kwargs.get("query") or "").strip() + if not query: + return "Error: 'query' parameter is required." + limit = max(1, min(int(kwargs.get("k") or 8), 25)) + + embedding = self._embed(query) + if embedding is None: + return "Entity search is unavailable." + + store = self._get_store() + lines: List[str] = [] + for source_id in self._sources(): + for row in store.search_nodes_by_embedding(source_id, embedding, k=limit): + similarity = 1.0 - float(row.get("distance") or 0.0) + description = (row.get("description") or "").strip() + suffix = f" — {description[:160]}" if description else "" + lines.append(f"- {row['name']} (match {similarity:.2f}){suffix}") + if not lines: + return f"No entities found for {query!r}." + return "Entities:\n" + "\n".join(lines[:limit]) + + def _get_relationships(self, **kwargs) -> str: + entity = str(kwargs.get("entity") or "").strip() + if not entity: + return "Error: 'entity' parameter is required." + + store = self._get_store() + lines: List[str] = [] + for source_id in self._sources(): + for edge in store.entity_relationships(source_id, entity): + relation = edge.get("type") or "related to" + lines.append(f"- {edge['source']} --{relation}--> {edge['target']}") + if not lines: + return ( + f"No relationships found for {entity!r}. Try search_entities first " + "to get the exact name used in the graph." + ) + return f"Relationships for {entity!r}:\n" + "\n".join(lines) + + def _read_entity_pages(self, **kwargs) -> str: + entity = str(kwargs.get("entity") or "").strip() + if not entity: + return "Error: 'entity' parameter is required." + + store = self._get_store() + parts: List[str] = [] + for source_id in self._sources(): + for page in store.entity_pages(source_id, entity): + metadata = page.get("metadata") or {} + title = ( + metadata.get("file_path") + or metadata.get("title") + or metadata.get("source") + or "document" + ) + text = (page.get("text") or "")[:MAX_PAGE_CHARS] + doc = {"title": title, "text": text, "source": metadata.get("source", "")} + if doc not in self.retrieved_docs: + self.retrieved_docs.append(doc) + parts.append(f"--- {title} ---\n{text}") + if not parts: + return f"No documents mention {entity!r}." + return "\n\n".join(parts) + + # -- metadata ------------------------------------------------------------ + def get_actions_metadata(self): + return [ + { + "name": "search_entities", + "description": ( + "Find named things in the knowledge graph — services, components, " + "settings, people — whose names resemble a query. Use this first to " + "learn the exact name the graph uses before asking for its " + "relationships." + ), + "parameters": { + "properties": { + "query": { + "type": "string", + "description": "What to look for, e.g. a service or component name.", + "filled_by_llm": True, + "required": True, + }, + "k": { + "type": "integer", + "description": "How many entities to return (default 8).", + "filled_by_llm": True, + "required": False, + }, + } + }, + }, + { + "name": "get_relationships", + "description": ( + "List what an entity is connected to, as 'source --relation--> target'. " + "This is how you answer a question about something the question does not " + "name: look up what it points at, then read that thing's pages." + ), + "parameters": { + "properties": { + "entity": { + "type": "string", + "description": "Exact entity name, as returned by search_entities.", + "filled_by_llm": True, + "required": True, + } + } + }, + }, + { + "name": "read_entity_pages", + "description": ( + "Read the documentation an entity appears in, the page it is about " + "first. Use this once you know which entity holds the answer." + ), + "parameters": { + "properties": { + "entity": { + "type": "string", + "description": "Exact entity name, as returned by search_entities.", + "filled_by_llm": True, + "required": True, + } + } + }, + }, + ] + + def get_config_requirements(self): + return {} + + +def build_graph_tool_entry() -> Dict: + """The synthetic ``tools_dict`` entry for the graph tool.""" + tool = GraphSearchTool({}) + actions = [] + for action in tool.get_actions_metadata(): + entry = dict(action) + entry["active"] = True + actions.append(entry) + return {"name": "graph_search", "actions": actions} + + +def sources_have_graph(source: Dict) -> bool: + """Whether any active source actually has a graph to search.""" + active = source.get("active_docs") or [] + if isinstance(active, str): + active = [active] + if not active: + return False + try: + from docsgpt.graphrag.store import GraphStore + + counts = GraphStore().count_nodes_many([str(a) for a in active]) + return any(count > 0 for count in counts.values()) + except Exception as e: # noqa: BLE001 + logger.debug(f"Could not check for graphs: {e}") + return False + + +def add_graph_search_tool(tools_dict: Dict, retriever_config: Dict) -> None: + """Add the graph tool when the agent's search-tool sources include a graph. + + No setting of its own: ``retriever_config`` already carries exactly the + sources the agent may *search* — the ones a source owner exposed as a + search tool, or every source for an agentic/research agent — so the graph + tool follows that same per-source exposure choice. A graph source left at + ``prefetch`` in a classic agent is used for ranking only. + """ + if not settings.GRAPHRAG_ENABLED: + return + source = retriever_config.get("source") or {} + if not source.get("active_docs") or not sources_have_graph(source): + return + + entry = build_graph_tool_entry() + # The executor resolves tools by ``id``; this one is synthetic (no DB row). + entry["id"] = GRAPH_TOOL_ID + entry["config"] = {"source": source} + tools_dict[GRAPH_TOOL_ID] = entry + + +def build_graph_tool_config(source: Dict, **_ignored: Any) -> Dict: + """Config for :class:`GraphSearchTool` — it only needs the source ids.""" + return {"source": source} diff --git a/tests/conftest.py b/tests/conftest.py index 2d2e338b..876e6d89 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -183,6 +183,21 @@ def _no_worker_delegation(monkeypatch): monkeypatch.setattr("docsgpt.cache._pubsub_redis_creation_failed", True) +@pytest.fixture(autouse=True) +def _graphrag_off_by_default(monkeypatch): + """Run with GraphRAG at its shipped default (off), as CI does. + + Every agent that gets a search tool checks its sources for a graph, and + that check reads the configured vector database. A dev ``.env`` enabling + GraphRAG sent unrelated agent tests to the developer's real database and + left a pool to it in ``pgconn._POOLS``, failing a live test that asserts it + owns the only pool. Tests that exercise GraphRAG turn it on themselves. + """ + from docsgpt.core.settings import settings + + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False, raising=False) + + @pytest.fixture def mock_llm(): llm = Mock() diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py new file mode 100644 index 00000000..2f63cf13 --- /dev/null +++ b/tests/graphrag/test_graph_search_tool.py @@ -0,0 +1,172 @@ +"""The graph exposed to an agent as callable tools. + +Ranking with a graph never beat vector search in measurement; letting a model +*follow* an edge did, on content where the answer is two documents away. These +tests cover the contract that makes that possible — the tool must return the +relationships verbatim enough for the model to read a name out of them, and +must refuse clearly rather than silently when it has nothing to offer. +""" + +from __future__ import annotations + +from docsgpt.agents.tools.graph_search import ( + GRAPH_TOOL_ID, + GraphSearchTool, + add_graph_search_tool, + build_graph_tool_entry, +) +from docsgpt.core.settings import settings + +SOURCE = {"active_docs": ["src-1"]} + + +class _StubStore: + def __init__(self, nodes=None, relationships=None, pages=None): + self._nodes = nodes or [] + self._relationships = relationships or [] + self._pages = pages or [] + + def search_nodes_by_embedding(self, source_id, embedding, k=10): + return self._nodes[:k] + + def entity_relationships(self, source_id, name, limit=25): + return self._relationships + + def entity_pages(self, source_id, name, limit=4): + return self._pages + + +def _tool(monkeypatch, store, enabled=True): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", enabled) + tool = GraphSearchTool({"source": SOURCE}) + tool._store = store + monkeypatch.setattr(tool, "_embed", lambda text: [0.0, 0.1]) + return tool + + +class TestGating: + def test_reports_when_graphs_are_disabled(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore(), enabled=False) + + assert "not enabled" in tool.execute_action("search_entities", query="x") + + def test_reports_when_no_sources_are_configured(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + tool = GraphSearchTool({"source": {"active_docs": []}}) + + assert "No graph-backed sources" in tool.execute_action( + "search_entities", query="x" + ) + + def test_unknown_action_is_named(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore()) + + assert "Unknown action" in tool.execute_action("wander") + + +class TestActions: + def test_search_entities_lists_names_with_match_strength(self, monkeypatch): + tool = _tool( + monkeypatch, + _StubStore(nodes=[{"name": "Quill", "distance": 0.2, "description": "A store."}]), + ) + + result = tool.execute_action("search_entities", query="quill") + + assert "Quill" in result + assert "0.80" in result + + def test_search_entities_requires_a_query(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore()) + + assert "required" in tool.execute_action("search_entities", query=" ") + + def test_relationships_are_rendered_as_triples(self, monkeypatch): + """The model reads the *target* out of this line to take its next step, + so the target name has to survive rendering intact.""" + tool = _tool( + monkeypatch, + _StubStore( + relationships=[ + {"source": "Alder", "type": "streams_to", "target": "Quill", "description": ""} + ] + ), + ) + + result = tool.execute_action("get_relationships", entity="Alder") + + assert "Alder --streams_to--> Quill" in result + + def test_missing_relationships_suggest_the_next_step(self, monkeypatch): + """A dead end should point at search_entities rather than stop the agent.""" + tool = _tool(monkeypatch, _StubStore(relationships=[])) + + result = tool.execute_action("get_relationships", entity="Nope") + + assert "search_entities" in result + + def test_pages_are_titled_truncated_and_recorded(self, monkeypatch): + tool = _tool( + monkeypatch, + _StubStore( + pages=[{"metadata": {"file_path": "quill-store.md"}, "text": "x" * 5000}] + ), + ) + + result = tool.execute_action("read_entity_pages", entity="Quill") + + assert "--- quill-store.md ---" in result + assert len(result) < 3000 + # Accumulated so the answer can cite what the walk actually read. + assert tool.retrieved_docs[0]["title"] == "quill-store.md" + + def test_pages_absent_is_stated_plainly(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore(pages=[])) + + assert "No documents" in tool.execute_action("read_entity_pages", entity="Quill") + + +class TestWiring: + def test_entry_exposes_every_action(self): + entry = build_graph_tool_entry() + + assert {a["name"] for a in entry["actions"]} == { + "search_entities", + "get_relationships", + "read_entity_pages", + } + assert all(action["active"] for action in entry["actions"]) + + def test_not_added_when_graphs_are_disabled(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False) + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True + ) + tools: dict = {} + + add_graph_search_tool(tools, {"source": SOURCE}) + + assert tools == {} + + def test_not_added_when_the_sources_have_no_graph(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: False + ) + tools: dict = {} + + add_graph_search_tool(tools, {"source": SOURCE}) + + assert tools == {} + + def test_added_with_its_sentinel_id_and_source_config(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True + ) + tools: dict = {} + + add_graph_search_tool(tools, {"source": SOURCE}) + + assert tools[GRAPH_TOOL_ID]["id"] == GRAPH_TOOL_ID + assert tools[GRAPH_TOOL_ID]["config"]["source"] == SOURCE From 5e67f8c9276fefbf449c63092dba42306825f65a Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:41 +0100 Subject: [PATCH 085/130] feat(settings): graph retrieval options in a graph source's retrieval settings Exposes the three per-source graph options in the source's retrieval settings, shown only when the retriever is graphrag: where the walk starts (entities or relationships), whether passages join the walk, and whether vector hits are blended in. Defaults match the backend's measured-best configuration and are filled in for sources saved before the options existed. A note points at the "search tool" exposure, which is what offers the graph to an agent. --- frontend/src/locale/de.json | 13 +++ frontend/src/locale/en.json | 13 +++ frontend/src/locale/es.json | 13 +++ frontend/src/locale/jp.json | 13 +++ frontend/src/locale/ru.json | 13 +++ frontend/src/locale/zh-TW.json | 13 +++ frontend/src/locale/zh.json | 13 +++ frontend/src/models/misc.ts | 12 ++ .../components/RetrievalOptions.test.tsx | 43 +++++++ .../settings/components/RetrievalOptions.tsx | 110 ++++++++++++++++++ 10 files changed, 256 insertions(+) diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json index e64ed52f..b914fbfd 100644 --- a/frontend/src/locale/de.json +++ b/frontend/src/locale/de.json @@ -268,6 +268,19 @@ }, "exposureHint": "Lade diese Quelle vorab in den Prompt oder lass den Agenten sie bei Bedarf als Werkzeug durchsuchen." }, + "graphRetrieval": { + "title": "Graph-Abruf", + "tag": "ohne Neuimport", + "seedStrategy": "Suche beginnt bei", + "seedStrategyHint": "Entitäten eignen sich für die meisten Dokumente. Beziehungen erreichen auch Entitäten, die in der Frage nicht vorkommen – ideal für Inhalte, die beschreiben, wie Dinge zusammenhängen.", + "seedEntities": "Entitäten (empfohlen)", + "seedRelationships": "Beziehungen", + "passageNodes": "Textabschnitte in die Suche einbeziehen", + "passageNodesHint": "Ein Abschnitt wird gefunden, wenn er zur Frage passt oder mit etwas Passendem verbunden ist. Am besten bei Graphen, die mit dieser Version erstellt wurden.", + "blendVector": "Mit Vektorsuche kombinieren", + "blendVectorHint": "Ergänzt Ergebnisse der Vektorsuche, damit kein Abschnitt verloren geht, den der Graph übersieht.", + "agentToolHint": "Agenten können diesen Beziehungen auch selbst folgen, wenn die Bereitstellung dieser Quelle „Suchwerkzeug auf Abruf“ ist oder ein agentischer Agent sie nutzt." + }, "prescreen": { "enable": "LLM-Vorfilterung aktivieren", "warning": "Ruft eine größere Kandidatenmenge ab und filtert sie mit einem LLM. Das erhöht Latenz und Kosten pro Anfrage.", diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index 03f0f0dd..ba67a1cb 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -272,6 +272,19 @@ }, "exposureHint": "Pre-fetch this source into the prompt, or let the agent search it on demand as a tool." }, + "graphRetrieval": { + "title": "Graph retrieval", + "tag": "no re-ingest", + "seedStrategy": "Start the walk from", + "seedStrategyHint": "Entities suit most documents. Relationships can reach an entity the question never names, and suit content that describes how things connect.", + "seedEntities": "Entities (recommended)", + "seedRelationships": "Relationships", + "passageNodes": "Include passages in the walk", + "passageNodesHint": "Lets a passage be found both by matching the question and by being connected to what does. Works best on graphs built with this version.", + "blendVector": "Blend with vector search", + "blendVectorHint": "Adds plain vector search results, so a passage the graph misses is not lost.", + "agentToolHint": "Agents can also follow these relationships themselves when this source's exposure is “On-demand search tool”, or when an agentic agent uses it." + }, "prescreen": { "enable": "Enable LLM prescreen", "warning": "Fetches a larger candidate set and uses an LLM to filter it. This adds query-time latency and cost.", diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json index c7b4f794..f30cadd7 100644 --- a/frontend/src/locale/es.json +++ b/frontend/src/locale/es.json @@ -268,6 +268,19 @@ }, "exposureHint": "Precarga esta fuente en el prompt, o deja que el agente la busque bajo demanda como herramienta." }, + "graphRetrieval": { + "title": "Recuperación por grafo", + "tag": "sin reingesta", + "seedStrategy": "Iniciar el recorrido desde", + "seedStrategyHint": "Las entidades funcionan para la mayoría de documentos. Las relaciones pueden llegar a una entidad que la pregunta no menciona; son ideales para contenido que describe cómo se conectan las cosas.", + "seedEntities": "Entidades (recomendado)", + "seedRelationships": "Relaciones", + "passageNodes": "Incluir fragmentos en el recorrido", + "passageNodesHint": "Un fragmento puede encontrarse por coincidir con la pregunta o por estar conectado con lo que coincide. Funciona mejor en grafos creados con esta versión.", + "blendVector": "Combinar con búsqueda vectorial", + "blendVectorHint": "Añade resultados de la búsqueda vectorial para no perder fragmentos que el grafo pase por alto.", + "agentToolHint": "Los agentes también pueden seguir estas relaciones por sí mismos cuando la exposición de esta fuente es «Herramienta de búsqueda bajo demanda» o cuando la usa un agente agéntico." + }, "prescreen": { "enable": "Habilitar preselección con LLM", "warning": "Obtiene un conjunto de candidatos más grande y usa un LLM para filtrarlo. Esto añade latencia y costo por consulta.", diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json index dc022fe8..e475f954 100644 --- a/frontend/src/locale/jp.json +++ b/frontend/src/locale/jp.json @@ -268,6 +268,19 @@ }, "exposureHint": "このソースをプロンプトに事前取得するか、エージェントがツールとして必要に応じて検索できるようにします。" }, + "graphRetrieval": { + "title": "グラフ検索", + "tag": "再取り込み不要", + "seedStrategy": "探索の開始点", + "seedStrategyHint": "ほとんどのドキュメントにはエンティティが適しています。リレーションは質問に登場しないエンティティにも到達でき、物事のつながりを説明するコンテンツに向いています。", + "seedEntities": "エンティティ(推奨)", + "seedRelationships": "リレーション", + "passageNodes": "パッセージを探索に含める", + "passageNodesHint": "質問に一致するパッセージだけでなく、一致したものとつながるパッセージも見つけられます。このバージョン以降に構築したグラフで最も効果的です。", + "blendVector": "ベクトル検索と組み合わせる", + "blendVectorHint": "ベクトル検索の結果を加え、グラフが見落としたパッセージも失わないようにします。", + "agentToolHint": "このソースの公開方法が「オンデマンド検索ツール」の場合、またはエージェント型エージェントが使用する場合、エージェントはこれらのリレーションを自ら辿ることもできます。" + }, "prescreen": { "enable": "LLMプリスクリーニングを有効にする", "warning": "より多くの候補を取得し、LLMでフィルタリングします。クエリ時のレイテンシとコストが増加します。", diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json index dd7aef1e..ba064d54 100644 --- a/frontend/src/locale/ru.json +++ b/frontend/src/locale/ru.json @@ -268,6 +268,19 @@ }, "exposureHint": "Предзагружать этот источник в промпт или позволить агенту искать по нему по мере необходимости как по инструменту." }, + "graphRetrieval": { + "title": "Поиск по графу", + "tag": "без повторной загрузки", + "seedStrategy": "Начинать обход с", + "seedStrategyHint": "Сущности подходят для большинства документов. Связи позволяют дойти до сущности, которая не упоминается в вопросе, — хорошо для контента о том, как всё связано.", + "seedEntities": "Сущностей (рекомендуется)", + "seedRelationships": "Связей", + "passageNodes": "Включать фрагменты в обход", + "passageNodesHint": "Фрагмент находится, если он соответствует вопросу или связан с тем, что соответствует. Лучше всего работает на графах, построенных в этой версии.", + "blendVector": "Сочетать с векторным поиском", + "blendVectorHint": "Добавляет результаты векторного поиска, чтобы не терять фрагменты, пропущенные графом.", + "agentToolHint": "Агенты также могут сами проходить по этим связям, если для источника выбран режим «Инструмент поиска по запросу» или его использует агентный агент." + }, "prescreen": { "enable": "Включить предварительный отбор LLM", "warning": "Извлекается расширенный набор кандидатов, который затем фильтруется с помощью LLM. Это увеличивает задержку и стоимость запроса.", diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json index 16236cec..c8c3e7e1 100644 --- a/frontend/src/locale/zh-TW.json +++ b/frontend/src/locale/zh-TW.json @@ -268,6 +268,19 @@ }, "exposureHint": "將此來源預先載入提示中,或讓代理以工具形式隨選搜尋。" }, + "graphRetrieval": { + "title": "圖譜檢索", + "tag": "無需重新匯入", + "seedStrategy": "走訪起點", + "seedStrategyHint": "實體適用於大多數文件。關係可以到達問題中未提及的實體,適合描述事物之間如何關聯的內容。", + "seedEntities": "實體(建議)", + "seedRelationships": "關係", + "passageNodes": "將段落納入走訪", + "passageNodesHint": "段落既可因符合問題而被找到,也可因與符合內容相連而被找到。在此版本之後建立的圖譜上效果最佳。", + "blendVector": "與向量檢索結合", + "blendVectorHint": "加入向量檢索結果,避免遺漏圖譜未找到的段落。", + "agentToolHint": "當此來源的公開方式為「隨選搜尋工具」,或由代理型代理使用時,代理也可以自行沿著這些關係查找。" + }, "prescreen": { "enable": "啟用 LLM 預篩選", "warning": "會擷取較大的候選集合並使用 LLM 篩選。這將增加查詢延遲與成本。", diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json index 5b4cf5a6..40e8fa0e 100644 --- a/frontend/src/locale/zh.json +++ b/frontend/src/locale/zh.json @@ -268,6 +268,19 @@ }, "exposureHint": "将此来源预取到提示词中,或让代理按需将其作为工具进行搜索。" }, + "graphRetrieval": { + "title": "图谱检索", + "tag": "无需重新导入", + "seedStrategy": "遍历起点", + "seedStrategyHint": "实体适用于大多数文档。关系可以到达问题中未提及的实体,适合描述事物之间如何关联的内容。", + "seedEntities": "实体(推荐)", + "seedRelationships": "关系", + "passageNodes": "将段落纳入遍历", + "passageNodesHint": "段落既可因匹配问题被找到,也可因与匹配内容相连而被找到。在此版本之后构建的图谱上效果最佳。", + "blendVector": "与向量检索结合", + "blendVectorHint": "加入向量检索结果,避免遗漏图谱未找到的段落。", + "agentToolHint": "当此来源的公开方式为“按需搜索工具”,或由智能体型代理使用时,代理也可以自行沿这些关系查找。" + }, "prescreen": { "enable": "启用 LLM 预筛选", "warning": "会获取更大的候选集并使用 LLM 进行过滤。这会增加查询时的延迟和成本。", diff --git a/frontend/src/models/misc.ts b/frontend/src/models/misc.ts index fc8b0782..c0c5efb2 100644 --- a/frontend/src/models/misc.ts +++ b/frontend/src/models/misc.ts @@ -32,6 +32,17 @@ export type SourcePrescreenConfig = { max_keep?: number; // default 8, <= candidate_k }; +// Where the graph walk starts: matching entities, or matching relationships +// ("A streams_to B"), which can reach an entity the question never names. +export type GraphSeedStrategy = 'entities' | 'relationships'; + +// Query-time graph retrieval knobs (graphrag only; live, no re-ingest). +export type SourceGraphRetrievalConfig = { + seed_strategy?: GraphSeedStrategy; // default 'entities' + passage_nodes?: boolean; // default true + blend_vector?: boolean; // default true +}; + // Query-time retrieval knobs (live; no re-ingest needed). export type SourceRetrievalConfig = { retriever?: string; // default 'classic' (only option for now) @@ -40,6 +51,7 @@ export type SourceRetrievalConfig = { score_threshold?: number | null; // default null rephrase_query?: boolean; // default true prescreen?: SourcePrescreenConfig | null; // null = off + graph?: SourceGraphRetrievalConfig; // graphrag retriever only }; // Ingest-time GraphRAG extraction knobs (only used when kind === 'graphrag'). diff --git a/frontend/src/settings/components/RetrievalOptions.test.tsx b/frontend/src/settings/components/RetrievalOptions.test.tsx index 92142b23..264c1794 100644 --- a/frontend/src/settings/components/RetrievalOptions.test.tsx +++ b/frontend/src/settings/components/RetrievalOptions.test.tsx @@ -146,6 +146,11 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => { batch_size: 5, max_keep: 10, }, + graph: { + seed_strategy: 'relationships', + passage_nodes: false, + blend_vector: false, + }, }, graph: { extraction_model: null, @@ -157,6 +162,44 @@ describe('round-trip configToOptions(optionsToConfig(x)) == x', () => { }); }); +describe('graph retrieval options', () => { + it('defaults to the measured-best configuration', () => { + expect(DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph).toEqual({ + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: true, + }); + }); + + it('fills the defaults for a source saved before the options existed', () => { + const opts = configToOptions({ retrieval: { retriever: 'graphrag' } }); + expect(opts.retrieval.graph).toEqual( + DEFAULT_RETRIEVAL_OPTIONS.retrieval.graph, + ); + }); + + it('honors stored options and fills only the missing ones', () => { + const opts = configToOptions({ + retrieval: { graph: { seed_strategy: 'relationships' } }, + }); + expect(opts.retrieval.graph).toEqual({ + seed_strategy: 'relationships', + passage_nodes: true, + blend_vector: true, + }); + }); + + it('writes the options into the retrieval block', () => { + const v = clone(DEFAULT_RETRIEVAL_OPTIONS); + v.retrieval.graph.blend_vector = false; + expect(optionsToConfig(v).retrieval?.graph).toEqual({ + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: false, + }); + }); +}); + describe('isPrescreenConfigValid', () => { const withPrescreen = ( chunks: number, diff --git a/frontend/src/settings/components/RetrievalOptions.tsx b/frontend/src/settings/components/RetrievalOptions.tsx index e2723358..b9dec5d6 100644 --- a/frontend/src/settings/components/RetrievalOptions.tsx +++ b/frontend/src/settings/components/RetrievalOptions.tsx @@ -16,6 +16,7 @@ import { import { Switch } from '../../components/ui/switch'; import type { ChunkingStrategy, + GraphSeedStrategy, RetrievalExposure, SourceConfig, } from '../../models/misc'; @@ -59,6 +60,11 @@ export type RetrievalOptionsValue = { batch_size: number; max_keep: number; }; + graph: { + seed_strategy: GraphSeedStrategy; + passage_nodes: boolean; + blend_vector: boolean; + }; }; graph: { extraction_model: string | null; @@ -85,6 +91,12 @@ export const DEFAULT_RETRIEVAL_OPTIONS: RetrievalOptionsValue = { enabled: false, ...DEFAULT_PRESCREEN, }, + // The configuration that measured best across the corpora tested. + graph: { + seed_strategy: 'entities', + passage_nodes: true, + blend_vector: true, + }, }, graph: { extraction_model: null, @@ -204,6 +216,7 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue { const chunking = config?.chunking ?? {}; const retrieval = config?.retrieval ?? {}; const prescreen = retrieval.prescreen ?? null; + const retrievalGraph = retrieval.graph ?? {}; const graph = config?.graph ?? {}; const d = DEFAULT_RETRIEVAL_OPTIONS; return { @@ -228,6 +241,14 @@ export function configToOptions(config?: SourceConfig): RetrievalOptionsValue { batch_size: prescreen?.batch_size ?? DEFAULT_PRESCREEN.batch_size, max_keep: prescreen?.max_keep ?? DEFAULT_PRESCREEN.max_keep, }, + graph: { + seed_strategy: + retrievalGraph.seed_strategy ?? d.retrieval.graph.seed_strategy, + passage_nodes: + retrievalGraph.passage_nodes ?? d.retrieval.graph.passage_nodes, + blend_vector: + retrievalGraph.blend_vector ?? d.retrieval.graph.blend_vector, + }, }, graph: { extraction_model: graph.extraction_model ?? d.graph.extraction_model, @@ -271,6 +292,11 @@ export function optionsToConfig(value: RetrievalOptionsValue): SourceConfig { max_keep: ps.max_keep, } : null, + graph: { + seed_strategy: value.retrieval.graph.seed_strategy, + passage_nodes: value.retrieval.graph.passage_nodes, + blend_vector: value.retrieval.graph.blend_vector, + }, }, graph: { extraction_model: value.graph.extraction_model?.trim() @@ -399,6 +425,12 @@ export default function RetrievalOptions({ }); }; + const setGraphRetrieval = ( + patch: Partial, + ) => { + setRetrieval({ graph: { ...value.retrieval.graph, ...patch } }); + }; + const modelOptions = useMemo(() => { const builtin: Model[] = []; const user: Model[] = []; @@ -611,6 +643,84 @@ export default function RetrievalOptions({ )}
+ {/* Graph retrieval group (graphrag only; live, so shown when testing too) */} + {isGraphRAG && ( +
+ +

+ {tr('graphRetrieval.agentToolHint')} +

+ +
+ + + + + + + setGraphRetrieval({ passage_nodes: checked }) + } + /> + + + + + setGraphRetrieval({ blend_vector: checked }) + } + /> + +
+
+ )} + {/* Graph extraction group (graphrag only; re-ingest required to apply) */} {isGraphRAG && !queryOnly && (
From 15bdda8554cc332230310d2f7fdfa9f9e9af80eb Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:42 +0100 Subject: [PATCH 086/130] fix(worker): know you are in a worker from any thread, not only the task's Celery records the executing task on the thread that runs it, so a thread that task starts sees none. The embeddings client and read_document both decided "am I in a worker?" from that alone, and from any other thread took the web-process branch: dispatch to the worker they were running in and block on the result. Celery refuses that get() ("Never call result.get() within a task!"), so the embed failed and latched the 30s dispatch cooldown for every caller after it; with joins allowed, read_document would instead wait on a parsing queue only its own busy process serves. Threads inside tasks are not hypothetical: per-source retrieval fans out to a pool, so a scheduled or webhook agent searching several sources embedded from pool threads. Graph extraction did too, which failed every chunk of a build. in_worker() in celery_init answers for the whole process. Celery's task_join_will_block is process-wide and set for every blocking pool (prefork, solo, threads) -- exactly the condition under which dispatch-and-wait goes wrong; eventlet/gevent leave it unset, so the task's own thread still counts through current_worker_task. Verified with real workers on each blocking pool: from a thread a task started, the old check dispatched and hit the error, the new one embedded locally. --- docsgpt/agents/tools/read_document.py | 9 ++-- docsgpt/celery_init.py | 23 ++++++++ docsgpt/vectorstore/embeddings_delegated.py | 11 ++-- tests/agents/tools/test_read_document_tool.py | 12 ++--- tests/test_celery.py | 52 +++++++++++++++++++ .../vectorstore/test_embeddings_delegated.py | 26 ++++++++++ 6 files changed, 117 insertions(+), 16 deletions(-) diff --git a/docsgpt/agents/tools/read_document.py b/docsgpt/agents/tools/read_document.py index ab9d2294..94e488df 100644 --- a/docsgpt/agents/tools/read_document.py +++ b/docsgpt/agents/tools/read_document.py @@ -17,8 +17,6 @@ import signal import threading from typing import Any, Callable, Dict, List, Optional -from celery import current_task - from docsgpt.agents.tools.artifact_ref import resolve_artifact_id from docsgpt.agents.tools.attachment_bridge import ( AttachmentBridgeError, @@ -26,6 +24,7 @@ from docsgpt.agents.tools.attachment_bridge import ( match_attachment, ) from docsgpt.agents.tools.base import Tool +from docsgpt.celery_init import in_worker from docsgpt.core.json_schema_utils import ( JsonSchemaValidationError, normalize_json_schema_payload, @@ -229,9 +228,9 @@ class ReadDocumentTool(Tool): # (floored at DOCUMENT_PARSE_TIMEOUT). timeout = parse_timeout_for_size(self._input_size) - # ``current_task`` is a Celery proxy: truthy only while this runs inside a worker task, - # falsy in the web process (the bare proxy is NOT identity-None, so test truthiness). - if current_task: + # Process-wide, not the thread-local ``current_task``: a thread a task starts has no + # task of its own, and dispatching from there is the self-deadlock described above. + if in_worker(): from docsgpt.worker import run_parse_document try: diff --git a/docsgpt/celery_init.py b/docsgpt/celery_init.py index 4272eac1..5e11d3e6 100644 --- a/docsgpt/celery_init.py +++ b/docsgpt/celery_init.py @@ -172,6 +172,29 @@ def _run_version_check(*args, **kwargs): celery = make_celery() celery.config_from_object("docsgpt.celeryconfig") + +def in_worker() -> bool: + """True anywhere in a Celery worker process, on any thread. + + ``current_worker_task`` alone is not enough: Celery records the executing + task on the thread that runs it, so a thread the task starts sees none and + would take the web-process branch — dispatching to the worker it is running + in and blocking on the result. Celery refuses that ``get()`` ("Never call + result.get() within a task!"), or, where joins are allowed, it waits on a + queue only this busy process serves. + + ``task_join_will_block`` is process-wide and set for every blocking pool + (prefork, solo, threads) — exactly the condition under which dispatching + and waiting goes wrong. eventlet/gevent leave it unset, so the task's own + thread still counts through ``current_worker_task``. + + Returns: + bool: Whether this call is running inside a worker process. + """ + from celery.result import task_join_will_block + + return task_join_will_block() or celery.current_worker_task is not None + #: Task-name prefix the package carried before the rename to ``docsgpt``. diff --git a/docsgpt/vectorstore/embeddings_delegated.py b/docsgpt/vectorstore/embeddings_delegated.py index 75f730b8..bc184f34 100644 --- a/docsgpt/vectorstore/embeddings_delegated.py +++ b/docsgpt/vectorstore/embeddings_delegated.py @@ -10,8 +10,9 @@ Celery and the vector comes back. The API pays a broker round trip per query and no resident model. Inside a worker there is nothing to delegate to -- dispatching would queue work -behind the task already running and wait on itself -- so a call made while a -task is executing runs locally, on a model this process loads once and caches. +behind the task already running and wait on itself -- so a call made anywhere in +a worker process, including from a thread a task started, runs locally, on a +model this process loads once and caches. ``DOCUMENT_PARSE_QUEUE`` exists for the same reason on the parsing side. Production deployments should point ``EMBEDDINGS_BASE_URL`` at a real embedding @@ -79,11 +80,11 @@ def _forget(result) -> None: def _in_worker() -> bool: - """True when a Celery task is executing in this process.""" + """True anywhere in a Celery worker process -- on any thread, not only the task's.""" try: - from docsgpt.celery_init import celery + from docsgpt.celery_init import in_worker - return celery.current_worker_task is not None + return in_worker() except Exception: return False diff --git a/tests/agents/tools/test_read_document_tool.py b/tests/agents/tools/test_read_document_tool.py index 8913495d..f0e7b177 100644 --- a/tests/agents/tools/test_read_document_tool.py +++ b/tests/agents/tools/test_read_document_tool.py @@ -357,9 +357,9 @@ def test_malformed_json_schema_rejected_before_enqueue(monkeypatch): @pytest.mark.unit def test_dispatch_inline_when_in_worker(monkeypatch): _stub_repo(monkeypatch, found=True, conv="conv-1", run=None) - # Inside a worker current_task is truthy -> parse inline, never enqueue (else the + # Inside a worker -> parse inline, never enqueue (else the # parsing queue self-deadlocks the worker that also serves it). - monkeypatch.setattr(rd, "current_task", object()) + monkeypatch.setattr(rd, "in_worker", lambda: True) import docsgpt.api.user.tasks as tasks monkeypatch.setattr( @@ -387,8 +387,8 @@ def test_dispatch_inline_when_in_worker(monkeypatch): @pytest.mark.unit def test_dispatch_enqueues_when_not_in_worker(monkeypatch): _stub_repo(monkeypatch, found=True, conv="conv-1", run=None) - # Web process: current_task falsy -> dispatch to the parsing queue, never inline. - monkeypatch.setattr(rd, "current_task", None) + # Web process -> dispatch to the parsing queue, never inline. + monkeypatch.setattr(rd, "in_worker", lambda: False) captured = _patch_task(monkeypatch, payload={"status": "ok", "content": "queued", "truncated": False}) import docsgpt.worker as worker @@ -414,9 +414,9 @@ _TIMED_OUT = "document parsing timed out after" def _inline(monkeypatch, run_parse, *, timeout=0.2) -> ReadDocumentTool: - """Drive the inline branch (current_task truthy) with a patched parse window.""" + """Drive the inline (in-worker) branch with a patched parse window.""" _stub_repo(monkeypatch, found=True, conv="conv-1", run=None) - monkeypatch.setattr(rd, "current_task", object()) + monkeypatch.setattr(rd, "in_worker", lambda: True) import docsgpt.api.user.tasks as tasks monkeypatch.setattr( diff --git a/tests/test_celery.py b/tests/test_celery.py index c5b692df..5a3e66be 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -276,3 +276,55 @@ class TestReclaimIsSkippedForEmbeds: from docsgpt.vectorstore.embeddings_delegated import EMBED_TASK assert EMBED_TASK in _NO_RECLAIM_TASKS + + +@pytest.mark.unit +class TestInWorker: + """Whether code runs inside a worker must not depend on which thread asks. + + Celery records the executing task on the thread that runs it, so a thread + that task starts sees no task at all. Code deciding "am I in the worker?" + from that alone takes the web-process branch there: it dispatches to the + worker it is running in and blocks on the result, which Celery refuses + ("Never call result.get() within a task!") or, where joins are allowed, + waits on a queue only this busy process serves. + """ + + @staticmethod + def _ask_from_a_new_thread(): + import threading + + from docsgpt.celery_init import in_worker + + seen = [] + thread = threading.Thread(target=lambda: seen.append(in_worker())) + thread.start() + thread.join() + return seen[0] + + def test_false_outside_a_worker(self): + from docsgpt.celery_init import in_worker + + assert in_worker() is False + assert self._ask_from_a_new_thread() is False + + def test_true_on_a_thread_started_inside_a_worker(self): + # Blocking pools (prefork, solo, threads) mark the whole process as one + # where joining a task would block; ``denied_join_result`` sets exactly + # that flag. + from celery.result import denied_join_result + + with denied_join_result(): + assert self._ask_from_a_new_thread() is True + + def test_true_on_the_task_thread_of_a_non_blocking_pool(self): + # eventlet/gevent pools leave the process flag unset; the thread + # running the task still knows it is in one. + from unittest.mock import PropertyMock + + from docsgpt.celery_init import celery, in_worker + + with patch.object( + type(celery), "current_worker_task", new_callable=PropertyMock, return_value=object() + ): + assert in_worker() is True diff --git a/tests/vectorstore/test_embeddings_delegated.py b/tests/vectorstore/test_embeddings_delegated.py index 060717ea..b9759510 100644 --- a/tests/vectorstore/test_embeddings_delegated.py +++ b/tests/vectorstore/test_embeddings_delegated.py @@ -72,6 +72,32 @@ class TestInsideAWorker: assert vector == [1.0, 2.0] celery.send_task.assert_not_called() + def test_a_thread_started_inside_the_worker_embeds_locally(self): + """The task's own thread is not the only one in a worker. + + Graph extraction and per-source retrieval both fan out to thread pools + inside tasks. The check used to read the task off the current thread + only, so from those threads it dispatched to the worker it was running + in -- and Celery refuses that ``get()`` inside a worker, failing the + call and latching the 30s dispatch cooldown for every caller after it. + """ + from celery.result import denied_join_result + + from docsgpt.celery_init import celery + + local = MagicMock() + local.embed_documents.return_value = [[1.0, 2.0]] + client = DelegatedEmbeddings("some/model") + vectors = [] + with denied_join_result(): + with patch("docsgpt.vectorstore.base.build_local_embeddings", return_value=local): + with patch.object(celery, "send_task") as send_task: + thread = threading.Thread(target=lambda: vectors.append(client.embed_query("hi"))) + thread.start() + thread.join() + assert vectors == [[1.0, 2.0]] + send_task.assert_not_called() + def test_the_local_model_is_built_once(self): local = MagicMock() local.embed_documents.return_value = [[1.0]] From 5285bb4115e050fa10051a7db522de3bfd5dd70c Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:07:42 +0100 Subject: [PATCH 087/130] fix(graphrag): give a failed chunk one more attempt before giving up A chunk whose extraction failed was marked failed and the build moved on. The checkpoint treats failed chunks as pending, but nothing ever reran the build, so one transient error -- a provider hiccup, a single response that did not parse -- left a permanent hole in the graph until someone rebuilt the whole source. Failed chunks now get one more attempt after the rest of the build, so a burst of rate limiting has time to pass. Extraction errors, unparseable responses and failed writes are all retried; a chunk is marked failed only when its retry fails too, so it costs at most two calls. The chunks given up on are logged by id, and progress still ends at the total. --- docsgpt/graphrag/extraction.py | 129 ++++++++++++--------- tests/graphrag/test_extraction.py | 181 ++++++++++++++++++++++++++++-- 2 files changed, 247 insertions(+), 63 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index db8ea27a..9539f4fd 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -170,13 +170,13 @@ def _extract_chunk( ) except Exception as exc: logger.warning( - "Graph extraction call failed for chunk %s, skipping: %s", chunk_id, exc + "Graph extraction call failed for chunk %s: %s", chunk_id, exc ) return None parsed = _parse_extraction(response) if parsed is None: logger.warning( - "Graph extraction returned unparseable output for chunk %s; marking it failed.", + "Graph extraction returned unparseable output for chunk %s.", chunk_id, ) return parsed @@ -203,8 +203,10 @@ def extract_graph_for_source( Resumable and idempotent: chunks already marked ``done`` are skipped via the ``graph_ingest_progress`` checkpoint, so a retry never re-extracts (and never re-bills). Processes at most the resolved chunk cap; excess chunks are - reported under ``skipped_over_cap``. A malformed response or an LLM error on - a single chunk marks it ``failed`` and continues — the pipeline never crashes. + reported under ``skipped_over_cap``. A malformed response, an LLM error or a + failed write on a single chunk is retried once after the rest of the build; + a chunk that fails again is marked ``failed`` and the run continues — the + pipeline never crashes. Each chunk is written in a single transaction with one batched embedding call (entity + relationship-endpoint names together). @@ -273,13 +275,9 @@ def extract_graph_for_source( """One chunk's LLM extraction — the only step run concurrently. A chunk spends almost all of its time waiting on the model, so that is - what runs in the pool. Everything else stays on the calling thread: - graph writes, so transactions and the progress checkpoint are exactly - what they were serially, and embedding. Inside a Celery worker the - embeddings client decides to embed locally from the task on the - *current thread's* stack; a pool thread has none, so it would instead - dispatch an embed task to the worker and wait on it, which Celery - refuses inside a task — failing every chunk of the build. + what runs in the pool. Graph writes and embedding stay on the calling + thread, so transactions and the progress checkpoint are exactly what + they were serially and the pool never touches the embeddings client. """ chunk, chunk_id = item text = _chunk_text(chunk) @@ -294,56 +292,85 @@ def extract_graph_for_source( relationships = _build_relationships(extracted["relationships"]) except Exception as exc: logger.warning( - "Graph extraction failed for chunk %s, skipping: %s", chunk_id, exc + "Graph extraction failed for chunk %s: %s", chunk_id, exc ) return chunk_id, "failed", None return chunk_id, "ok", (entities, relationships) + def _write(chunk_id, status, payload) -> bool: + """Apply one prepared chunk to the graph; False when it did not land.""" + nonlocal node_upserts, edges, chunks_processed + if status == "empty": + store.mark_chunk(source_id, chunk_id, "done") + chunks_processed += 1 + return True + if status == "failed": + return False + + entities, relationships = payload + try: + name_embeddings = _embed_names(embedding, entities, relationships) + _embed_facts(embedding, relationships) + chunk_nodes, chunk_edges = store.apply_chunk( + source_id, chunk_id, entities, relationships, name_embeddings + ) + except Exception as exc: + logger.warning( + "Graph extraction embed/write failed for chunk %s: %s", chunk_id, exc + ) + return False + # ``apply_chunk`` marks the chunk done inside the transaction that + # writes its rows, so the checkpoint cannot disagree with the graph + # and a replayed write cannot apply the chunk twice. + node_upserts += chunk_nodes + edges += chunk_edges + chunks_processed += 1 + return True + workers = max(1, int(getattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 1) or 1)) pool = None if workers > 1 and len(to_process) > 1: pool = ThreadPoolExecutor(max_workers=workers) - # ``map`` yields in submission order, so chunks are still applied in the - # order they were given and a run stays reproducible. - prepared = pool.map(_prepare, to_process) - else: - prepared = (_prepare(item) for item in to_process) + + def _pass(items): + """Extract and write ``items``; return the ones that did not land.""" + if pool is not None: + # ``map`` yields in submission order, so chunks are still applied in + # the order they were given and a run stays reproducible. + prepared = pool.map(_prepare, items) + else: + prepared = (_prepare(item) for item in items) + missed = [] + for item, (chunk_id, status, payload) in zip(items, prepared): + if not _write(chunk_id, status, payload): + missed.append(item) + _report() + return missed try: - for chunk_id, status, payload in prepared: - if status == "empty": - store.mark_chunk(source_id, chunk_id, "done") - chunks_processed += 1 - _report() - continue - if status == "failed": - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 - _report() - continue - - entities, relationships = payload - try: - # On this thread, not in the pool — see ``_prepare``. - name_embeddings = _embed_names(embedding, entities, relationships) - _embed_facts(embedding, relationships) - chunk_nodes, chunk_edges = store.apply_chunk( - source_id, chunk_id, entities, relationships, name_embeddings - ) - node_upserts += chunk_nodes - edges += chunk_edges - # ``apply_chunk`` marks the chunk done inside the transaction that - # writes its rows, so the checkpoint cannot disagree with the graph - # and a replayed write cannot apply the chunk twice. - chunks_processed += 1 - except Exception as exc: - logger.warning( - "Graph extraction embed/write failed for chunk %s, skipping: %s", - chunk_id, - exc, - ) - store.mark_chunk(source_id, chunk_id, "failed") - failed_chunks += 1 + missed = _pass(to_process) + if missed: + # A failure is usually transient — a provider error, one response + # that did not parse — and the checkpoint only picks it up on a + # rerun nothing schedules. One more attempt, after the rest of the + # build so a burst of rate limiting has passed, and no more: a + # chunk that cannot be extracted costs at most two calls. + logger.info( + "Graph extraction retrying %d failed chunk(s) for source %s", + len(missed), + source_id, + ) + missed = _pass(missed) + for _, chunk_id in missed: + store.mark_chunk(source_id, chunk_id, "failed") + failed_chunks = len(missed) + if missed: + logger.warning( + "Graph extraction gave up on %d chunk(s) for source %s after a retry: %s", + failed_chunks, + source_id, + ", ".join(str(chunk_id) for _, chunk_id in missed), + ) _report() finally: if pool is not None: diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index eda8dd7c..0573cb7c 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -88,6 +88,37 @@ class _StubLLM: return response +class _ScriptedLLM: + """Stub LLM answering per chunk, so results do not depend on call order. + + The extraction pool runs calls concurrently, so a stub that hands out + responses in call order gives each chunk whichever response its thread + happened to grab first. ``script`` maps a chunk's text to the responses + for that chunk, consumed one per call. + """ + + def __init__(self, script): + self._script = {text: list(responses) for text, responses in script.items()} + self.model_id = "stub-model" + self.calls = [] + self._token_usage_source = None + self._request_id = None + + def gen(self, model=None, messages=None, **kwargs): + text = messages[-1]["content"].removeprefix("\n").removesuffix("\n") + self.calls.append(text) + responses = self._script.get(text) + if not responses: + raise AssertionError(f"unexpected extraction call for {text!r}") + response = responses.pop(0) + if isinstance(response, Exception): + raise response + return response + + def calls_for(self, text): + return self.calls.count(text) + + class _StubEmbedding: """Stub embeddings model producing deterministic fixed-dim vectors.""" @@ -365,11 +396,12 @@ class TestExtractionLive: ): """Only the LLM call may run in the extraction pool, never embedding. - Inside a Celery worker the embeddings client decides to embed locally - from the task on the *current thread's* stack. A pool thread has none, - so from there it dispatches an embed task to the worker and waits on - it — which Celery refuses inside a task, so every chunk of a graph - build failed. + Embedding from the pool is what broke every graph build inside a + worker: the embeddings client used to decide "embed locally" from the + task on the *current thread's* stack, so a pool thread dispatched to + the worker instead and Celery refused the wait. ``in_worker`` is + process-wide now, but the pool still has no reason to touch the + embeddings client — it exists to overlap model latency. """ import threading @@ -505,11 +537,13 @@ class TestExtractionLive: entities=[{"name": "Ada", "type": "person", "description": "d"}], relationships=[], ) - llm = _StubLLM([ - "not json at all", - RuntimeError("model exploded"), - good, - ]) + # Each failing chunk fails its retry too; one that recovers on + # retry is covered in ``TestFailedChunksAreRetried``. + llm = _ScriptedLLM({ + "garbage": ["not json at all", "still not json"], + "boom": [RuntimeError("model exploded"), RuntimeError("model exploded again")], + "Ada.": [good], + }) _install_stub_llm(monkeypatch, llm) summary = extract_graph_for_source( @@ -762,7 +796,7 @@ class TestFailedChunksAreReported: import logging store = self._fake_store(monkeypatch, ["c1"]) - _install_stub_llm(monkeypatch, _StubLLM(["not json at all"])) + _install_stub_llm(monkeypatch, _StubLLM(["not json at all", "still not json"])) with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): summary = extract_graph_for_source( @@ -785,7 +819,10 @@ class TestFailedChunksAreReported: import logging self._fake_store(monkeypatch, ["c7"]) - _install_stub_llm(monkeypatch, _StubLLM([RuntimeError("model exploded")])) + _install_stub_llm( + monkeypatch, + _StubLLM([RuntimeError("model exploded"), RuntimeError("model exploded again")]), + ) with caplog.at_level(logging.WARNING, logger="docsgpt.graphrag.extraction"): extract_graph_for_source( @@ -800,6 +837,126 @@ class TestFailedChunksAreReported: assert any("c7" in message for message in messages), messages +@pytest.mark.unit +class TestFailedChunksAreRetried: + """A chunk that fails once gets one more attempt before the build ends. + + Failures are recorded as ``failed`` and the checkpoint treats them as + pending, but nothing ever ran the build again, so a single transient error + — a provider hiccup, one response that did not parse — left a permanent + hole in the graph until someone rebuilt the whole source. Retries run + after the rest of the build, which gives a burst of rate limiting time to + pass, and are bounded at one per chunk so a chunk that can never be + extracted costs at most two calls. + """ + + GOOD = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + + def _fake_store(self, monkeypatch, chunk_ids): + from unittest.mock import MagicMock + + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = list(chunk_ids) + store.apply_chunk.return_value = (1, 0) + store.count_nodes.return_value = 1 + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda *a, **k: store + ) + return store + + def _run(self, chunks, progress=None): + return extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=chunks, + config=SourceConfig(), + request_id="req-retry", + progress_cb=progress, + ) + + @staticmethod + def _marked_failed(store): + return [c.args[1] for c in store.mark_chunk.call_args_list if c.args[2] == "failed"] + + def test_a_transient_failure_is_retried_and_written(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + llm = _ScriptedLLM({"flaky": [RuntimeError("rate limited"), self.GOOD]}) + _install_stub_llm(monkeypatch, llm) + + summary = self._run([_chunk("c1", "flaky")]) + + assert summary["failed_chunks"] == 0 + assert summary["chunks_processed"] == 1 + assert store.apply_chunk.call_args.args[1] == "c1" + assert self._marked_failed(store) == [] + + def test_an_unparseable_response_is_retried(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + _install_stub_llm(monkeypatch, _ScriptedLLM({"odd": ["not json", self.GOOD]})) + + summary = self._run([_chunk("c1", "odd")]) + + assert summary["failed_chunks"] == 0 + assert self._marked_failed(store) == [] + + def test_a_chunk_that_fails_again_is_marked_failed_once(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + llm = _ScriptedLLM({"broken": ["not json", "still not json"]}) + _install_stub_llm(monkeypatch, llm) + + summary = self._run([_chunk("c1", "broken")]) + + assert summary["failed_chunks"] == 1 + assert summary["chunks_processed"] == 0 + assert llm.calls_for("broken") == 2 + assert self._marked_failed(store) == ["c1"] + + def test_only_failed_chunks_are_retried(self, monkeypatch, stub_embedding): + from docsgpt.core.settings import settings + + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + self._fake_store(monkeypatch, ["c1", "c2", "c3"]) + llm = _ScriptedLLM({ + "one": [self.GOOD], + "two": [RuntimeError("timeout"), self.GOOD], + "three": [self.GOOD], + }) + _install_stub_llm(monkeypatch, llm) + + summary = self._run([_chunk("c1", "one"), _chunk("c2", "two"), _chunk("c3", "three")]) + + assert summary["chunks_processed"] == 3 + assert summary["failed_chunks"] == 0 + assert (llm.calls_for("one"), llm.calls_for("two"), llm.calls_for("three")) == (1, 2, 1) + + def test_a_failed_write_is_retried(self, monkeypatch, stub_embedding): + store = self._fake_store(monkeypatch, ["c1"]) + store.apply_chunk.side_effect = [RuntimeError("write failed"), (1, 0)] + _install_stub_llm(monkeypatch, _ScriptedLLM({"text": [self.GOOD, self.GOOD]})) + + summary = self._run([_chunk("c1", "text")]) + + assert summary["failed_chunks"] == 0 + assert summary["chunks_processed"] == 1 + assert self._marked_failed(store) == [] + + def test_progress_ends_at_the_total(self, monkeypatch, stub_embedding): + self._fake_store(monkeypatch, ["c1", "c2"]) + _install_stub_llm(monkeypatch, _ScriptedLLM({ + "fine": [self.GOOD], + "broken": ["not json", "still not json"], + })) + events = [] + + self._run([_chunk("c1", "fine"), _chunk("c2", "broken")], progress=events.append) + + assert all(e["current"] <= e["total"] for e in events) + assert events[-1]["current"] == events[-1]["total"] == 2 + + @pytest.mark.integration class TestSummaryNodeCount: """``nodes`` must describe the graph, not the number of upserts.""" From 022cf69b491048b74c84fff29d9b53b334a3d653 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 088/130] fix(graphrag): fold an entity's singular and plural onto one key canonical_name dropped "es" from every -ches/-ses/-zes plural, so "caches" became "cach" while "cache" stayed "cache" -- the singular and plural landed on two nodes, which is the split the function exists to prevent. The same rule split the words documentation uses most: databases/database, responses/ response, releases/release, sizes/size. An "-es" plural cannot say whether it is cache + "s" or batch + "es", so instead of guessing, both sides now meet at the stem: the singular endings -che/-she/-se/-ze/-xe drop their "e" the way the plurals drop "es", and a singular's -ie folds to -y as -ies already did (cookie/cookies). The key is a merge key that is never shown, so it only has to agree, not be a word. alias, canvas, atlas and bias join the words that only look plural. Found by the naming tests CI was missing: the module had no direct tests. --- docsgpt/graphrag/naming.py | 39 ++++++++++++---- tests/graphrag/test_naming.py | 88 +++++++++++++++++++++++++++++++++++ 2 files changed, 118 insertions(+), 9 deletions(-) create mode 100644 tests/graphrag/test_naming.py diff --git a/docsgpt/graphrag/naming.py b/docsgpt/graphrag/naming.py index 3dbfa1cb..2a173eea 100644 --- a/docsgpt/graphrag/naming.py +++ b/docsgpt/graphrag/naming.py @@ -13,6 +13,9 @@ plural. Cautious matters: this corpus contains ``postgres``, ``kubernetes``, ``https`` and ``aws``, none of which are plurals, so a naive "strip trailing s" would corrupt them into new entities rather than merge anything. +The result is a merge key, never shown to anyone, so it only has to be the +same for a word's singular and plural — not to be a word itself. + Always on: every graph is built with canonical names. """ @@ -33,18 +36,30 @@ _NOT_PLURAL = frozenset( "analysis", "basis", "axis", "https", "rss", "less", "express", "redis", "nats", "kibana", "elasticsearch", "os", "ios", "macos", "always", "sometimes", "series", "docs", "ops", "devops", "sse", + "alias", "canvas", "atlas", "bias", "pandas", } ) +#: Plural endings that drop ``es``, and the singular endings that meet them. +#: ``caches`` cannot say whether it is ``cache`` + "s" or ``cach`` + "es" +#: (as ``batches`` is ``batch`` + "es"), so rather than guess, both +#: ``caches`` and ``cache`` fold to ``cach`` — as ``databases``/``database`` +#: fold to ``databas``. Hardly any real word differs from one of these singulars +#: by its final "e" alone, so the fold merges next to nothing it should not. +_ES_PLURAL = ("ches", "shes", "ses", "zes", "xes") +_E_SINGULAR = ("che", "she", "se", "ze", "xe") + def _singular(word: str) -> str: - """Best-effort singular of one word, biased hard towards leaving it alone. + """Fold one word so its singular and plural share a key, else leave it alone. - Only the endings that are unambiguous in this domain are touched: - ``-ies`` -> ``-y`` (``policies``), ``-ses``/``-xes``/``-zes``/``-ches``/ - ``-shes`` -> drop ``es`` (``indexes``, ``batches``), and a bare trailing - ``s`` on a word long enough to be safe. Everything in :data:`_NOT_PLURAL`, - and anything ending in ``ss``/``us``/``is``, is returned unchanged. + ``-ies`` and a singular's ``-ie`` both fold to ``-y`` (``policies``, + ``cookies``/``cookie``). The ``-es`` endings in :data:`_ES_PLURAL` drop + ``es`` and the singular endings in :data:`_E_SINGULAR` drop their ``e``, so + both sides of an ambiguous plural meet (``caches``/``cache`` -> ``cach``). + Otherwise a bare trailing ``s`` is dropped on a word long enough to be + safe. Everything in :data:`_NOT_PLURAL`, and anything ending in + ``ss``/``us``/``is``, is returned unchanged. """ if len(word) < 4 or word in _NOT_PLURAL: return word @@ -52,8 +67,12 @@ def _singular(word: str) -> str: return word if word.endswith("ies") and len(word) > 4: return word[:-3] + "y" - if word.endswith(("ses", "xes", "zes", "ches", "shes")): + if word.endswith(_ES_PLURAL): return word[:-2] + if word.endswith("ie") and len(word) > 4: + return word[:-2] + "y" + if word.endswith(_E_SINGULAR): + return word[:-1] if word.endswith("s"): return word[:-1] return word @@ -66,12 +85,14 @@ def canonical_name(name: str) -> str: name: The entity name as the model wrote it. Returns: - A lowercase, punctuation-free, singularised key. Returns ``""`` for an - empty or punctuation-only name, which callers treat as "no entity". + A lowercase, punctuation-free key shared by a name's singular and + plural. Returns ``""`` for an empty or punctuation-only name, which + callers treat as "no entity". Examples: ``VECTOR_STORE`` and ``Vector stores`` -> ``vector store``; ``.env file`` and ``env_file`` -> ``env file``; + ``cache`` and ``caches`` -> ``cach``; ``postgres`` stays ``postgres``. """ if not name: diff --git a/tests/graphrag/test_naming.py b/tests/graphrag/test_naming.py new file mode 100644 index 00000000..ff14f36a --- /dev/null +++ b/tests/graphrag/test_naming.py @@ -0,0 +1,88 @@ +"""Tests for canonical entity naming (the key graph nodes are merged on). + +Two failure directions matter. Too little folding splits one entity across +nodes ("agent" / "agents", "VECTOR_STORE" / "vector stores"), so the walk never +connects what the text connects. Too much folding invents entities: stripping +the "s" off ``postgres`` or ``redis`` would merge nothing and create a node no +chunk ever named. +""" + +from __future__ import annotations + +import pytest + +from docsgpt.graphrag.naming import canonical_name, normalize_entity_name + + +@pytest.mark.unit +class TestCanonicalName: + @pytest.mark.parametrize( + "variants, key", + [ + (["VECTOR_STORE", "Vector store", "vector stores", "vector-stores"], "vector store"), + ([".env file", "env_file", "ENV FILE"], "env file"), + (["Celery worker", "Celery workers"], "celery worker"), + (["agent", "Agents", "agents!"], "agent"), + ], + ) + def test_orthographic_variants_share_one_key(self, variants, key): + assert {canonical_name(v) for v in variants} == {key} + + @pytest.mark.parametrize( + "singular, plural", + [ + ("policy", "policies"), + ("index", "indexes"), + ("batch", "batches"), + ("hash", "hashes"), + ("class", "classes"), + ("process", "processes"), + ("status", "statuses"), + ("bus", "buses"), + ("alias", "aliases"), + ("document", "documents"), + ("service", "services"), + # Singulars ending in "e" whose plural also ends in "-es": the + # plural alone cannot say whether to drop "s" or "es". + ("cache", "caches"), + ("database", "databases"), + ("response", "responses"), + ("release", "releases"), + ("case", "cases"), + ("size", "sizes"), + ("cookie", "cookies"), + ], + ) + def test_singular_and_plural_share_one_key(self, singular, plural): + assert canonical_name(singular) == canonical_name(plural) + + def test_the_key_need_not_be_a_word(self): + # It is a merge key, never shown: "cache" and "caches" meet at the + # stem an "-es" plural cannot see past, rather than guessing a form. + assert canonical_name("caches") == "cach" + assert canonical_name("batches") == "batch" + + @pytest.mark.parametrize( + "word", + ["postgres", "kubernetes", "redis", "https", "status", "analysis", "access", "docs", "series"], + ) + def test_words_that_only_look_plural_are_left_alone(self, word): + assert canonical_name(word) == word + + @pytest.mark.parametrize("word", ["class", "corpus", "thesis"]) + def test_ss_us_is_endings_are_never_stripped(self, word): + assert canonical_name(word) == word + + @pytest.mark.parametrize("word", ["aws", "ids", "ops"]) + def test_short_words_are_left_alone(self, word): + assert canonical_name(word) == word + + def test_each_word_of_a_phrase_is_folded(self): + assert canonical_name("Postgres Replicas") == "postgres replica" + + @pytest.mark.parametrize("name", [None, "", " ", "!!!", "--_--"]) + def test_a_name_with_nothing_left_is_no_entity(self, name): + assert canonical_name(name) == "" + + def test_normalize_entity_name_is_the_canonical_key(self): + assert normalize_entity_name("Vector Stores") == canonical_name("Vector Stores") From a20402a6dc09acb3a51e31f05437cf54be208b35 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 089/130] fix(graphrag): give each extraction thread its own LLM Provider-reported usage is kept on the LLM instance (_last_usage) and claimed by whichever call finishes next. With GRAPHRAG_EXTRACTION_WORKERS > 1 (the default is 8) every extraction thread shared one instance, so a call could claim another call's provider counts while its own fell back to the estimate: token_usage rows, and the cost they bill, could be attributed to the wrong call and summed wrong. Each pool thread now builds its own extraction LLM on first use. The calling thread's instance is still built up front, so a misconfigured model fails the run before any chunk is touched. --- docsgpt/graphrag/extraction.py | 24 +++++++++++--- tests/graphrag/test_extraction.py | 55 +++++++++++++++++++++++++++++++ 2 files changed, 75 insertions(+), 4 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 9539f4fd..5ddbe1e1 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -227,6 +227,7 @@ def extract_graph_for_source( source's graph holds after the run — not how many upserts ran, which counts the same entity once per chunk it appears in. """ + import threading from concurrent.futures import ThreadPoolExecutor from docsgpt.graphrag.store import GraphStore @@ -246,9 +247,24 @@ def extract_graph_for_source( embedding = get_embeddings() - llm = _build_extraction_llm( - _resolve_extraction_model(config), user, request_id - ) + model_id = _resolve_extraction_model(config) + # Built here first so a misconfigured model fails the run before any + # chunk is touched; this instance serves the calling thread. + thread_llm = threading.local() + thread_llm.llm = _build_extraction_llm(model_id, user, request_id) + + def _llm(): + """This thread's extraction LLM. + + Provider-reported usage is kept on the LLM instance (``_last_usage``) + and claimed by whichever call finishes next, so two calls in flight on + one instance can bill each other's tokens. Each pool thread therefore + builds its own. + """ + llm = getattr(thread_llm, "llm", None) + if llm is None: + llm = thread_llm.llm = _build_extraction_llm(model_id, user, request_id) + return llm node_upserts = 0 edges = 0 @@ -284,7 +300,7 @@ def extract_graph_for_source( if not text: return chunk_id, "empty", None - extracted = _extract_chunk(llm, text, chunk_id) + extracted = _extract_chunk(_llm(), text, chunk_id) if extracted is None: return chunk_id, "failed", None try: diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index 0573cb7c..05bb0a40 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -606,6 +606,61 @@ class TestExtractionTokenUsage: assert built._request_id == "req-99" assert captured["model_id"] == "stub-model" + def test_concurrent_extraction_calls_never_share_an_llm(self, monkeypatch, stub_embedding): + """Provider usage is recorded on the LLM instance (``_last_usage``) and + claimed by whichever call finishes next, so two calls in flight on one + instance can bill each other's tokens. Each extraction thread needs its + own instance.""" + import threading + import time + from unittest.mock import MagicMock + + from docsgpt.core.settings import settings + + payload = _extraction_json( + entities=[{"name": "Ada", "type": "person", "description": "d"}], + relationships=[], + ) + + class _ThreadRecordingLLM: + model_id = "stub-model" + + def __init__(self): + self.threads = set() + + def gen(self, model=None, messages=None, **kwargs): + self.threads.add(threading.get_ident()) + time.sleep(0.01) # keep calls overlapping + return payload + + built = [] + + def _create(*args, **kwargs): + llm = _ThreadRecordingLLM() + built.append(llm) + return llm + + monkeypatch.setattr(extraction_module.LLMCreator, "create_llm", staticmethod(_create)) + monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4) + store = MagicMock(name="GraphStore") + store.pending_chunks.return_value = [f"c{i}" for i in range(8)] + store.apply_chunk.return_value = (1, 0) + store.count_nodes.return_value = 1 + monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", lambda *a, **k: store) + + summary = extract_graph_for_source( + str(uuid.uuid4()), + user="owner-1", + chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(8)], + config=SourceConfig(), + request_id="req-threads", + ) + + assert summary["chunks_processed"] == 8 + used = [llm for llm in built if llm.threads] + assert len(used) > 1, "calls did not run concurrently" + assert all(len(llm.threads) == 1 for llm in used) + @pytest.mark.unit class TestModelResolution: From 877609dfb6b9b64c767cc3879aa6ea61cf3c8a68 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 090/130] fix(worker): record worker state at startup, for every pool in_worker() read task_join_will_block, which eventlet and gevent leave unset, and the task's current_worker_task, which they scope to one greenlet -- so a greenlet a task spawned in those pools still took the web-process branch and dispatched to its own worker, where it could queue behind its parent and time out. The worker's own startup now records it: worker_init fires in every worker's main process before the pool starts (where solo, threads, eventlet and gevent run tasks, and what prefork children fork from), and worker_process_init in each prefork child. worker_ready would be too late -- prefork children are forked before it fires. The two existing checks stay for anything that runs tasks without that startup. --- docsgpt/celery_init.py | 45 +++++++++++++++++++++++++++++++----------- tests/test_celery.py | 22 +++++++++++++++++++++ 2 files changed, 56 insertions(+), 11 deletions(-) diff --git a/docsgpt/celery_init.py b/docsgpt/celery_init.py index 5e11d3e6..2eda5a3a 100644 --- a/docsgpt/celery_init.py +++ b/docsgpt/celery_init.py @@ -13,6 +13,7 @@ from celery.signals import ( setup_logging, task_postrun, task_prerun, + worker_init, worker_process_init, worker_ready, ) @@ -173,27 +174,49 @@ celery = make_celery() celery.config_from_object("docsgpt.celeryconfig") +#: Set once this process starts as a worker; see :func:`_mark_worker_process`. +_IS_WORKER_PROCESS = False + + +@worker_init.connect +@worker_process_init.connect +def _mark_worker_process(*args, **kwargs): + """Record that this process runs tasks, for :func:`in_worker`. + + ``worker_init`` fires in every worker's main process before its pool + starts: that is where solo, threads, eventlet and gevent run tasks, and + what prefork children fork from. ``worker_process_init`` covers prefork + children however they were started. + """ + global _IS_WORKER_PROCESS + _IS_WORKER_PROCESS = True + + def in_worker() -> bool: - """True anywhere in a Celery worker process, on any thread. + """True anywhere in a Celery worker process, on any thread or greenlet. ``current_worker_task`` alone is not enough: Celery records the executing - task on the thread that runs it, so a thread the task starts sees none and - would take the web-process branch — dispatching to the worker it is running - in and blocking on the result. Celery refuses that ``get()`` ("Never call - result.get() within a task!"), or, where joins are allowed, it waits on a - queue only this busy process serves. + task on the thread (or greenlet) that runs it, so one the task starts sees + none and would take the web-process branch — dispatching to the worker it + is running in and blocking on the result. Celery refuses that ``get()`` + ("Never call result.get() within a task!"), or, where joins are allowed, + it waits on a queue that only this busy process may be able to serve. - ``task_join_will_block`` is process-wide and set for every blocking pool - (prefork, solo, threads) — exactly the condition under which dispatching - and waiting goes wrong. eventlet/gevent leave it unset, so the task's own - thread still counts through ``current_worker_task``. + The worker's own startup (:func:`_mark_worker_process`) answers for every + pool. ``task_join_will_block`` — process-wide, set for every blocking pool + — and the task's own ``current_worker_task`` still count for a process + that runs tasks without having gone through that startup. Returns: bool: Whether this call is running inside a worker process. """ from celery.result import task_join_will_block - return task_join_will_block() or celery.current_worker_task is not None + return ( + _IS_WORKER_PROCESS + or task_join_will_block() + or celery.current_worker_task is not None + ) #: Task-name prefix the package carried before the rename to ``docsgpt``. diff --git a/tests/test_celery.py b/tests/test_celery.py index 5a3e66be..897b2623 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -317,6 +317,28 @@ class TestInWorker: with denied_join_result(): assert self._ask_from_a_new_thread() is True + def test_true_on_any_thread_of_a_non_blocking_pool_worker(self, monkeypatch): + # eventlet/gevent leave the join flag unset and scope the current task + # to one greenlet, so only the worker's own startup can say this + # process is a worker. The lifecycle signal records that for every + # thread and greenlet in it. + import docsgpt.celery_init as celery_init + + monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) + assert self._ask_from_a_new_thread() is False + + celery_init.worker_init.send(sender=None) + + assert self._ask_from_a_new_thread() is True + + def test_prefork_children_record_it_on_their_own_start(self, monkeypatch): + import docsgpt.celery_init as celery_init + + monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) + celery_init.worker_process_init.send(sender=None) + + assert self._ask_from_a_new_thread() is True + def test_true_on_the_task_thread_of_a_non_blocking_pool(self): # eventlet/gevent pools leave the process flag unset; the thread # running the task still knows it is in one. From b32d27c9129c2efedd323b9ab032cf07b040d28a Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 091/130] fix(agents): cite the pages the graph tool read The graph tool records every page read_entity_pages returns in retrieved_docs, but agents only ever collected internal_search's. A turn answered from graph pages therefore emitted no sources: on the multi-hop e2e run, a question answered from two read_entity_pages calls came back with none. _search_tool_docs gathers both tools' documents the way the executor caches them, and both the classic/agentic collector and the research agent's per-step citations use it. Checked against the real executor and a real graph: a graph-only turn now cites the pages it read. --- docsgpt/agents/base.py | 28 +++++++++++++++++++++------- docsgpt/agents/research_agent.py | 14 ++++---------- tests/agents/test_classic_agent.py | 20 ++++++++++++++++++++ tests/agents/test_research_agent.py | 14 ++++++++++++++ 4 files changed, 59 insertions(+), 17 deletions(-) diff --git a/docsgpt/agents/base.py b/docsgpt/agents/base.py index d731a415..030d76fc 100644 --- a/docsgpt/agents/base.py +++ b/docsgpt/agents/base.py @@ -955,16 +955,30 @@ class BaseAgent(ABC): ) self.retrieved_docs = scrubbed - def _collect_internal_sources(self) -> None: - """Merge the cached InternalSearchTool's docs into ``retrieved_docs``, - deduped, preserving any pre-fetched docs so a mixed-exposure agent cites - both pre-fetched and tool-retrieved sources (not just the tool's).""" + def _search_tool_docs(self) -> List[Dict]: + """Documents this run's search tools read: internal search and the graph tool. + + Both record what they surface in ``retrieved_docs``; a page read from + the graph carries the answer as much as a search hit does, so both are + cited. Tools are looked up the way the executor caches them. + """ + from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID from docsgpt.agents.tools.internal_search import INTERNAL_TOOL_ID executor = getattr(self, "tool_executor", None) loaded = getattr(executor, "_loaded_tools", None) or {} - tool = loaded.get(f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}") - if not (tool and getattr(tool, "retrieved_docs", None)): + docs: List[Dict] = [] + for name, tool_id in (("internal_search", INTERNAL_TOOL_ID), ("graph_search", GRAPH_TOOL_ID)): + tool = loaded.get(f"{name}:{tool_id}:{self.user or ''}") + docs.extend(getattr(tool, "retrieved_docs", None) or []) + return docs + + def _collect_internal_sources(self) -> None: + """Merge the search tools' docs into ``retrieved_docs``, deduped, + preserving any pre-fetched docs so a mixed-exposure agent cites both + pre-fetched and tool-retrieved sources (not just the tools').""" + tool_docs = self._search_tool_docs() + if not tool_docs: return def _key(d): @@ -974,7 +988,7 @@ class BaseAgent(ABC): merged = list(self.retrieved_docs or []) seen = {_key(d) for d in merged} - for doc in tool.retrieved_docs: + for doc in tool_docs: k = _key(doc) if k not in seen: seen.add(k) diff --git a/docsgpt/agents/research_agent.py b/docsgpt/agents/research_agent.py index ee92507a..de96cf35 100644 --- a/docsgpt/agents/research_agent.py +++ b/docsgpt/agents/research_agent.py @@ -7,10 +7,7 @@ from typing import Dict, Generator, List, Optional from docsgpt.agents.base import BaseAgent from docsgpt.agents.tool_executor import ToolExecutor from docsgpt.agents.tools.graph_search import add_graph_search_tool -from docsgpt.agents.tools.internal_search import ( - INTERNAL_TOOL_ID, - add_internal_search_tool, -) +from docsgpt.agents.tools.internal_search import add_internal_search_tool from docsgpt.agents.tools.wiki import add_wiki_tool from docsgpt.agents.tools.think import THINK_TOOL_ENTRY, THINK_TOOL_ID from docsgpt.logging import LogContext @@ -622,12 +619,9 @@ class ResearchAgent(BaseAgent): return messages, search_returned_empty def _collect_step_sources(self): - """Collect sources from InternalSearchTool and register with CitationManager.""" - cache_key = f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}" - tool = self.tool_executor._loaded_tools.get(cache_key) - if tool and hasattr(tool, "retrieved_docs"): - for doc in tool.retrieved_docs: - self.citations.add(doc) + """Register the search tools' docs (internal search and graph pages) with CitationManager.""" + for doc in self._search_tool_docs(): + self.citations.add(doc) # ------------------------------------------------------------------ # Phase 3: Synthesis diff --git a/tests/agents/test_classic_agent.py b/tests/agents/test_classic_agent.py index b73e1145..6b10e79d 100644 --- a/tests/agents/test_classic_agent.py +++ b/tests/agents/test_classic_agent.py @@ -313,6 +313,26 @@ class TestClassicAgentSearchExposure: "Tool Doc", ] + def test_collect_internal_sources_includes_graph_pages( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + # Pages the graph tool read carry the answer as much as search hits do, + # so they are cited the same way. + from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID + + retriever_config = {"source": {"active_docs": ["b"]}} + agent = ClassicAgent(retriever_config=retriever_config, **agent_base_params) + search = Mock() + search.retrieved_docs = [{"text": "Found", "title": "Search Doc", "source": "b"}] + graph = Mock() + graph.retrieved_docs = [{"text": "Quill is a store.", "title": "quill.md", "source": "b"}] + user = agent.user or "" + agent.tool_executor._loaded_tools[f"internal_search:{INTERNAL_TOOL_ID}:{user}"] = search + agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{user}"] = graph + + agent._collect_internal_sources() + assert [d["title"] for d in agent.retrieved_docs] == ["Search Doc", "quill.md"] + def test_collect_internal_sources_dedupes( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): diff --git a/tests/agents/test_research_agent.py b/tests/agents/test_research_agent.py index 5b76fe84..cd9aabd9 100644 --- a/tests/agents/test_research_agent.py +++ b/tests/agents/test_research_agent.py @@ -779,6 +779,20 @@ class TestCollectStepSources: assert len(agent.citations.citations) == 2 + def test_collects_pages_the_graph_tool_read( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID + + agent = ResearchAgent(**agent_base_params) + graph = Mock() + graph.retrieved_docs = [{"source": "s3", "title": "quill.md", "text": "Quill"}] + agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{agent.user or ''}"] = graph + + agent._collect_step_sources() + + assert len(agent.citations.citations) == 1 + def test_no_tool_no_error( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): From 5e06470f9ec616c320844caa48949fb4f4145894 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:33:14 +0100 Subject: [PATCH 092/130] test(graphrag): cover the graph reads, graph tool and vector blend without pgvector CI has no pgvector, so every live graph test skips there and the queries this branch added ran in no CI job at all. These pin what holds without a database: the four new read queries bind every value (the entity name comes from an LLM tool call) and map their rows; empty input runs no query; a failed query returns nothing and releases its connection. The graph tool's plumbing, the sources_have_graph gate and the hybrid path's vector ranking get the same. Also drops a redundant chained comparison flagged by code scanning. --- tests/graphrag/test_graph_search_tool.py | 76 ++++++++++++ tests/graphrag/test_retriever_default_path.py | 64 ++++++++++ tests/graphrag/test_store.py | 111 +++++++++++++++++- 3 files changed, 250 insertions(+), 1 deletion(-) diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py index 2f63cf13..559c911f 100644 --- a/tests/graphrag/test_graph_search_tool.py +++ b/tests/graphrag/test_graph_search_tool.py @@ -170,3 +170,79 @@ class TestWiring: assert tools[GRAPH_TOOL_ID]["id"] == GRAPH_TOOL_ID assert tools[GRAPH_TOOL_ID]["config"]["source"] == SOURCE + + +class TestPlumbing: + def test_a_single_source_id_and_empty_entries_are_accepted(self): + tool = GraphSearchTool({"source": {"active_docs": "src-1"}}) + assert tool._sources() == ["src-1"] + + tool = GraphSearchTool({"source": {"active_docs": ["src-1", "", None]}}) + assert tool._sources() == ["src-1"] + + def test_the_store_is_built_once_and_reused(self, monkeypatch): + built = [] + monkeypatch.setattr( + "docsgpt.graphrag.store.GraphStore", lambda: built.append(object()) or built[-1] + ) + tool = GraphSearchTool({"source": SOURCE}) + + assert tool._get_store() is tool._get_store() + assert len(built) == 1 + + def test_an_embedding_failure_makes_entity_search_unavailable(self, monkeypatch): + def _broken_embeddings(): + raise RuntimeError("no model") + + monkeypatch.setattr("docsgpt.vectorstore.base.get_embeddings", _broken_embeddings) + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + tool = GraphSearchTool({"source": SOURCE}) + tool._store = _StubStore() + + assert tool.execute_action("search_entities", query="quill") == "Entity search is unavailable." + + def test_no_matching_entities_is_stated_plainly(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore()) + + assert "No entities found" in tool.execute_action("search_entities", query="quill") + + def test_a_failing_store_is_reported_not_raised(self, monkeypatch): + class _BrokenStore(_StubStore): + def entity_relationships(self, source_id, name, limit=25): + raise RuntimeError("connection lost") + + tool = _tool(monkeypatch, _BrokenStore()) + + assert tool.execute_action("get_relationships", entity="Quill") == "The graph lookup failed." + + +class TestSourcesHaveGraph: + """Whether to offer the tool at all: only when some source has a graph.""" + + def _patch_counts(self, monkeypatch, counts=None, error=None): + class _Store: + def count_nodes_many(self, source_ids): + if error: + raise error + return {s: counts.get(s, 0) for s in source_ids} + + monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", _Store) + + def test_true_when_any_source_has_nodes(self, monkeypatch): + from docsgpt.agents.tools.graph_search import sources_have_graph + + self._patch_counts(monkeypatch, {"b": 12}) + assert sources_have_graph({"active_docs": ["a", "b"]}) is True + + def test_false_when_no_source_has_nodes(self, monkeypatch): + from docsgpt.agents.tools.graph_search import sources_have_graph + + self._patch_counts(monkeypatch, {}) + assert sources_have_graph({"active_docs": "a"}) is False + + def test_false_without_sources_or_when_the_check_fails(self, monkeypatch): + from docsgpt.agents.tools.graph_search import sources_have_graph + + assert sources_have_graph({"active_docs": []}) is False + self._patch_counts(monkeypatch, error=RuntimeError("no pgvector")) + assert sources_have_graph({"active_docs": ["a"]}) is False diff --git a/tests/graphrag/test_retriever_default_path.py b/tests/graphrag/test_retriever_default_path.py index c11a0f0b..d50a2b16 100644 --- a/tests/graphrag/test_retriever_default_path.py +++ b/tests/graphrag/test_retriever_default_path.py @@ -107,3 +107,67 @@ class TestPerSourceOptions: assert retriever.vector_calls == 0 assert VECTOR_ONLY not in _texts(docs) + + +class TestVectorRanking: + """The vector half of the blend, keyed on chunk text since hits carry no id.""" + + class _VectorStore: + def __init__(self, hits=None, error=None): + self.hits = hits or [] + self.error = error + self.searched = None + self.closed = False + + def search(self, question, k, query_vector=None): + self.searched = (question, k, query_vector) + if self.error: + raise self.error + return self.hits + + def close(self): + self.closed = True + + @staticmethod + def _real_retriever(monkeypatch, store): + from types import SimpleNamespace + + retriever = object.__new__(GraphRAGRetriever) + retriever.chunks = 3 + retriever._classic = SimpleNamespace(_get_rephrased_question=lambda: "where does Alder stream?") + monkeypatch.setattr( + "docsgpt.vectorstore.vector_creator.VectorCreator.create_vectorstore", + lambda *args, **kwargs: store, + ) + return retriever + + def test_object_and_dict_hits_become_text_and_metadata(self, monkeypatch): + from types import SimpleNamespace + + store = self._VectorStore( + hits=[ + SimpleNamespace(page_content="Alder streams to Quill.", metadata={"title": "alder.md"}), + {"text": "Quill is compacted every six hours.", "metadata": {"title": "quill.md"}}, + {"page_content": "A passage without metadata."}, + {"metadata": {"title": "no text"}}, + ] + ) + retriever = self._real_retriever(monkeypatch, store) + + ranked = retriever._vector_ranking("src", [0.1, 0.2]) + + assert ranked == [ + ("Alder streams to Quill.", {"title": "alder.md"}), + ("Quill is compacted every six hours.", {"title": "quill.md"}), + ("A passage without metadata.", {}), + ] + # The rephrased question and the shared query vector, with room to fuse. + assert store.searched == ("where does Alder stream?", 20, [0.1, 0.2]) + assert store.closed + + def test_a_failed_search_ranks_nothing_and_still_closes_the_store(self, monkeypatch): + store = self._VectorStore(error=RuntimeError("pgvector down")) + retriever = self._real_retriever(monkeypatch, store) + + assert retriever._vector_ranking("src", [0.1]) == [] + assert store.closed diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 02a315db..27cfddc2 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -335,7 +335,7 @@ class TestGraphStoreLive: store.set_node_degrees(source_id) recomputed = store.get_node_by_normalized(source_id, "solo")["degree"] - assert recomputed == incremental == 0 + assert recomputed == incremental finally: store.delete_by_source(source_id) @@ -615,6 +615,115 @@ class TestGraphStoreParameterization: assert params[-1] == embedding +@pytest.mark.unit +class TestGraphReadQueries: + """The reads behind fact seeding and the agent's graph tool, without a DB. + + The live class covers what these return from real rows; these pin the + contract that holds without one. The entity name reaching + ``entity_relationships``/``entity_pages`` comes from an LLM tool call, so + it must only ever travel as a bound parameter. + """ + + def _store(self, rows=(), fail=False): + store = GraphStore.__new__(GraphStore) + cursor = MagicMock() + cursor.fetchall.return_value = list(rows) + if fail: + cursor.execute.side_effect = RuntimeError("relation does not exist") + conn = MagicMock() + conn.cursor.return_value = cursor + store._get_connection = lambda: conn + return store, cursor, conn + + def test_fact_seeds_bind_every_value_and_read_weight_as_distance(self): + store, cursor, _ = self._store(rows=[("n1", "Quill", "a store", 0.8), ("n2", "Alder", None, None)]) + sid = str(uuid.uuid4()) + embedding = _embedding(0.3) + + rows = store.seed_nodes_from_facts(sid, embedding, fact_limit=0, limit=3) + + sql, params = cursor.execute.call_args.args + assert sid not in sql and str(embedding) not in sql + # Limits are clamped to at least one before binding. + assert params == (embedding, sid, embedding, 1, sid, 3) + assert rows[0] == {"id": "n1", "name": "Quill", "description": "a store", "distance": pytest.approx(0.2)} + assert rows[1]["distance"] == 1.0 + + def test_fact_seeds_need_a_query_vector(self): + store, cursor, _ = self._store() + assert store.seed_nodes_from_facts(str(uuid.uuid4()), []) == [] + cursor.execute.assert_not_called() + + def test_relationships_bind_the_name_as_a_pattern(self): + store, cursor, _ = self._store(rows=[("Alder", "streams_to", "Quill", "audit events")]) + sid = str(uuid.uuid4()) + name = "Quill'; DROP TABLE graph_nodes; --" + + rows = store.entity_relationships(sid, f" {name} ", limit=500) + + sql, params = cursor.execute.call_args.args + assert name not in sql + assert params == (sid, f"%{name}%", f"%{name}%", 500) + assert rows == [ + {"source": "Alder", "type": "streams_to", "target": "Quill", "description": "audit events"} + ] + + def test_pages_prefer_the_entity_itself_over_a_mention(self): + store, cursor, _ = self._store(rows=[({"title": "quill.md"}, "Quill is a store."), (None, None)]) + sid = str(uuid.uuid4()) + + pages = store.entity_pages(sid, "Quill", limit=0) + + sql, params = cursor.execute.call_args.args + assert "Quill" not in sql + # Exact name, name plus a qualifier ("Quill Store"), substring fallback, + # text-opens-with ordering, then the clamped limit. + assert params == ("quill", "quill %", sid, sid, "quill", "quill %", "%Quill%", "Quill%", 1) + assert pages == [{"metadata": {"title": "quill.md"}, "text": "Quill is a store."}, {"metadata": {}, "text": ""}] + + def test_chunk_similarities_are_restricted_to_the_reached_chunks(self): + store, cursor, _ = self._store(rows=[("11", 0.75)]) + sid = str(uuid.uuid4()) + embedding = _embedding(0.9) + + scores = store.chunk_similarities(sid, [11, "12"], embedding) + + sql, params = cursor.execute.call_args.args + assert "= ANY(%s)" in sql and sid not in sql + assert params == (embedding, sid, ["11", "12"]) + assert scores == {"11": 0.75} + + @pytest.mark.parametrize( + "call", + [ + lambda s: s.entity_relationships("sid", " "), + lambda s: s.entity_pages("sid", ""), + lambda s: s.chunk_similarities("sid", [], [0.1]), + lambda s: s.chunk_similarities("sid", ["1"], []), + ], + ) + def test_empty_input_runs_no_query(self, call): + store, cursor, _ = self._store() + assert not call(store) + cursor.execute.assert_not_called() + + @pytest.mark.parametrize( + "call", + [ + lambda s: s.seed_nodes_from_facts("sid", [0.1]), + lambda s: s.entity_relationships("sid", "Quill"), + lambda s: s.entity_pages("sid", "Quill"), + lambda s: s.chunk_similarities("sid", ["1"], [0.1]), + ], + ) + def test_a_failed_query_returns_nothing_and_releases_the_connection(self, call): + store, cursor, conn = self._store(fail=True) + assert not call(store) + cursor.close.assert_called_once() + conn.rollback.assert_called_once() + + @pytest.mark.unit class TestEmbeddingDim: """The graph table dimension is derived from the configured model (FIX 1).""" From 558f803fed56e1f217af444ed409708f327e1817 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 14:53:18 +0100 Subject: [PATCH 093/130] test(worker): send the real startup signals, and import celery_init one way The lifecycle tests imported docsgpt.celery_init as a module next to the file's from-imports, which code scanning flags. Patch the flag by path and send the signals from celery.signals -- the objects a worker actually fires. --- tests/test_celery.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/tests/test_celery.py b/tests/test_celery.py index 897b2623..e240871c 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -322,20 +322,20 @@ class TestInWorker: # to one greenlet, so only the worker's own startup can say this # process is a worker. The lifecycle signal records that for every # thread and greenlet in it. - import docsgpt.celery_init as celery_init + from celery.signals import worker_init - monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) + monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False) assert self._ask_from_a_new_thread() is False - celery_init.worker_init.send(sender=None) + worker_init.send(sender=None) assert self._ask_from_a_new_thread() is True def test_prefork_children_record_it_on_their_own_start(self, monkeypatch): - import docsgpt.celery_init as celery_init + from celery.signals import worker_process_init - monkeypatch.setattr(celery_init, "_IS_WORKER_PROCESS", False) - celery_init.worker_process_init.send(sender=None) + monkeypatch.setattr("docsgpt.celery_init._IS_WORKER_PROCESS", False) + worker_process_init.send(sender=None) assert self._ask_from_a_new_thread() is True From b3d7e63ae32228e6914c5543d86e23e3d260e597 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 20:10:54 +0100 Subject: [PATCH 094/130] fix(graphrag): compose graph SQL through psycopg instead of f-strings Four graph-store queries formatted table and column names into the SQL string: get_chunk_texts and delete_by_source (Bandit B608, alerts #582/#583 on main) and entity_pages/chunk_similarities from this branch (#661/#662, dismissed). The names were validated by _safe_identifier, so none was injectable, but each query was still a string built at runtime. They are now fixed statements composed with psycopg.sql: identifiers go in as sql.Identifier, values stay bound. Identifiers are lower-cased before quoting, because PGVectorStore writes the same names unquoted and Postgres folds those to lower case -- a quoted mixed-case name would address a different table. Bandit reports nothing for docsgpt/graphrag now; the queries return the same rows against a real graph as before. --- docsgpt/graphrag/store.py | 80 ++++++++++++++++++++++--------- tests/graphrag/test_store.py | 25 ++++++++-- tests/retriever/test_graph_rag.py | 15 ++++-- 3 files changed, 90 insertions(+), 30 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 134998ba..ab2b143a 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -18,6 +18,7 @@ import uuid from typing import Any, Dict, List, Optional import psycopg +from psycopg import sql from psycopg.types.json import Jsonb from docsgpt.core.settings import settings @@ -52,6 +53,18 @@ def _safe_identifier(name: str) -> str: return name +def _identifier(name: str) -> sql.Identifier: + """``name`` as a quoted identifier, folded the way Postgres folds it unquoted. + + Composing identifiers through psycopg keeps every query a fixed statement + with bound values: nothing is formatted into the SQL string. The fold + matters because ``PGVectorStore`` writes these names unquoted, which + Postgres lower-cases, while a quoted identifier keeps its case; folding + first keeps both stores addressing the same table. + """ + return sql.Identifier(_safe_identifier(name).lower()) + + def _pgvector_identifiers() -> tuple[str, str, str, str]: """Resolve ``(table, text_col, metadata_col, source_col)`` from ``PGVectorStore``. @@ -1228,18 +1241,25 @@ class GraphStore: cursor = conn.cursor() try: cursor.execute( - f""" - SELECT d.{metadata_col}, d.{text_col}, - (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject - FROM graph_node_chunks gc - JOIN graph_nodes n ON n.id = gc.node_id - JOIN {table} d ON d.id::text = gc.chunk_id - WHERE gc.source_id = %s AND d.{source_col} = %s - AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) - GROUP BY d.{metadata_col}, d.{text_col}, is_subject - ORDER BY is_subject DESC, (d.{text_col} ILIKE %s) DESC - LIMIT %s; - """, + sql.SQL( + """ + SELECT d.{metadata}, d.{text}, + (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject + FROM graph_node_chunks gc + JOIN graph_nodes n ON n.id = gc.node_id + JOIN {table} d ON d.id::text = gc.chunk_id + WHERE gc.source_id = %s AND d.{source} = %s + AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) + GROUP BY d.{metadata}, d.{text}, is_subject + ORDER BY is_subject DESC, (d.{text} ILIKE %s) DESC + LIMIT %s; + """ + ).format( + metadata=_identifier(metadata_col), + text=_identifier(text_col), + table=_identifier(table), + source=_identifier(source_col), + ), ( clean.lower(), f"{clean.lower()} %", source_id, source_id, @@ -1277,11 +1297,17 @@ class GraphStore: cursor = conn.cursor() try: cursor.execute( - f""" - SELECT id::text, 1 - ({vector_col} <=> %s::vector) - FROM {table} - WHERE {source_col} = %s AND id::text = ANY(%s); - """, + sql.SQL( + """ + SELECT id::text, 1 - ({vector} <=> %s::vector) + FROM {table} + WHERE {source} = %s AND id::text = ANY(%s); + """ + ).format( + vector=_identifier(vector_col), + table=_identifier(table), + source=_identifier(source_col), + ), (query_embedding, source_id, [str(c) for c in chunk_ids]), ) return {row[0]: float(row[1]) for row in cursor.fetchall()} @@ -1312,10 +1338,17 @@ class GraphStore: cursor = conn.cursor() try: cursor.execute( - f""" - SELECT id, {text_col}, {metadata_col} FROM {table} - WHERE {source_col} = %s AND id::text = ANY(%s); - """, + sql.SQL( + """ + SELECT id, {text}, {metadata} FROM {table} + WHERE {source} = %s AND id::text = ANY(%s); + """ + ).format( + text=_identifier(text_col), + metadata=_identifier(metadata_col), + table=_identifier(table), + source=_identifier(source_col), + ), (source_id, [str(c) for c in chunk_ids]), ) return { @@ -1501,7 +1534,10 @@ class GraphStore: "graph_ingest_progress", ): cursor.execute( - f"DELETE FROM {table} WHERE source_id = %s;", (source_id,) + sql.SQL("DELETE FROM {} WHERE source_id = %s;").format( + sql.Identifier(table) + ), + (source_id,), ) conn.commit() except Exception as e: diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 27cfddc2..355bed09 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -558,16 +558,23 @@ class TestGraphStoreParameterization: return store, cursor def test_delete_by_source_binds_source_id(self): + from psycopg import sql as pgsql + store, cursor = self._store_with_mock_conn() sid = str(uuid.uuid4()) store.delete_by_source(sid) + tables = [] for call in cursor.execute.call_args_list: - sql = call.args[0] + query = call.args[0] params = call.args[1] if len(call.args) > 1 else None + assert isinstance(query, pgsql.Composable) + sql = query.as_string() assert "WHERE source_id = %s" in sql assert sid not in sql assert params == (sid,) + tables.append(sql.split('"')[1]) + assert tables == ["graph_node_chunks", "graph_edges", "graph_nodes", "graph_ingest_progress"] def test_search_binds_embedding_and_source(self): store, cursor = self._store_with_mock_conn() @@ -636,6 +643,14 @@ class TestGraphReadQueries: store._get_connection = lambda: conn return store, cursor, conn + def test_identifiers_are_quoted_as_postgres_folds_them_unquoted(self): + # PGVectorStore writes these names unquoted, which Postgres folds to + # lower case; quoting keeps case, so the fold happens first or the two + # stores would address different tables. + assert store_module._identifier("Documents").as_string() == '"documents"' + with pytest.raises(ValueError): + store_module._identifier('documents"; DROP TABLE graph_nodes; --') + def test_fact_seeds_bind_every_value_and_read_weight_as_distance(self): store, cursor, _ = self._store(rows=[("n1", "Quill", "a store", 0.8), ("n2", "Alder", None, None)]) sid = str(uuid.uuid4()) @@ -675,7 +690,9 @@ class TestGraphReadQueries: pages = store.entity_pages(sid, "Quill", limit=0) - sql, params = cursor.execute.call_args.args + query, params = cursor.execute.call_args.args + sql = query.as_string() + assert 'JOIN "documents" d' in sql and 'd."source_id" = %s' in sql assert "Quill" not in sql # Exact name, name plus a qualifier ("Quill Store"), substring fallback, # text-opens-with ordering, then the clamped limit. @@ -689,7 +706,9 @@ class TestGraphReadQueries: scores = store.chunk_similarities(sid, [11, "12"], embedding) - sql, params = cursor.execute.call_args.args + query, params = cursor.execute.call_args.args + sql = query.as_string() + assert '1 - ("embedding" <=> %s::vector)' in sql and 'FROM "documents"' in sql assert "= ANY(%s)" in sql and sid not in sql assert params == (embedding, sid, ["11", "12"]) assert scores == {"11": 0.75} diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index e4d1bf46..71c996d5 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -472,11 +472,16 @@ class TestGetChunkTexts: sid = str(uuid.uuid4()) store.get_chunk_texts(sid, ["1", "2"]) - sql, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1] - assert f"FROM {table}" in sql - assert text_col in sql - assert metadata_col in sql - assert f"{source_col} = %s" in sql + from psycopg import sql as pgsql + + query, params = cursor.execute.call_args.args[0], cursor.execute.call_args.args[1] + # Identifiers are composed and quoted by psycopg, never formatted in. + assert isinstance(query, pgsql.Composable) + sql = query.as_string() + assert f'FROM "{table}"' in sql + assert f'"{text_col}"' in sql + assert f'"{metadata_col}"' in sql + assert f'"{source_col}" = %s' in sql assert "id::text = ANY(%s)" in sql assert sid not in sql assert params == (sid, ["1", "2"]) From ecf02d0d60be3eb234ac1db8367a9a1ed444d7f7 Mon Sep 17 00:00:00 2001 From: Alex Date: Sat, 19 Sep 2026 20:27:27 +0100 Subject: [PATCH 095/130] fix(graphrag): take a source's write lock per chunk, and keep zero weights Two builds of one source can overlap: the extraction lease is keyed by the source's updated_at, and enabling a graph updates the source before it dispatches, so a rebuild started while the last build runs gets a new key and a lease of its own. Both builds could then pass a chunk's "done" check before either committed and apply it twice -- doc_freq bumped twice, reproduced with two live writers. A reset could also land in the middle of a chunk. apply_chunk and delete_by_source now take a transaction-scoped advisory lock keyed by the source before touching a row, as the schema bootstrap already does for DDL. A single build's writes were already serial, so it loses nothing; overlapping builds take turns chunk by chunk, and the second sees the first's "done" row and returns (0, 0). apply_chunk also still defaulted with `rel.get("weight") or 1.0`, turning an explicit zero into a full-strength edge -- the conversion 3f774d81 removed from add_edge and the ranker but missed here. Only a missing weight defaults now. --- docsgpt/graphrag/store.py | 28 ++++++++++- tests/graphrag/test_store.py | 94 +++++++++++++++++++++++++++++++++++- 2 files changed, 119 insertions(+), 3 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index ab2b143a..1804d929 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -108,6 +108,22 @@ def _is_connection_lost(exc: BaseException) -> bool: return isinstance(exc, (psycopg.OperationalError, psycopg.InterfaceError)) +def _lock_source(cursor, source_id: str) -> None: + """Serialize graph writes for one source until this transaction ends. + + Writes within one build are already serial, but two builds of the same + source can overlap: a rebuild dispatched while the last one is still + running gets a new idempotency key, so its lease does not stop it. Without + this, both could pass a chunk's "done" check before either commits and + apply it twice. A transaction-scoped advisory lock keyed by the source + makes them take turns chunk by chunk; the lock is released on commit or + rollback, and a hash collision only makes two sources take turns. + """ + cursor.execute( + "SELECT pg_advisory_xact_lock(hashtext(%s));", (f"graphrag:source:{source_id}",) + ) + + def _safe_rollback(conn) -> None: """Roll back, tolerating a connection too broken to roll back.""" try: @@ -672,7 +688,10 @@ class GraphStore: # committed, and the retry then replays this write: doc_freq # would be bumped twice and a second logical edge inserted # (graph_edges has no uniqueness constraint). The progress row - # below is written in this transaction, so a replay sees it. + # below is written in this transaction, so a replay sees it — + # and so does an overlapping build, once the source lock makes + # it wait for this one to commit. + _lock_source(cursor, source_id) cursor.execute( "SELECT status FROM graph_ingest_progress " "WHERE source_id = %s AND chunk_id = %s;", @@ -713,7 +732,9 @@ class GraphStore: dst_id, type=rel.get("type"), description=rel.get("description"), - weight=float(rel.get("weight") or 1.0), + # Only a missing weight defaults: 0 is a real one, + # and the ranker drops non-positive edges. + weight=1.0 if rel.get("weight") is None else float(rel["weight"]), source_chunk_ids=[chunk_id], fact_embedding=rel.get("fact_embedding"), ) @@ -1527,6 +1548,9 @@ class GraphStore: conn = self._get_connection() cursor = conn.cursor() try: + # A reset while a build is still writing must not land in the + # middle of one of its chunks. + _lock_source(cursor, source_id) for table in ( "graph_node_chunks", "graph_edges", diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 355bed09..062db1f2 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -557,6 +557,43 @@ class TestGraphStoreParameterization: store._tables_ensured = True return store, cursor + def test_graph_writes_for_a_source_are_serialized(self): + # A chunk write and a reset each take the source's transaction-scoped + # advisory lock before touching a row, so overlapping builds of one + # source cannot interleave inside a chunk. + store, cursor = self._store_with_mock_conn() + cursor.fetchone.return_value = None + sid = str(uuid.uuid4()) + + store.apply_chunk(sid, "c1", [], [], {}) + first_sql, first_params = cursor.execute.call_args_list[0].args + assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql + assert first_params == (f"graphrag:source:{sid}",) + + cursor.execute.reset_mock() + store.delete_by_source(sid) + first_sql, first_params = cursor.execute.call_args_list[0].args + assert "pg_advisory_xact_lock(hashtext(%s))" in first_sql + assert first_params == (f"graphrag:source:{sid}",) + + def test_apply_chunk_keeps_an_explicit_zero_weight(self, monkeypatch): + store, cursor = self._store_with_mock_conn() + cursor.fetchone.side_effect = [None, ["n1"], ["n2"]] + weights = [] + + def _capture(cursor, source_id, src, dst, type=None, description=None, weight=1.0, **kwargs): + weights.append(weight) + return "e1", True + + monkeypatch.setattr(store, "_add_edge", _capture) + store.apply_chunk( + "sid", "c1", [], + [{"source": "A", "target": "B", "weight": 0}, {"source": "A", "target": "B"}], + {}, + ) + # Zero is a real weight; only a missing one defaults. + assert weights == [0.0, 1.0] + def test_delete_by_source_binds_source_id(self): from psycopg import sql as pgsql @@ -565,7 +602,9 @@ class TestGraphStoreParameterization: store.delete_by_source(sid) tables = [] - for call in cursor.execute.call_args_list: + lock, *deletes = cursor.execute.call_args_list + assert "pg_advisory_xact_lock" in lock.args[0] + for call in deletes: query = call.args[0] params = call.args[1] if len(call.args) > 1 else None assert isinstance(query, pgsql.Composable) @@ -1342,6 +1381,59 @@ class TestApplyChunkIsReplaySafe: finally: store.delete_by_source(source_id) + def test_overlapping_applies_of_one_chunk_write_it_once(self, store, postgresql, monkeypatch): + """Two builds of one source can overlap: a rebuild dispatched while the + last one runs gets a new lease key. Both may reach the same chunk at + once, and the second must wait for the first to commit instead of + passing the done check while the first is still in flight.""" + import threading + import time + + source_id = str(uuid.uuid4()) + entities = [{"name": "Ada", "normalized_name": "ada", "type": "person", "description": "d"}] + relationships = [ + {"source": "Ada", "target": "Engine", "type": "worked_on", "description": "x", "weight": 2.0} + ] + embeddings = {"ada": _embedding(0.1), "engine": _embedding(0.2)} + writers = [GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) for _ in range(2)] + real_upsert = GraphStore._upsert_node + + def _slow_upsert(self, *args, **kwargs): + time.sleep(0.3) # hold the first writer inside its transaction + return real_upsert(self, *args, **kwargs) + + monkeypatch.setattr(GraphStore, "_upsert_node", _slow_upsert) + results = [] + + def _apply(writer): + results.append(writer.apply_chunk(source_id, "c1", entities, relationships, embeddings)) + + try: + threads = [threading.Thread(target=_apply, args=(w,)) for w in writers] + threads[0].start() + time.sleep(0.05) + threads[1].start() + for thread in threads: + thread.join() + + assert sorted(results) == [(0, 0), (1, 1)] + assert store.get_node_by_normalized(source_id, "ada")["doc_freq"] == 1 + assert len(store.get_graph_overview(source_id)["edges"]) == 1 + finally: + for writer in writers: + writer.close() + store.delete_by_source(source_id) + + def test_a_zero_weight_relationship_stays_zero(self, store): + source_id = str(uuid.uuid4()) + relationships = [{"source": "Ada", "target": "Engine", "type": "mentions", "weight": 0.0}] + try: + store.apply_chunk(source_id, "c1", [], relationships, {}) + edges = store.get_graph_overview(source_id)["edges"] + assert [edge["weight"] for edge in edges] == [0.0] + finally: + store.delete_by_source(source_id) + def test_a_different_chunk_still_applies(self, store): """The guard is per chunk, not a blanket 'already saw this source'.""" source_id = str(uuid.uuid4()) From 22ecc0aee3918e560e84dfdd7a2479f002dd2431 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sat, 19 Sep 2026 23:39:58 +0100 Subject: [PATCH 096/130] feat(pat): add personal access token storage and settings A personal_access_tokens table (migration 0032) holds scoped user-level API credentials. Only the SHA-256 of the secret is stored, like device session tokens. Lookups exclude revoked and expired tokens and the tokens of deactivated users. PAT_* settings cover the feature switch, default and maximum lifetime, the operator opt-in for non-expiring tokens and the per-user cap. --- .env-template | 8 + docs/content/Deploying/Settings-Reference.mdx | 30 ++++ .../versions/0032_personal_access_tokens.py | 72 ++++++++ docsgpt/core/settings/auth.py | 24 +++ docsgpt/storage/db/models.py | 41 +++++ .../db/repositories/personal_access_tokens.py | 156 +++++++++++++++++ .../test_personal_access_tokens.py | 163 ++++++++++++++++++ 7 files changed, 494 insertions(+) create mode 100644 docsgpt/alembic/versions/0032_personal_access_tokens.py create mode 100644 docsgpt/storage/db/repositories/personal_access_tokens.py create mode 100644 tests/storage/db/repositories/test_personal_access_tokens.py diff --git a/.env-template b/.env-template index eb218c71..5e868f5e 100644 --- a/.env-template +++ b/.env-template @@ -101,3 +101,11 @@ MICROSOFT_AUTHORITY=https://{tenantId}.ciamlogin.com/{tenantId} # pair with OIDC_USER_ID_CLAIM=email so SCIM userName matches the OIDC user id) # SCIM_ENABLED=false # SCIM_TOKEN= + +# Personal access tokens (scoped API tokens for CLI and CI/CD; Settings → Access Tokens). +# Available with AUTH_TYPE=oidc or unset. +# PAT_ENABLED=true +# PAT_DEFAULT_LIFETIME_DAYS=90 +# PAT_MAX_LIFETIME_DAYS=365 +# PAT_ALLOW_NON_EXPIRING=false +# PAT_MAX_PER_USER=25 diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 32181adb..9266a90e 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -132,6 +132,36 @@ Type `str`, default unset. Bearer token for IdP SCIM clients (required when SCIM is enabled). +### `PAT_ENABLED` + +Type `bool`, default `true`. + +Allow users to create personal access tokens. Tokens are only issued under AUTH_TYPE=oidc or unset (None); simple_jwt and session_jwt have no stable user identity to bind a token to. + +### `PAT_DEFAULT_LIFETIME_DAYS` + +Type `int`, default `90`, must be `> 0`. + +Lifetime of a personal access token created without an explicit expiry. + +### `PAT_MAX_LIFETIME_DAYS` + +Type `int`, default `365`, must be `> 0`. + +Longest lifetime a user may request for a personal access token. + +### `PAT_ALLOW_NON_EXPIRING` + +Type `bool`, default `false`. + +Let users create personal access tokens that never expire. Off by default. + +### `PAT_MAX_PER_USER` + +Type `int`, default `25`, must be `> 0`. + +Maximum number of live personal access tokens per user. + ## LLM providers diff --git a/docsgpt/alembic/versions/0032_personal_access_tokens.py b/docsgpt/alembic/versions/0032_personal_access_tokens.py new file mode 100644 index 00000000..7d5090bf --- /dev/null +++ b/docsgpt/alembic/versions/0032_personal_access_tokens.py @@ -0,0 +1,72 @@ +"""0032 personal access tokens — scoped, user-level API credentials. + +A personal access token (PAT) authenticates its owner against the management +API for CLI and CI/CD use. Only the SHA-256 of the secret is stored, mirroring +``devices.token_hash``: the plaintext is shown once at creation and a database +leak cannot reconstruct it. ``token_prefix`` keeps the first characters so a +user can tell their tokens apart in the UI. + +``scopes`` is the server-side grant list (never read from the credential +itself). ``resource_filter`` optionally narrows a resource family to specific +ids, e.g. ``{"agents": [""]}``; an absent family is unrestricted within +the token's scopes. ``expires_at`` is NULL only when the operator allows +non-expiring tokens. + +``user_id`` is the auth ``sub``; no FK or trigger, mirroring ``devices`` and +``user_roles`` so a token row never blocks user deletion. + +Revision ID: 0032_personal_access_tokens +Revises: 0031_token_usage_cache_tokens +""" + +from typing import Sequence, Union + +from alembic import op + + +revision: str = "0032_personal_access_tokens" +down_revision: Union[str, None] = "0031_token_usage_cache_tokens" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.execute( + """ + CREATE TABLE IF NOT EXISTS personal_access_tokens ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id TEXT NOT NULL, + name TEXT NOT NULL, + token_hash TEXT NOT NULL, + token_prefix TEXT NOT NULL, + scopes TEXT[] NOT NULL DEFAULT '{}', + resource_filter JSONB NOT NULL DEFAULT '{}'::jsonb, + status TEXT NOT NULL DEFAULT 'active' + CHECK (status IN ('active', 'revoked')), + expires_at TIMESTAMPTZ, + last_used_at TIMESTAMPTZ, + last_used_ip TEXT, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + revoked_at TIMESTAMPTZ, + revoke_reason TEXT + ); + """ + ) + # Looked up on every PAT-authenticated request. + op.execute( + "CREATE UNIQUE INDEX IF NOT EXISTS personal_access_tokens_hash_uidx " + "ON personal_access_tokens(token_hash);" + ) + # Names are unique among a user's live tokens; a revoked name can be reused. + op.execute( + "CREATE UNIQUE INDEX IF NOT EXISTS personal_access_tokens_user_name_uidx " + "ON personal_access_tokens(user_id, name) WHERE status = 'active';" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS personal_access_tokens_user_idx " + "ON personal_access_tokens(user_id, created_at DESC);" + ) + + +def downgrade() -> None: + op.execute("DROP TABLE IF EXISTS personal_access_tokens;") diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py index 85c57b84..7471dc5a 100644 --- a/docsgpt/core/settings/auth.py +++ b/docsgpt/core/settings/auth.py @@ -85,6 +85,28 @@ class AuthSettings(SettingsGroup): default=None, description="Bearer token for IdP SCIM clients (required when SCIM is enabled)." ) + # Personal access tokens: scoped user-level API credentials for CLI and CI/CD use. + PAT_ENABLED: bool = Field( + default=True, + description=( + "Allow users to create personal access tokens. Tokens are only issued under AUTH_TYPE=oidc or " + "unset (None); simple_jwt and session_jwt have no stable user identity to bind a token to." + ), + ) + PAT_DEFAULT_LIFETIME_DAYS: int = Field( + default=90, gt=0, description="Lifetime of a personal access token created without an explicit expiry." + ) + PAT_MAX_LIFETIME_DAYS: int = Field( + default=365, gt=0, description="Longest lifetime a user may request for a personal access token." + ) + PAT_ALLOW_NON_EXPIRING: bool = Field( + default=False, + description="Let users create personal access tokens that never expire. Off by default.", + ) + PAT_MAX_PER_USER: int = Field( + default=25, gt=0, description="Maximum number of live personal access tokens per user." + ) + @field_validator("AUTH_TYPE", mode="before") @classmethod def _normalize_auth_type(cls, v): @@ -99,4 +121,6 @@ class AuthSettings(SettingsGroup): raise ValueError(f"AUTH_TYPE=oidc requires settings: {', '.join(missing)}") if self.SCIM_ENABLED and not self.SCIM_TOKEN: raise ValueError("SCIM_ENABLED requires settings: SCIM_TOKEN") + if self.PAT_DEFAULT_LIFETIME_DAYS > self.PAT_MAX_LIFETIME_DAYS: + raise ValueError("PAT_DEFAULT_LIFETIME_DAYS must not exceed PAT_MAX_LIFETIME_DAYS") return self diff --git a/docsgpt/storage/db/models.py b/docsgpt/storage/db/models.py index 143c3b07..65527c76 100644 --- a/docsgpt/storage/db/models.py +++ b/docsgpt/storage/db/models.py @@ -1079,3 +1079,44 @@ device_auto_approve_patterns_table = Table( Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()), UniqueConstraint("device_id", "user_id", "pattern", name="device_auto_approve_uidx"), ) + +# --- Personal access tokens (migration 0032) -------------------------------- +# Scoped user-level API credentials. Only the SHA-256 of the secret is stored. + +personal_access_tokens_table = Table( + "personal_access_tokens", + metadata, + Column("id", UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid()), + Column("user_id", Text, nullable=False), + Column("name", Text, nullable=False), + Column("token_hash", Text, nullable=False), + Column("token_prefix", Text, nullable=False), + Column("scopes", ARRAY(Text), nullable=False, server_default="{}"), + Column("resource_filter", JSONB, nullable=False, server_default=text("'{}'::jsonb")), + Column("status", Text, nullable=False, server_default="active"), + Column("expires_at", DateTime(timezone=True)), + Column("last_used_at", DateTime(timezone=True)), + Column("last_used_ip", Text), + Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()), + Column("revoked_at", DateTime(timezone=True)), + Column("revoke_reason", Text), + CheckConstraint("status IN ('active', 'revoked')", name="personal_access_tokens_status_check"), +) + +Index( + "personal_access_tokens_hash_uidx", + personal_access_tokens_table.c.token_hash, + unique=True, +) +Index( + "personal_access_tokens_user_name_uidx", + personal_access_tokens_table.c.user_id, + personal_access_tokens_table.c.name, + unique=True, + postgresql_where=personal_access_tokens_table.c.status == "active", +) +Index( + "personal_access_tokens_user_idx", + personal_access_tokens_table.c.user_id, + personal_access_tokens_table.c.created_at.desc(), +) diff --git a/docsgpt/storage/db/repositories/personal_access_tokens.py b/docsgpt/storage/db/repositories/personal_access_tokens.py new file mode 100644 index 00000000..56fb5caa --- /dev/null +++ b/docsgpt/storage/db/repositories/personal_access_tokens.py @@ -0,0 +1,156 @@ +"""Repository for the ``personal_access_tokens`` table.""" + +from __future__ import annotations + +import json +from datetime import datetime +from typing import Optional + +from sqlalchemy import Connection, text + +from docsgpt.storage.db.base_repository import row_to_dict + + +# token_hash never leaves the repository except through find_active_by_hash. +_PUBLIC_COLUMNS = ( + "id, user_id, name, token_prefix, scopes, resource_filter, status, " + "expires_at, last_used_at, last_used_ip, created_at, revoked_at, revoke_reason" +) + + +class PersonalAccessTokensRepository: + """CRUD for personal access tokens. Callers hash the secret; only the hash is stored.""" + + def __init__(self, conn: Connection) -> None: + self._conn = conn + + def create( + self, + user_id: str, + name: str, + *, + token_hash: str, + token_prefix: str, + scopes: list[str], + resource_filter: Optional[dict] = None, + expires_at: Optional[datetime] = None, + ) -> dict: + row = self._conn.execute( + text( + f""" + INSERT INTO personal_access_tokens ( + user_id, name, token_hash, token_prefix, scopes, + resource_filter, expires_at + ) VALUES ( + :user_id, :name, :token_hash, :token_prefix, :scopes, + CAST(:resource_filter AS jsonb), :expires_at + ) RETURNING {_PUBLIC_COLUMNS} + """ + ), + { + "user_id": user_id, + "name": name, + "token_hash": token_hash, + "token_prefix": token_prefix, + "scopes": list(scopes), + "resource_filter": json.dumps(resource_filter or {}), + "expires_at": expires_at, + }, + ).fetchone() + return row_to_dict(row) + + def get(self, token_id: str, user_id: Optional[str] = None) -> Optional[dict]: + sql = f"SELECT {_PUBLIC_COLUMNS} FROM personal_access_tokens WHERE id = CAST(:id AS uuid)" + params: dict = {"id": token_id} + if user_id is not None: + sql += " AND user_id = :user_id" + params["user_id"] = user_id + row = self._conn.execute(text(sql), params).fetchone() + return row_to_dict(row) if row is not None else None + + def list_for_user(self, user_id: str, *, include_revoked: bool = False) -> list[dict]: + sql = f"SELECT {_PUBLIC_COLUMNS} FROM personal_access_tokens WHERE user_id = :user_id" + if not include_revoked: + sql += " AND status = 'active'" + sql += " ORDER BY created_at DESC" + result = self._conn.execute(text(sql), {"user_id": user_id}) + return [row_to_dict(r) for r in result.fetchall()] + + def count_active(self, user_id: str) -> int: + """Live tokens only: revoked and expired rows don't count against the per-user cap.""" + return self._conn.execute( + text( + "SELECT count(*) FROM personal_access_tokens " + "WHERE user_id = :user_id AND status = 'active' " + "AND (expires_at IS NULL OR expires_at > now())" + ), + {"user_id": user_id}, + ).scalar_one() + + def name_in_use(self, user_id: str, name: str) -> bool: + return ( + self._conn.execute( + text( + "SELECT 1 FROM personal_access_tokens " + "WHERE user_id = :user_id AND name = :name AND status = 'active' LIMIT 1" + ), + {"user_id": user_id, "name": name}, + ).fetchone() + is not None + ) + + def find_active_by_hash(self, token_hash: str) -> Optional[dict]: + """Resolve the credential on each request. + + Revoked and expired tokens never match, and neither do the tokens of a + deactivated user (admin or SCIM), so deactivation needs no token sweep + and reactivation restores them. + """ + row = self._conn.execute( + text( + f"SELECT {_PUBLIC_COLUMNS} FROM personal_access_tokens pat " + "WHERE token_hash = :token_hash AND status = 'active' " + "AND (expires_at IS NULL OR expires_at > now()) " + "AND NOT EXISTS (SELECT 1 FROM users u " + "WHERE u.user_id = pat.user_id AND u.active = false) " + "LIMIT 1" + ), + {"token_hash": token_hash}, + ).fetchone() + return row_to_dict(row) if row is not None else None + + def touch_last_used(self, token_id: str, ip: Optional[str], *, min_interval_seconds: int = 60) -> None: + """Record use, at most once per ``min_interval_seconds`` so hot tokens don't write per request.""" + self._conn.execute( + text( + "UPDATE personal_access_tokens " + "SET last_used_at = now(), last_used_ip = :ip " + "WHERE id = CAST(:id AS uuid) AND (last_used_at IS NULL " + "OR last_used_at <= now() - make_interval(secs => :min_interval))" + ), + {"id": token_id, "ip": ip, "min_interval": min_interval_seconds}, + ) + + def revoke(self, token_id: str, user_id: Optional[str] = None, *, reason: str = "user_revoked") -> bool: + """Revoke one token. ``user_id=None`` is the admin path (any owner).""" + sql = ( + "UPDATE personal_access_tokens " + "SET status = 'revoked', revoked_at = now(), revoke_reason = :reason " + "WHERE id = CAST(:id AS uuid) AND status = 'active'" + ) + params: dict = {"id": token_id, "reason": reason} + if user_id is not None: + sql += " AND user_id = :user_id" + params["user_id"] = user_id + return self._conn.execute(text(sql), params).rowcount > 0 + + def revoke_all_for_user(self, user_id: str, *, reason: str = "admin_revoked") -> int: + result = self._conn.execute( + text( + "UPDATE personal_access_tokens " + "SET status = 'revoked', revoked_at = now(), revoke_reason = :reason " + "WHERE user_id = :user_id AND status = 'active'" + ), + {"user_id": user_id, "reason": reason}, + ) + return result.rowcount diff --git a/tests/storage/db/repositories/test_personal_access_tokens.py b/tests/storage/db/repositories/test_personal_access_tokens.py new file mode 100644 index 00000000..0102e625 --- /dev/null +++ b/tests/storage/db/repositories/test_personal_access_tokens.py @@ -0,0 +1,163 @@ +"""Tests for PersonalAccessTokensRepository against a real Postgres.""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone + +import pytest +from sqlalchemy import text + +from docsgpt.storage.db.repositories.personal_access_tokens import ( + PersonalAccessTokensRepository, +) + + +def _create(repo, user_id="u1", name="ci", token_hash="h1", **kwargs): + kwargs.setdefault("scopes", ["agents:read"]) + return repo.create( + user_id, name, token_hash=token_hash, token_prefix="dgpt_pat_abc123", **kwargs + ) + + +class TestCreateAndRead: + def test_create_returns_public_columns_only(self, pg_conn): + row = _create( + PersonalAccessTokensRepository(pg_conn), + resource_filter={"agents": ["00000000-0000-0000-0000-000000000001"]}, + ) + assert "token_hash" not in row + assert row["status"] == "active" + assert row["scopes"] == ["agents:read"] + assert row["resource_filter"] == {"agents": ["00000000-0000-0000-0000-000000000001"]} + assert row["expires_at"] is None + + def test_get_is_owner_scoped(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + row = _create(repo) + assert repo.get(str(row["id"]), "u1")["name"] == "ci" + assert repo.get(str(row["id"]), "someone-else") is None + assert repo.get(str(row["id"]))["user_id"] == "u1" + + def test_list_hides_revoked_by_default(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + kept = _create(repo, name="kept", token_hash="h1") + gone = _create(repo, name="gone", token_hash="h2") + repo.revoke(str(gone["id"]), "u1") + assert [r["id"] for r in repo.list_for_user("u1")] == [kept["id"]] + assert len(repo.list_for_user("u1", include_revoked=True)) == 2 + assert repo.list_for_user("u2") == [] + + +class TestUniqueness: + def test_duplicate_active_name_rejected(self, pg_conn): + from sqlalchemy.exc import IntegrityError + + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo, token_hash="h1") + with pytest.raises(IntegrityError), pg_conn.begin_nested(): + _create(repo, token_hash="h2") + + def test_revoked_name_can_be_reused(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + first = _create(repo, token_hash="h1") + repo.revoke(str(first["id"]), "u1") + assert not repo.name_in_use("u1", "ci") + assert _create(repo, token_hash="h2")["name"] == "ci" + + def test_same_name_for_other_user_is_fine(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo, user_id="u1", token_hash="h1") + _create(repo, user_id="u2", token_hash="h2") + assert repo.name_in_use("u1", "ci") and repo.name_in_use("u2", "ci") + + +class TestFindActiveByHash: + def test_finds_live_token(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo) + found = repo.find_active_by_hash("h1") + assert found["user_id"] == "u1" + assert "token_hash" not in found + + def test_unknown_hash(self, pg_conn): + assert PersonalAccessTokensRepository(pg_conn).find_active_by_hash("nope") is None + + def test_revoked_token_never_matches(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + row = _create(repo) + repo.revoke(str(row["id"]), "u1") + assert repo.find_active_by_hash("h1") is None + + def test_expired_token_never_matches(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo, expires_at=datetime.now(timezone.utc) - timedelta(seconds=1)) + assert repo.find_active_by_hash("h1") is None + + def test_future_expiry_matches(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo, expires_at=datetime.now(timezone.utc) + timedelta(days=1)) + assert repo.find_active_by_hash("h1") is not None + + def test_deactivated_user_token_never_matches(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo) + pg_conn.execute( + text("INSERT INTO users (user_id, active) VALUES ('u1', false)") + ) + assert repo.find_active_by_hash("h1") is None + pg_conn.execute(text("UPDATE users SET active = true WHERE user_id = 'u1'")) + assert repo.find_active_by_hash("h1") is not None + + +class TestRevoke: + def test_revoke_is_owner_scoped(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + row = _create(repo) + assert repo.revoke(str(row["id"]), "someone-else") is False + assert repo.revoke(str(row["id"]), "u1") is True + assert repo.revoke(str(row["id"]), "u1") is False + stored = repo.get(str(row["id"])) + assert stored["status"] == "revoked" + assert stored["revoke_reason"] == "user_revoked" + assert stored["revoked_at"] is not None + + def test_admin_revoke_needs_no_owner(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + row = _create(repo) + assert repo.revoke(str(row["id"]), reason="admin_revoked") is True + assert repo.get(str(row["id"]))["revoke_reason"] == "admin_revoked" + + def test_revoke_all_for_user(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo, name="a", token_hash="h1") + _create(repo, name="b", token_hash="h2") + _create(repo, user_id="u2", token_hash="h3") + assert repo.revoke_all_for_user("u1") == 2 + assert repo.list_for_user("u1") == [] + assert len(repo.list_for_user("u2")) == 1 + + +class TestCountAndUsage: + def test_count_active_ignores_revoked_and_expired(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo, name="live", token_hash="h1") + _create( + repo, + name="expired", + token_hash="h2", + expires_at=datetime.now(timezone.utc) - timedelta(days=1), + ) + revoked = _create(repo, name="revoked", token_hash="h3") + repo.revoke(str(revoked["id"]), "u1") + assert repo.count_active("u1") == 1 + + def test_touch_last_used_is_throttled(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + row = _create(repo) + repo.touch_last_used(str(row["id"]), "10.0.0.1") + first = repo.get(str(row["id"])) + assert first["last_used_ip"] == "10.0.0.1" + repo.touch_last_used(str(row["id"]), "10.0.0.2") + assert repo.get(str(row["id"]))["last_used_ip"] == "10.0.0.1" + repo.touch_last_used(str(row["id"]), "10.0.0.3", min_interval_seconds=0) + assert repo.get(str(row["id"]))["last_used_ip"] == "10.0.0.3" From 174130607a099c7bc9f704d8db16f76df9666e4b Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sat, 19 Sep 2026 23:39:58 +0100 Subject: [PATCH 097/130] feat(pat): authenticate personal access tokens handle_auth resolves a dgpt_pat_ bearer against the database instead of decoding it as a JWT, for both the Flask and the ASGI routes. Scopes and the resource filter always come from the token row, and the claims that mark a PAT are stripped from decoded JWTs so a session token cannot pose as one. ASGI routes reject tokens unless they name the scope that admits them, and /api/user/me reports what the calling token may do. --- docsgpt/api/asgi_auth.py | 20 ++- docsgpt/api/async_sse.py | 2 +- docsgpt/api/pat/__init__.py | 0 docsgpt/api/pat/tokens.py | 234 +++++++++++++++++++++++++++++ docsgpt/api/user/me/routes.py | 11 ++ docsgpt/auth.py | 23 +++ tests/api/test_asgi_auth.py | 42 ++++++ tests/api/test_pat_tokens.py | 275 ++++++++++++++++++++++++++++++++++ 8 files changed, 604 insertions(+), 3 deletions(-) create mode 100644 docsgpt/api/pat/__init__.py create mode 100644 docsgpt/api/pat/tokens.py create mode 100644 tests/api/test_pat_tokens.py diff --git a/docsgpt/api/asgi_auth.py b/docsgpt/api/asgi_auth.py index 93f08a23..aa2ac5bd 100644 --- a/docsgpt/api/asgi_auth.py +++ b/docsgpt/api/asgi_auth.py @@ -16,27 +16,43 @@ from starlette.requests import Request from starlette.responses import JSONResponse from docsgpt.api.oidc.denylist import is_denied as oidc_session_denied +from docsgpt.api.pat.tokens import is_pat from docsgpt.auth import handle_auth from docsgpt.core import log_context from docsgpt.core.settings import settings -async def authenticate(request: Request) -> Tuple[Optional[dict], Optional[JSONResponse]]: +async def authenticate( + request: Request, *, pat_scope: Optional[str] = None +) -> Tuple[Optional[dict], Optional[JSONResponse]]: """Decode the caller's JWT the way Flask's ``authenticate_request`` does. Args: request: The incoming Starlette request. + pat_scope: Scope a personal access token needs for this route. Left + unset, the route rejects PATs outright (deny by default, matching + the Flask rule table in ``docsgpt/api/pat/rules.py``). Returns: tuple: ``(claims, None)`` for an authenticated caller, ``(None, None)`` when no token was sent (the route decides whether that is allowed), or ``(None, response)`` carrying the 401 to return. """ - decoded = handle_auth(request) + # A personal access token resolves against Postgres; keep that sync read off the event loop. + decoded = await anyio.to_thread.run_sync(handle_auth, request) if not decoded: return None, None if "error" in decoded: return None, JSONResponse(decoded, status_code=401) + if is_pat(decoded): + # A PAT lookup already excludes revoked tokens and deactivated users, + # so the session denylist below does not apply to it. + if pat_scope is None or pat_scope not in (decoded.get("scopes") or []): + return None, JSONResponse( + {"success": False, "message": "Token lacks the required scope", "error": "insufficient_scope"}, + status_code=403, + ) + return decoded, None # The denylist is a sync Redis read; keep it off the event loop. if settings.AUTH_TYPE == "oidc" and await anyio.to_thread.run_sync(oidc_session_denied, decoded): return None, JSONResponse( diff --git a/docsgpt/api/async_sse.py b/docsgpt/api/async_sse.py index a5064e2c..1bb3af1d 100644 --- a/docsgpt/api/async_sse.py +++ b/docsgpt/api/async_sse.py @@ -94,7 +94,7 @@ async def stream_message_events(request: Request) -> Response: """ # Same JWT decoder and OIDC revocation check as the Flask routes. With # AUTH_TYPE unset the caller resolves to ``{"sub": "local"}``. - decoded, error = await authenticate(request) + decoded, error = await authenticate(request, pat_scope="chat:run") if error is not None: return error user_id = decoded.get("sub") if isinstance(decoded, dict) else None diff --git a/docsgpt/api/pat/__init__.py b/docsgpt/api/pat/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/docsgpt/api/pat/tokens.py b/docsgpt/api/pat/tokens.py new file mode 100644 index 00000000..fa4cb9d0 --- /dev/null +++ b/docsgpt/api/pat/tokens.py @@ -0,0 +1,234 @@ +"""Personal access tokens: format, scope catalog and the per-request verifier. + +A PAT is ``dgpt_pat_`` + 32 random bytes (urlsafe). Only its SHA-256 is stored +(``personal_access_tokens.token_hash``), the same shape as device session +tokens. Scopes and the resource filter are always read from the database row; +nothing about a token's authority is encoded in the credential itself. +""" + +from __future__ import annotations + +import hashlib +import logging +import secrets +import uuid +from datetime import datetime, timedelta, timezone +from typing import Any, Optional + +from docsgpt.core.settings import settings +from docsgpt.storage.db.repositories.personal_access_tokens import ( + PersonalAccessTokensRepository, +) +from docsgpt.storage.db.session import db_readonly, db_session + +logger = logging.getLogger(__name__) + +TOKEN_PREFIX = "dgpt_pat_" +# Characters of the secret kept in ``token_prefix`` so users can tell tokens apart. +_DISPLAY_CHARS = 6 + +AUTH_METHOD_PAT = "pat" + +#: Every grantable scope with the description shown in the UI and docs. +SCOPES: dict[str, str] = { + "agents:read": "View agents, folders, guardrail events and export agent definitions", + "agents:write": "Create, update, delete, share and import (apply) agents and folders", + "agents:keys": "Regenerate agent API keys and read incoming webhook URLs", + "sources:read": "View sources, their files, chunks and ingestion task status", + "sources:write": "Upload, ingest, sync, edit and delete sources and chunks", + "prompts:read": "View prompts", + "prompts:write": "Create, update and delete prompts", + "tools:read": "View configured tools", + "tools:write": "Create, update and delete tools and MCP servers", + "models:read": "View available and custom models", + "models:write": "Create, update, test and delete custom models", + "workflows:read": "View workflows", + "workflows:write": "Create, update and delete workflows", + "schedules:read": "View agent schedules and their runs", + "schedules:write": "Create, update, run and delete agent schedules", + "conversations:read": "View conversations and messages", + "conversations:write": "Rename, delete and give feedback on conversations", + "analytics:read": "View usage analytics and logs", + "teams:read": "View teams, members and resource shares", + "chat:run": "Ask agents and search sources (answer, stream, search); used for benchmarking", +} + +#: Resource families whose tokens can be narrowed to specific ids. +FILTERABLE_FAMILIES = ("agents", "sources", "prompts", "tools", "workflows") +_MAX_FILTER_IDS = 200 + + +def auth_type_supports_pats() -> bool: + """PATs bind to a stable user id, which simple_jwt/session_jwt don't have.""" + return bool(settings.PAT_ENABLED) and settings.AUTH_TYPE in (None, "oidc") + + +def generate_token() -> tuple[str, str, str]: + """Mint a token. Returns ``(plaintext, sha256_hex, display_prefix)``.""" + secret = secrets.token_urlsafe(32) + token = TOKEN_PREFIX + secret + return token, hash_token(token), TOKEN_PREFIX + secret[:_DISPLAY_CHARS] + + +def hash_token(token: str) -> str: + return hashlib.sha256(token.encode("utf-8")).hexdigest() + + +def looks_like_pat(value: Optional[str]) -> bool: + return bool(value) and value.startswith(TOKEN_PREFIX) + + +def redact(value: Optional[str]) -> str: + """Log-safe form of a credential: the display prefix only.""" + if not value: + return "" + if looks_like_pat(value): + return value[: len(TOKEN_PREFIX) + _DISPLAY_CHARS] + "…" + return value[:4] + "…" + + +def expand_scopes(scopes) -> set[str]: + """Granted scopes plus what they imply (``x:write`` includes ``x:read``).""" + granted = set(scopes or []) + for scope in list(granted): + family, _, action = scope.partition(":") + if action == "write" and f"{family}:read" in SCOPES: + granted.add(f"{family}:read") + return granted + + +def normalize_scopes(raw: Any) -> list[str]: + """Validate a requested scope list. Raises ``ValueError`` with a user-facing message.""" + if not isinstance(raw, list) or not raw: + raise ValueError("scopes must be a non-empty list") + unknown = sorted({s for s in raw if not isinstance(s, str) or s not in SCOPES}, key=str) + if unknown: + raise ValueError(f"Unknown scopes: {', '.join(map(str, unknown))}") + return sorted(set(raw)) + + +def normalize_resource_filter(raw: Any, scopes: list[str]) -> dict[str, list[str]]: + """Validate ``{"": ["", ...]}``. Raises ``ValueError`` with a user-facing message. + + A family may only be restricted when the token holds a scope in it; + otherwise the restriction would be dead weight that reads as protection. + """ + if raw in (None, {}): + return {} + if not isinstance(raw, dict): + raise ValueError("resource_filter must be an object") + families = {s.partition(":")[0] for s in scopes} + # chat:run acts on agents and sources, so both may be restricted alongside it. + if "chat" in families: + families.update({"agents", "sources"}) + out: dict[str, list[str]] = {} + for family, ids in raw.items(): + if family not in FILTERABLE_FAMILIES: + raise ValueError( + f"resource_filter supports only: {', '.join(FILTERABLE_FAMILIES)}" + ) + if family not in families: + raise ValueError(f"resource_filter.{family} needs a {family} scope on the token") + if not isinstance(ids, list) or not ids: + raise ValueError(f"resource_filter.{family} must be a non-empty list of ids") + if len(ids) > _MAX_FILTER_IDS: + raise ValueError(f"resource_filter.{family} allows at most {_MAX_FILTER_IDS} ids") + normalized = [] + for value in ids: + try: + normalized.append(str(uuid.UUID(str(value)))) + except (ValueError, AttributeError, TypeError): + raise ValueError(f"resource_filter.{family} contains an invalid id: {value!r}") + out[family] = sorted(set(normalized)) + return out + + +def resolve_expiry(expires_in_days: Any) -> Optional[datetime]: + """Map the requested lifetime to ``expires_at``. Raises ``ValueError`` with a user-facing message. + + ``None`` means "use the default"; ``0`` asks for a non-expiring token, + which only an operator setting can allow. + """ + if expires_in_days is None: + days = settings.PAT_DEFAULT_LIFETIME_DAYS + elif isinstance(expires_in_days, bool) or not isinstance(expires_in_days, int): + raise ValueError("expires_in_days must be an integer") + elif expires_in_days == 0: + if not settings.PAT_ALLOW_NON_EXPIRING: + raise ValueError("Non-expiring tokens are disabled on this server") + return None + elif expires_in_days < 0: + raise ValueError("expires_in_days must be positive") + else: + days = expires_in_days + if days > settings.PAT_MAX_LIFETIME_DAYS: + raise ValueError(f"expires_in_days must not exceed {settings.PAT_MAX_LIFETIME_DAYS}") + return datetime.now(timezone.utc) + timedelta(days=days) + + +def _client_ip(request) -> Optional[str]: + # Flask exposes remote_addr; Starlette exposes client.host. + ip = getattr(request, "remote_addr", None) + if ip: + return ip + client = getattr(request, "client", None) + return getattr(client, "host", None) + + +_TOUCH_INTERVAL_SECONDS = 60 + + +def _usage_is_stale(last_used_at: Any) -> bool: + """True when ``last_used_at`` is old enough to be worth a write transaction.""" + if not last_used_at: + return True + try: + seen = last_used_at if isinstance(last_used_at, datetime) else datetime.fromisoformat(str(last_used_at)) + except ValueError: + return True + if seen.tzinfo is None: + seen = seen.replace(tzinfo=timezone.utc) + return (datetime.now(timezone.utc) - seen).total_seconds() >= _TOUCH_INTERVAL_SECONDS + + +_INVALID = {"message": "Authentication error: invalid token", "error": "invalid_token"} + + +def authenticate_pat(token: str, request) -> dict: + """Resolve a PAT into the claims dict the rest of the app reads. + + Fails closed: an unknown, revoked or expired token, a disabled feature or + a database error all yield the same ``invalid_token`` error. + """ + if not auth_type_supports_pats(): + return dict(_INVALID) + try: + with db_readonly() as conn: + row = PersonalAccessTokensRepository(conn).find_active_by_hash(hash_token(token)) + except Exception: + logger.error("PAT lookup failed for %s", redact(token), exc_info=True) + return dict(_INVALID) + if not row: + logger.warning("Rejected personal access token %s", redact(token)) + return dict(_INVALID) + if _usage_is_stale(row.get("last_used_at")): + try: + with db_session() as conn: + PersonalAccessTokensRepository(conn).touch_last_used( + str(row["id"]), _client_ip(request), min_interval_seconds=_TOUCH_INTERVAL_SECONDS + ) + except Exception: + # Usage telemetry must never fail a request. + logger.debug("PAT last-used update failed", exc_info=True) + return { + "sub": row["user_id"], + "auth_method": AUTH_METHOD_PAT, + "pat_id": str(row["id"]), + "pat_name": row["name"], + "scopes": sorted(expand_scopes(row.get("scopes"))), + "resource_filter": row.get("resource_filter") or {}, + } + + +def is_pat(decoded_token: Optional[dict]) -> bool: + return bool(decoded_token) and decoded_token.get("auth_method") == AUTH_METHOD_PAT diff --git a/docsgpt/api/user/me/routes.py b/docsgpt/api/user/me/routes.py index 34e108ad..67e31dfb 100644 --- a/docsgpt/api/user/me/routes.py +++ b/docsgpt/api/user/me/routes.py @@ -12,6 +12,8 @@ from __future__ import annotations from flask import jsonify, make_response, request from flask_restx import Namespace, Resource +from docsgpt.api.pat.tokens import is_pat + me_ns = Namespace("me", description="Current user identity and roles", path="/api") @@ -31,4 +33,13 @@ class MeResource(Resource): value = decoded_token.get(field) if value: body[field] = value + if is_pat(decoded_token): + # Lets a CLI or pipeline confirm what its token is allowed to do. + body["auth_method"] = "pat" + body["token"] = { + "id": decoded_token.get("pat_id"), + "name": decoded_token.get("pat_name"), + "scopes": decoded_token.get("scopes") or [], + "resource_filter": decoded_token.get("resource_filter") or {}, + } return make_response(jsonify(body), 200) diff --git a/docsgpt/auth.py b/docsgpt/auth.py index d56140d7..a38623a5 100644 --- a/docsgpt/auth.py +++ b/docsgpt/auth.py @@ -4,7 +4,28 @@ from jose.exceptions import ExpiredSignatureError from docsgpt.core.settings import settings +# Claims only the PAT verifier may set. Dropped from decoded JWTs so a session +# token can never present itself as a (differently scoped) personal access token. +_PAT_ONLY_CLAIMS = ("auth_method", "pat_id", "pat_name", "scopes", "resource_filter") + + +def _bearer_value(request): + header = request.headers.get("Authorization") + if not header or not isinstance(header, str): + return None + scheme, _, value = header.partition(" ") + return value.strip() if scheme.lower() == "bearer" and value else header.strip() + + def handle_auth(request, data={}): + # Personal access tokens are opaque (not JWTs) and resolve against the + # database in every auth mode that supports them, including AUTH_TYPE unset. + from docsgpt.api.pat.tokens import authenticate_pat, looks_like_pat + + bearer = _bearer_value(request) + if looks_like_pat(bearer): + return authenticate_pat(bearer, request) + if settings.AUTH_TYPE in ["simple_jwt", "session_jwt", "oidc"]: jwt_token = request.headers.get("Authorization") if not jwt_token: @@ -26,6 +47,8 @@ def handle_auth(request, data={}): # requirement is scoped to oidc. options={"verify_exp": is_oidc, "require_exp": is_oidc}, ) + for claim in _PAT_ONLY_CLAIMS: + decoded_token.pop(claim, None) return decoded_token except ExpiredSignatureError: return { diff --git a/tests/api/test_asgi_auth.py b/tests/api/test_asgi_auth.py index cf72eb43..341ded0d 100644 --- a/tests/api/test_asgi_auth.py +++ b/tests/api/test_asgi_auth.py @@ -114,3 +114,45 @@ def test_json_error_shape(): response = asgi_auth.json_error("Forbidden", 403) assert response.status_code == 403 assert json.loads(response.body) == {"success": False, "message": "Forbidden"} + + +_PAT_CLAIMS = { + "sub": "alice", + "auth_method": "pat", + "pat_id": "t1", + "scopes": ["chat:run"], + "resource_filter": {}, +} + + +@pytest.mark.unit +@pytest.mark.asyncio +class TestPersonalAccessTokens: + async def test_routes_reject_tokens_unless_they_name_a_scope(self): + with patch.object(asgi_auth, "handle_auth", return_value=dict(_PAT_CLAIMS)): + decoded, error = await asgi_auth.authenticate(_request()) + assert decoded is None + assert error.status_code == 403 + assert json.loads(error.body)["error"] == "insufficient_scope" + + async def test_token_with_the_route_scope_passes(self): + with patch.object(asgi_auth, "handle_auth", return_value=dict(_PAT_CLAIMS)): + decoded, error = await asgi_auth.authenticate(_request(), pat_scope="chat:run") + assert error is None + assert decoded["sub"] == "alice" + + async def test_token_without_the_route_scope_is_refused(self): + claims = dict(_PAT_CLAIMS, scopes=["agents:read"]) + with patch.object(asgi_auth, "handle_auth", return_value=claims): + decoded, error = await asgi_auth.authenticate(_request(), pat_scope="chat:run") + assert decoded is None + assert error.status_code == 403 + + async def test_session_denylist_is_not_consulted_for_tokens(self, monkeypatch): + monkeypatch.setattr(asgi_auth.settings, "AUTH_TYPE", "oidc") + with patch.object(asgi_auth, "handle_auth", return_value=dict(_PAT_CLAIMS)), patch.object( + asgi_auth, "oidc_session_denied", return_value=True + ) as denied: + decoded, error = await asgi_auth.authenticate(_request(), pat_scope="chat:run") + denied.assert_not_called() + assert error is None and decoded is not None diff --git a/tests/api/test_pat_tokens.py b/tests/api/test_pat_tokens.py new file mode 100644 index 00000000..1e38d0ce --- /dev/null +++ b/tests/api/test_pat_tokens.py @@ -0,0 +1,275 @@ +"""Unit tests for docsgpt/api/pat/tokens.py and the PAT branch of handle_auth.""" + +from __future__ import annotations + +from contextlib import contextmanager +from datetime import datetime, timedelta, timezone +from unittest.mock import Mock, patch + +import pytest + +from docsgpt.api.pat import tokens + + +def _request(authorization=None, ip="10.1.1.1"): + request = Mock() + request.headers = {"Authorization": authorization} if authorization else {} + request.remote_addr = ip + return request + + +@contextmanager +def _db(row): + repo = Mock() + repo.find_active_by_hash.return_value = row + + @contextmanager + def _conn(): + yield Mock() + + with patch.object(tokens, "db_readonly", _conn), patch.object( + tokens, "db_session", _conn + ), patch.object(tokens, "PersonalAccessTokensRepository", return_value=repo): + yield repo + + +_ROW = { + "id": "11111111-1111-1111-1111-111111111111", + "user_id": "alice", + "name": "ci", + "scopes": ["agents:write", "chat:run"], + "resource_filter": {"agents": ["22222222-2222-2222-2222-222222222222"]}, + "last_used_at": None, +} + + +@pytest.mark.unit +class TestTokenFormat: + def test_generate_returns_prefixed_secret_hash_and_display_prefix(self): + token, token_hash, display = tokens.generate_token() + assert token.startswith("dgpt_pat_") + assert len(token) > 40 + assert token_hash == tokens.hash_token(token) + assert len(token_hash) == 64 + assert token.startswith(display) and len(display) == len("dgpt_pat_") + 6 + + def test_tokens_are_unique(self): + assert tokens.generate_token()[0] != tokens.generate_token()[0] + + def test_redact_never_returns_the_secret(self): + token, _, display = tokens.generate_token() + assert tokens.redact(token) == display + "…" + assert tokens.redact("some-agent-key") == "some…" + assert tokens.redact(None) == "" + + @pytest.mark.parametrize( + "value,expected", + [("dgpt_pat_abc", True), ("eyJhbGciOi", False), ("", False), (None, False)], + ) + def test_looks_like_pat(self, value, expected): + assert tokens.looks_like_pat(value) is expected + + +@pytest.mark.unit +class TestScopes: + def test_write_implies_read(self): + assert tokens.expand_scopes(["agents:write"]) == {"agents:write", "agents:read"} + + def test_standalone_scopes_imply_nothing(self): + assert tokens.expand_scopes(["agents:keys", "chat:run"]) == {"agents:keys", "chat:run"} + + def test_normalize_sorts_and_dedupes(self): + assert tokens.normalize_scopes(["sources:read", "agents:read", "sources:read"]) == [ + "agents:read", + "sources:read", + ] + + @pytest.mark.parametrize("raw", [None, [], "agents:read", ["admin:all"], [1], ["agents:read", "x"]]) + def test_normalize_rejects_bad_input(self, raw): + with pytest.raises(ValueError): + tokens.normalize_scopes(raw) + + def test_no_admin_scope_exists(self): + assert not any(s.startswith("admin") for s in tokens.SCOPES) + + +@pytest.mark.unit +class TestResourceFilter: + UUID_A = "22222222-2222-2222-2222-222222222222" + + def test_empty_is_unrestricted(self): + assert tokens.normalize_resource_filter(None, ["agents:read"]) == {} + assert tokens.normalize_resource_filter({}, ["agents:read"]) == {} + + def test_canonicalizes_and_dedupes_ids(self): + out = tokens.normalize_resource_filter( + {"agents": [self.UUID_A.upper(), self.UUID_A]}, ["agents:read"] + ) + assert out == {"agents": [self.UUID_A]} + + def test_chat_scope_allows_agent_and_source_restrictions(self): + out = tokens.normalize_resource_filter( + {"agents": [self.UUID_A], "sources": [self.UUID_A]}, ["chat:run"] + ) + assert set(out) == {"agents", "sources"} + + @pytest.mark.parametrize( + "raw,scopes", + [ + ("agents", ["agents:read"]), + ({"conversations": [UUID_A]}, ["conversations:read"]), + ({"sources": [UUID_A]}, ["agents:read"]), + ({"agents": []}, ["agents:read"]), + ({"agents": "all"}, ["agents:read"]), + ({"agents": ["not-a-uuid"]}, ["agents:read"]), + ({"agents": [UUID_A] * 201}, ["agents:read"]), + ], + ) + def test_rejects_bad_filters(self, raw, scopes): + with pytest.raises(ValueError): + tokens.normalize_resource_filter(raw, scopes) + + +@pytest.mark.unit +class TestExpiryPolicy: + @pytest.fixture(autouse=True) + def _policy(self, monkeypatch): + monkeypatch.setattr(tokens.settings, "PAT_DEFAULT_LIFETIME_DAYS", 90) + monkeypatch.setattr(tokens.settings, "PAT_MAX_LIFETIME_DAYS", 365) + monkeypatch.setattr(tokens.settings, "PAT_ALLOW_NON_EXPIRING", False) + + def _days(self, expires_at): + return round((expires_at - datetime.now(timezone.utc)) / timedelta(days=1)) + + def test_default_lifetime(self): + assert self._days(tokens.resolve_expiry(None)) == 90 + + def test_explicit_lifetime(self): + assert self._days(tokens.resolve_expiry(7)) == 7 + + def test_max_lifetime_enforced(self): + assert self._days(tokens.resolve_expiry(365)) == 365 + with pytest.raises(ValueError): + tokens.resolve_expiry(366) + + def test_non_expiring_refused_unless_operator_allows(self, monkeypatch): + with pytest.raises(ValueError, match="disabled"): + tokens.resolve_expiry(0) + monkeypatch.setattr(tokens.settings, "PAT_ALLOW_NON_EXPIRING", True) + assert tokens.resolve_expiry(0) is None + + @pytest.mark.parametrize("raw", [-1, "30", 1.5, True]) + def test_rejects_bad_values(self, raw): + with pytest.raises(ValueError): + tokens.resolve_expiry(raw) + + +@pytest.mark.unit +class TestAuthenticatePat: + @pytest.fixture(autouse=True) + def _enabled(self, monkeypatch): + monkeypatch.setattr(tokens.settings, "PAT_ENABLED", True) + monkeypatch.setattr(tokens.settings, "AUTH_TYPE", "oidc") + + def test_valid_token_yields_claims_from_the_row(self): + with _db(dict(_ROW)) as repo: + claims = tokens.authenticate_pat("dgpt_pat_secret", _request()) + repo.find_active_by_hash.assert_called_once_with(tokens.hash_token("dgpt_pat_secret")) + assert claims == { + "sub": "alice", + "auth_method": "pat", + "pat_id": _ROW["id"], + "pat_name": "ci", + "scopes": ["agents:read", "agents:write", "chat:run"], + "resource_filter": _ROW["resource_filter"], + } + repo.touch_last_used.assert_called_once() + assert repo.touch_last_used.call_args.args[:2] == (_ROW["id"], "10.1.1.1") + + def test_unknown_token_is_invalid(self): + with _db(None): + assert tokens.authenticate_pat("dgpt_pat_x", _request())["error"] == "invalid_token" + + def test_lookup_failure_fails_closed(self): + with _db(dict(_ROW)) as repo: + repo.find_active_by_hash.side_effect = RuntimeError("db down") + assert tokens.authenticate_pat("dgpt_pat_x", _request())["error"] == "invalid_token" + + def test_usage_write_failure_never_fails_the_request(self): + with _db(dict(_ROW)) as repo: + repo.touch_last_used.side_effect = RuntimeError("db down") + assert tokens.authenticate_pat("dgpt_pat_x", _request())["sub"] == "alice" + + def test_recent_usage_skips_the_write(self): + row = dict(_ROW, last_used_at=datetime.now(timezone.utc).isoformat()) + with _db(row) as repo: + tokens.authenticate_pat("dgpt_pat_x", _request()) + repo.touch_last_used.assert_not_called() + + @pytest.mark.parametrize("auth_type", ["simple_jwt", "session_jwt"]) + def test_rejected_where_there_is_no_stable_user_identity(self, monkeypatch, auth_type): + monkeypatch.setattr(tokens.settings, "AUTH_TYPE", auth_type) + with _db(dict(_ROW)) as repo: + assert tokens.authenticate_pat("dgpt_pat_x", _request())["error"] == "invalid_token" + repo.find_active_by_hash.assert_not_called() + + def test_rejected_when_disabled(self, monkeypatch): + monkeypatch.setattr(tokens.settings, "PAT_ENABLED", False) + with _db(dict(_ROW)): + assert tokens.authenticate_pat("dgpt_pat_x", _request())["error"] == "invalid_token" + + def test_works_with_auth_disabled(self, monkeypatch): + monkeypatch.setattr(tokens.settings, "AUTH_TYPE", None) + with _db(dict(_ROW)): + assert tokens.authenticate_pat("dgpt_pat_x", _request())["sub"] == "alice" + + def test_starlette_request_ip(self): + request = Mock(spec=["headers", "client"]) + request.client.host = "10.9.9.9" + with _db(dict(_ROW)) as repo: + tokens.authenticate_pat("dgpt_pat_x", request) + assert repo.touch_last_used.call_args.args[1] == "10.9.9.9" + + +@pytest.mark.unit +class TestHandleAuthPatBranch: + def test_pat_bearer_goes_to_the_pat_verifier_in_any_auth_mode(self): + from docsgpt import auth + + for auth_type in (None, "oidc", "simple_jwt"): + with patch.object(auth.settings, "AUTH_TYPE", auth_type), patch( + "docsgpt.api.pat.tokens.authenticate_pat", return_value={"sub": "alice"} + ) as verifier: + request = _request("Bearer dgpt_pat_secret") + assert auth.handle_auth(request) == {"sub": "alice"} + verifier.assert_called_once_with("dgpt_pat_secret", request) + + def test_bearer_scheme_is_case_insensitive(self): + from docsgpt import auth + + with patch("docsgpt.api.pat.tokens.authenticate_pat", return_value={"sub": "a"}) as verifier: + auth.handle_auth(_request("bearer dgpt_pat_secret")) + verifier.assert_called_once() + + def test_jwt_cannot_smuggle_pat_claims(self): + from jose import jwt + + from docsgpt import auth + + forged = jwt.encode( + { + "sub": "mallory", + "auth_method": "pat", + "scopes": ["agents:write"], + "resource_filter": {}, + "pat_id": "x", + "pat_name": "x", + }, + "secret", + algorithm="HS256", + ) + with patch.object(auth.settings, "AUTH_TYPE", "simple_jwt"), patch.object( + auth.settings, "JWT_SECRET_KEY", "secret" + ): + decoded = auth.handle_auth(_request(f"Bearer {forged}")) + assert decoded == {"sub": "mallory"} From 3a74aa23f0c7566058c8fb404ccfae665e70c26c Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sat, 19 Sep 2026 23:39:58 +0100 Subject: [PATCH 098/130] feat(pat): token management API and admin revocation Users list, create and revoke their own tokens under /api/user/tokens; the plaintext is returned once at creation. Admins can list a user's tokens and revoke any token, and the admin revoke-sessions action now revokes the user's tokens too. Creation and revocation are written to auth_events. --- docsgpt/api/admin/routes.py | 14 ++- docsgpt/api/pat/routes.py | 229 ++++++++++++++++++++++++++++++++++++ 2 files changed, 242 insertions(+), 1 deletion(-) create mode 100644 docsgpt/api/pat/routes.py diff --git a/docsgpt/api/admin/routes.py b/docsgpt/api/admin/routes.py index 334ad936..0bad4772 100644 --- a/docsgpt/api/admin/routes.py +++ b/docsgpt/api/admin/routes.py @@ -23,6 +23,9 @@ from docsgpt.api.user.authz import ROLE_ADMIN, admin_required from docsgpt.storage.db.repositories.admin_stats import AdminStatsRepository from docsgpt.storage.db.repositories.auth_events import AuthEventsRepository from docsgpt.storage.db.repositories.device_audit_log import DeviceAuditLogRepository +from docsgpt.storage.db.repositories.personal_access_tokens import ( + PersonalAccessTokensRepository, +) from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository from docsgpt.storage.db.repositories.user_roles import UserRolesRepository from docsgpt.storage.db.repositories.users import UsersRepository @@ -246,12 +249,21 @@ class AdminUserSessionsResource(Resource): """Force-logout: revoke the user's live OIDC sessions (best-effort).""" ok = denylist.deny_user(user_id) with db_session() as conn: + # A forced logout that left API credentials alive would not be one. + tokens_revoked = PersonalAccessTokensRepository(conn).revoke_all_for_user( + user_id, reason="admin_sessions_revoked" + ) AuthEventsRepository(conn).insert( user_id, "admin_sessions_revoked", ip=request.remote_addr, user_agent=request.headers.get("User-Agent"), - metadata={"by": _actor(), "via": "admin_api", "persisted": ok}, + metadata={ + "by": _actor(), + "via": "admin_api", + "persisted": ok, + "personal_access_tokens_revoked": tokens_revoked, + }, ) return make_response(jsonify({"success": True, "revoked": ok}), 200) diff --git a/docsgpt/api/pat/routes.py b/docsgpt/api/pat/routes.py new file mode 100644 index 00000000..84360a1b --- /dev/null +++ b/docsgpt/api/pat/routes.py @@ -0,0 +1,229 @@ +"""Personal access token management. + +``/api/user/tokens`` lets a signed-in user list, create and revoke their own +tokens; ``/api/admin/...`` lets an admin inspect and revoke anyone's. None of +these routes accept a PAT (see ``docsgpt/api/pat/rules.py``), so a leaked +token can neither mint a replacement nor widen itself. +""" + +from __future__ import annotations + +import uuid + +from flask import jsonify, make_response, request +from flask_restx import Namespace, Resource +from sqlalchemy.exc import IntegrityError + +from docsgpt.api.pat.tokens import ( + FILTERABLE_FAMILIES, + SCOPES, + auth_type_supports_pats, + generate_token, + is_pat, + normalize_resource_filter, + normalize_scopes, + resolve_expiry, +) +from docsgpt.api.user.authz import admin_required +from docsgpt.core.settings import settings +from docsgpt.storage.db.repositories.auth_events import AuthEventsRepository +from docsgpt.storage.db.repositories.personal_access_tokens import ( + PersonalAccessTokensRepository, +) +from docsgpt.storage.db.session import db_readonly, db_session + +pat_ns = Namespace("tokens", description="Personal access tokens", path="/api") + +_MAX_NAME_LENGTH = 100 + + +def _error(message: str, status: int): + return make_response(jsonify({"success": False, "message": message}), status) + + +def _session_user_id(): + """The caller's id, or ``None`` for anonymous and PAT callers alike.""" + decoded = getattr(request, "decoded_token", None) + if not decoded or is_pat(decoded): + return None + return decoded.get("sub") + + +def _valid_uuid(value: str) -> bool: + try: + uuid.UUID(str(value)) + except (ValueError, AttributeError, TypeError): + return False + return True + + +def serialize_token(row: dict) -> dict: + return { + "id": str(row["id"]), + "name": row["name"], + "token_prefix": row["token_prefix"], + "scopes": list(row.get("scopes") or []), + "resource_filter": row.get("resource_filter") or {}, + "status": row["status"], + "expires_at": row.get("expires_at"), + "last_used_at": row.get("last_used_at"), + "last_used_ip": row.get("last_used_ip"), + "created_at": row.get("created_at"), + "revoked_at": row.get("revoked_at"), + } + + +def _policy() -> dict: + return { + "enabled": auth_type_supports_pats(), + "default_lifetime_days": settings.PAT_DEFAULT_LIFETIME_DAYS, + "max_lifetime_days": settings.PAT_MAX_LIFETIME_DAYS, + "allow_non_expiring": settings.PAT_ALLOW_NON_EXPIRING, + "max_per_user": settings.PAT_MAX_PER_USER, + "filterable_families": list(FILTERABLE_FAMILIES), + } + + +@pat_ns.route("/user/tokens") +class PersonalAccessTokens(Resource): + def get(self): + """List the caller's tokens with the scope catalog and the server's token policy.""" + user_id = _session_user_id() + if not user_id: + return _error("Authentication required", 401) + with db_readonly() as conn: + rows = PersonalAccessTokensRepository(conn).list_for_user(user_id) + return make_response( + jsonify( + { + "success": True, + "tokens": [serialize_token(r) for r in rows], + "scopes": [{"name": k, "description": v} for k, v in SCOPES.items()], + "policy": _policy(), + } + ), + 200, + ) + + def post(self): + """Create a token. The plaintext ``token`` is returned here and never again.""" + user_id = _session_user_id() + if not user_id: + return _error("Authentication required", 401) + if not auth_type_supports_pats(): + return _error("Personal access tokens are not available on this server", 403) + + body = request.get_json(silent=True) or {} + name = body.get("name") + if not isinstance(name, str) or not name.strip(): + return _error("name is required", 400) + name = name.strip() + if len(name) > _MAX_NAME_LENGTH: + return _error(f"name must be at most {_MAX_NAME_LENGTH} characters", 400) + try: + scopes = normalize_scopes(body.get("scopes")) + resource_filter = normalize_resource_filter(body.get("resource_filter"), scopes) + expires_at = resolve_expiry(body.get("expires_in_days")) + except ValueError as exc: + return _error(str(exc), 400) + + token, token_hash, token_prefix = generate_token() + try: + with db_session() as conn: + repo = PersonalAccessTokensRepository(conn) + if repo.count_active(user_id) >= settings.PAT_MAX_PER_USER: + return _error( + f"Token limit reached ({settings.PAT_MAX_PER_USER}); revoke one first", 409 + ) + if repo.name_in_use(user_id, name): + return _error("A token with this name already exists", 409) + row = repo.create( + user_id, + name, + token_hash=token_hash, + token_prefix=token_prefix, + scopes=scopes, + resource_filter=resource_filter, + expires_at=expires_at, + ) + AuthEventsRepository(conn).insert( + user_id, + "pat_created", + ip=request.remote_addr, + user_agent=request.headers.get("User-Agent"), + metadata={ + "token_id": str(row["id"]), + "name": name, + "scopes": scopes, + "resource_filter": resource_filter, + "expires_at": row.get("expires_at"), + }, + ) + except IntegrityError: + # Lost a race against a concurrent create with the same name. + return _error("A token with this name already exists", 409) + return make_response( + jsonify({"success": True, "token": token, "personal_access_token": serialize_token(row)}), + 201, + ) + + +@pat_ns.route("/user/tokens/") +class PersonalAccessToken(Resource): + def delete(self, token_id): + """Revoke one of the caller's tokens. Takes effect on the next request.""" + user_id = _session_user_id() + if not user_id: + return _error("Authentication required", 401) + if not _valid_uuid(token_id): + return _error("Token not found", 404) + with db_session() as conn: + revoked = PersonalAccessTokensRepository(conn).revoke(token_id, user_id) + if revoked: + AuthEventsRepository(conn).insert( + user_id, + "pat_revoked", + ip=request.remote_addr, + user_agent=request.headers.get("User-Agent"), + metadata={"token_id": token_id, "by": user_id}, + ) + if not revoked: + return _error("Token not found", 404) + return make_response(jsonify({"success": True}), 200) + + +@pat_ns.route("/admin/users//tokens") +class AdminUserTokens(Resource): + @admin_required + def get(self, user_id): + """List a user's tokens, revoked ones included.""" + with db_readonly() as conn: + rows = PersonalAccessTokensRepository(conn).list_for_user(user_id, include_revoked=True) + return make_response( + jsonify({"success": True, "tokens": [serialize_token(r) for r in rows]}), 200 + ) + + +@pat_ns.route("/admin/tokens/") +class AdminToken(Resource): + @admin_required + def delete(self, token_id): + """Revoke any user's token.""" + if not _valid_uuid(token_id): + return _error("Token not found", 404) + actor = (getattr(request, "decoded_token", None) or {}).get("sub") + with db_session() as conn: + repo = PersonalAccessTokensRepository(conn) + row = repo.get(token_id) + revoked = bool(row) and repo.revoke(token_id, reason="admin_revoked") + if revoked: + AuthEventsRepository(conn).insert( + row["user_id"], + "pat_revoked", + ip=request.remote_addr, + user_agent=request.headers.get("User-Agent"), + metadata={"token_id": token_id, "by": actor, "via": "admin_api"}, + ) + if not revoked: + return _error("Token not found", 404) + return make_response(jsonify({"success": True}), 200) From 82ec6cba0e0acff04b7ae4a9c90cf99c9d1bec01 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sat, 19 Sep 2026 23:39:58 +0100 Subject: [PATCH 099/130] feat(pat): enforce scopes and resource restrictions, deny by default A central rule table maps each route and method to the scope a token needs; a route that is not listed cannot be called with a token, and a test fails when a registered route is left unclassified. Token management, admin, team management, sign-in, device pairing and OAuth handshakes are never token reachable, and a token never carries the admin role. A token restricted to specific agents, sources, prompts, tools or workflows is held to its allowlist: ids are checked wherever a route carries them, listings are filtered, creation is refused, and routes whose rows cannot be tied to the allowlist are closed. Agent import checks the resolved target. --- docsgpt/api/pat/rules.py | 447 +++++++++++++++++++++++++ docsgpt/api/user/agents/portability.py | 31 ++ docsgpt/api/user/agents/routes.py | 3 +- docsgpt/api/user/prompts/routes.py | 4 +- docsgpt/api/user/sources/routes.py | 3 +- docsgpt/api/user/tools/routes.py | 4 + docsgpt/app.py | 19 +- tests/api/test_agent_portability.py | 69 ++++ tests/api/test_pat_routes.py | 244 ++++++++++++++ tests/api/test_pat_rules.py | 294 ++++++++++++++++ 10 files changed, 1114 insertions(+), 4 deletions(-) create mode 100644 docsgpt/api/pat/rules.py create mode 100644 tests/api/test_pat_routes.py create mode 100644 tests/api/test_pat_rules.py diff --git a/docsgpt/api/pat/rules.py b/docsgpt/api/pat/rules.py new file mode 100644 index 00000000..f340d698 --- /dev/null +++ b/docsgpt/api/pat/rules.py @@ -0,0 +1,447 @@ +"""What a personal access token may call: the scope and resource rule table. + +Authorization for PATs is central and deny by default. ``RULES`` maps a Flask +route (its rule string and method) to the scope it needs; a PAT request to a +route that is not listed is refused, so a new endpoint is unreachable by token +until someone classifies it here. ``tests/api/test_pat_rules.py`` fails when a +registered route is in neither ``RULES`` nor ``DENIED``. + +A token may also carry a resource filter (``{"agents": [ids]}``). For a +restricted family the rule must be able to prove the request stays inside the +allowlist: it names where the id travels (``ids``), or declares that the route +filters its own listing (``listing``), or delegates to the route +(``in_route``). Anything else, creation included, is refused. ``refs`` cover +ids of *other* families a route accepts (an agent update naming a source), and +``blocked_by`` closes routes whose rows hang off a family the rule cannot see +(a schedule belongs to an agent). + +Session (JWT) callers never pass through here. +""" + +from __future__ import annotations + +import json +import uuid +from dataclasses import dataclass +from typing import Any, Callable, Iterable, Optional + +VIEW, QUERY, JSON, FORM, BODY = "view", "query", "json", "form", "body" + +Locator = tuple[str, str] + + +@dataclass(frozen=True) +class Rule: + """Requirement for one route+method. ``scopes`` is any-of; empty means any valid token.""" + + scopes: tuple[str, ...] = () + family: Optional[str] = None + ids: tuple[Locator, ...] = () + refs: tuple[tuple[str, Locator], ...] = () + listing: bool = False + open: bool = False + in_route: bool = False + blocked_by: tuple[str, ...] = () + check: Optional[Callable[[Any, dict], Optional[str]]] = None + + +def _rule(scope: Optional[str] = None, *ids: Locator, any_of: tuple[str, ...] = (), **kwargs) -> Rule: + scopes = any_of or ((scope,) if scope else ()) + family = kwargs.pop("family", None) + if family is None and scope: + family = scope.partition(":")[0] + return Rule(scopes=scopes, family=family, ids=tuple(ids), **kwargs) + + +# Ids of other families that agent create/update accept in their JSON-or-form body. +_AGENT_BODY_REFS = ( + ("sources", (BODY, "source")), + ("sources", (BODY, "sources")), + ("prompts", (BODY, "prompt_id")), + ("tools", (BODY, "tools")), + ("workflows", (BODY, "workflow")), +) + + +def _chat_check(request, resource_filter: dict) -> Optional[str]: + """Keep a restricted token's chat traffic inside its allowlists. + + An agent brings its own sources, prompt and tools, which this table cannot + see, so a token restricted on any family must name an allowed agent (or, + when only sources are restricted, chat against allowed sources directly). + An agent ``api_key`` in the body would swap in an arbitrary agent. + """ + body = _json_body(request) + if body.get("api_key"): + return "A restricted token cannot chat with an agent API key; pass agent_id" + if body.get("workflow"): + # An inline workflow graph (builder preview) can reference any resource. + return "A restricted token cannot run an inline workflow" + agent_ids = _as_ids(body.get("agent_id")) + if "agents" in resource_filter: + if not agent_ids: + return "This token is restricted to specific agents; pass agent_id" + return None # the agent id itself is verified through ``refs`` + if agent_ids: + return "This token is restricted to specific resources and cannot run arbitrary agents" + return None + + +_CHAT = dict( + family=None, + refs=( + ("agents", (JSON, "agent_id")), + ("sources", (JSON, "active_docs")), + ("prompts", (JSON, "prompt_id")), + ("workflows", (JSON, "workflow_id")), + ), + check=_chat_check, +) + +RULES: dict[tuple[str, str], Rule] = { + # Identity and public metadata: any valid token. + ("/api/user/me", "GET"): _rule(open=True), + ("/api/health", "GET"): _rule(open=True), + ("/api/config", "GET"): _rule(open=True), + # Agents + ("/api/get_agent", "GET"): _rule("agents:read", (QUERY, "id")), + ("/api/get_agents", "GET"): _rule("agents:read", listing=True), + ("/api/pinned_agents", "GET"): _rule("agents:read"), + ("/api/shared_agents", "GET"): _rule("agents:read"), + ("/api/template_agents", "GET"): _rule("agents:read", open=True), + ("/api/export_agent", "GET"): _rule("agents:read", (QUERY, "id")), + ("/api/guardrails/catalog", "GET"): _rule("agents:read", open=True), + ("/api/guardrails/events", "GET"): _rule("agents:read", (QUERY, "agent_id")), + ("/api/guardrails/summary", "GET"): _rule("agents:read", (QUERY, "agent_id")), + ("/api/agents/folders/", "GET"): _rule("agents:read", open=True), + ("/api/agents/folders/", "GET"): _rule("agents:read"), + ("/api/create_agent", "POST"): _rule("agents:write", refs=_AGENT_BODY_REFS), + ("/api/update_agent/", "PUT"): _rule( + "agents:write", (VIEW, "agent_id"), refs=_AGENT_BODY_REFS + ), + ("/api/delete_agent", "DELETE"): _rule("agents:write", (QUERY, "id")), + ("/api/adopt_agent", "POST"): _rule("agents:write"), + ("/api/pin_agent", "POST"): _rule("agents:write", (QUERY, "id")), + ("/api/remove_shared_agent", "DELETE"): _rule("agents:write", (QUERY, "id")), + ("/api/share_agent", "PUT"): _rule("agents:write", (JSON, "id")), + ("/api/import_agent/plan", "POST"): _rule("agents:write", in_route=True), + ("/api/import_agent", "POST"): _rule("agents:write", in_route=True), + ("/api/agents/folders/", "POST"): _rule("agents:write"), + ("/api/agents/folders/", "PUT"): _rule("agents:write"), + ("/api/agents/folders/", "DELETE"): _rule("agents:write"), + ("/api/agents/folders/move_agent", "POST"): _rule("agents:write", (JSON, "agent_id")), + ("/api/agents/folders/bulk_move", "POST"): _rule("agents:write", (JSON, "agent_ids")), + ("/api/regenerate_agent_key/", "POST"): _rule("agents:keys", (VIEW, "agent_id")), + ("/api/agent_webhook", "GET"): _rule("agents:keys", (QUERY, "id")), + # Schedules hang off an agent. + ("/api/agents//schedules", "GET"): _rule( + "schedules:read", refs=(("agents", (VIEW, "agent_id")),) + ), + ("/api/agents//schedules", "POST"): _rule( + "schedules:write", refs=(("agents", (VIEW, "agent_id")),), blocked_by=("tools",) + ), + ("/api/schedules/", "GET"): _rule("schedules:read", blocked_by=("agents",)), + ("/api/schedules//runs", "GET"): _rule("schedules:read", blocked_by=("agents",)), + ("/api/schedules//runs/", "GET"): _rule( + "schedules:read", blocked_by=("agents",) + ), + ("/api/schedules/", "PUT"): _rule("schedules:write", blocked_by=("agents", "tools")), + ("/api/schedules/", "PATCH"): _rule("schedules:write", blocked_by=("agents", "tools")), + ("/api/schedules/", "DELETE"): _rule("schedules:write", blocked_by=("agents",)), + ("/api/schedules//run", "POST"): _rule("schedules:write", blocked_by=("agents",)), + # Sources + ("/api/sources", "GET"): _rule("sources:read", listing=True), + # Counted and paged in SQL, so it cannot be narrowed here; restricted tokens use /api/sources. + ("/api/sources/paginated", "GET"): _rule("sources:read"), + ("/api/directory_structure", "GET"): _rule("sources:read", (QUERY, "id")), + ("/api/get_chunks", "GET"): _rule("sources:read", (QUERY, "id")), + ("/api/sources//wiki/pages", "GET"): _rule("sources:read", (VIEW, "source_id")), + ("/api/sources//wiki/page", "GET"): _rule("sources:read", (VIEW, "source_id")), + ("/api/sources//graph", "GET"): _rule("sources:read", (VIEW, "source_id")), + ("/api/sources//graph/node/", "GET"): _rule( + "sources:read", (VIEW, "source_id") + ), + # Ingestion and attachment extraction both report through this poll. + ("/api/task_status", "GET"): _rule(any_of=("sources:read", "sources:write", "chat:run"), open=True), + ("/api/upload", "POST"): _rule("sources:write"), + ("/api/remote", "POST"): _rule("sources:write"), + ("/api/sources/wiki", "POST"): _rule("sources:write"), + ("/api/delete_old", "GET"): _rule("sources:write", (QUERY, "source_id")), + ("/api/manage_sync", "POST"): _rule("sources:write", (JSON, "source_id")), + ("/api/sync_source", "POST"): _rule("sources:write", (JSON, "source_id")), + ("/api/sources/reingest", "POST"): _rule("sources:write", (JSON, "source_id")), + ("/api/manage_source_files", "POST"): _rule("sources:write", (FORM, "source_id")), + ("/api/sources//config", "PATCH"): _rule("sources:write", (VIEW, "source_id")), + ("/api/sources//wiki/page", "PUT"): _rule("sources:write", (VIEW, "source_id")), + ("/api/sources//wiki/convert", "POST"): _rule("sources:write", (VIEW, "source_id")), + ("/api/sources//graphrag/enable", "POST"): _rule( + "sources:write", (VIEW, "source_id") + ), + ("/api/add_chunk", "POST"): _rule("sources:write", (JSON, "id")), + ("/api/update_chunk", "PUT"): _rule("sources:write", (JSON, "id")), + ("/api/delete_chunk", "DELETE"): _rule("sources:write", (QUERY, "id")), + # Prompts + ("/api/get_prompts", "GET"): _rule("prompts:read", listing=True), + ("/api/get_single_prompt", "GET"): _rule("prompts:read", (QUERY, "id")), + ("/api/create_prompt", "POST"): _rule("prompts:write"), + ("/api/update_prompt", "POST"): _rule("prompts:write", (JSON, "id")), + ("/api/delete_prompt", "POST"): _rule("prompts:write", (JSON, "id")), + # Tools + ("/api/available_tools", "GET"): _rule("tools:read", open=True), + ("/api/get_tools", "GET"): _rule("tools:read", listing=True), + ("/api/create_tool", "POST"): _rule("tools:write"), + ("/api/parse_spec", "POST"): _rule("tools:write", open=True), + ("/api/update_tool", "POST"): _rule("tools:write", (JSON, "id")), + ("/api/update_tool_config", "POST"): _rule("tools:write", (JSON, "id")), + ("/api/update_tool_actions", "POST"): _rule("tools:write", (JSON, "id")), + ("/api/update_tool_status", "POST"): _rule("tools:write", (JSON, "id")), + ("/api/delete_tool", "POST"): _rule("tools:write", (JSON, "id")), + ("/api/mcp_server/test", "POST"): _rule("tools:write", open=True), + ("/api/mcp_server/save", "POST"): _rule("tools:write", (JSON, "id")), + # Models + ("/api/models", "GET"): _rule(any_of=("models:read", "chat:run"), open=True), + ("/api/user/models", "GET"): _rule("models:read"), + ("/api/user/models/", "GET"): _rule("models:read"), + ("/api/user/models", "POST"): _rule("models:write"), + ("/api/user/models/", "PATCH"): _rule("models:write"), + ("/api/user/models/", "DELETE"): _rule("models:write"), + ("/api/user/models/test", "POST"): _rule("models:write"), + ("/api/user/models//test", "POST"): _rule("models:write"), + # Workflows + ("/api/workflows", "POST"): _rule("workflows:write"), + ("/api/workflows/", "GET"): _rule("workflows:read", (VIEW, "workflow_id")), + ("/api/workflows/", "PUT"): _rule("workflows:write", (VIEW, "workflow_id")), + ("/api/workflows/", "DELETE"): _rule("workflows:write", (VIEW, "workflow_id")), + # Conversations and analytics span every agent, so an agent-restricted token is kept out. + ("/api/get_conversations", "GET"): _rule("conversations:read", blocked_by=("agents",)), + ("/api/search_conversations", "GET"): _rule("conversations:read", blocked_by=("agents",)), + ("/api/get_single_conversation", "GET"): _rule("conversations:read", blocked_by=("agents",)), + ("/api/messages//tail", "GET"): _rule( + any_of=("conversations:read", "chat:run"), family=None + ), + ("/api/delete_conversation", "POST"): _rule("conversations:write", blocked_by=("agents",)), + ("/api/delete_all_conversations", "GET"): _rule("conversations:write", blocked_by=("agents",)), + ("/api/update_conversation_name", "POST"): _rule("conversations:write", blocked_by=("agents",)), + ("/api/feedback", "POST"): _rule("conversations:write", blocked_by=("agents",)), + ("/api/get_message_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), + ("/api/get_token_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), + ("/api/get_feedback_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), + ("/api/get_tool_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), + ("/api/get_schedule_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), + ("/api/get_user_logs", "POST"): _rule("analytics:read", blocked_by=("agents",)), + # Teams (read only) + ("/api/teams", "GET"): _rule("teams:read"), + ("/api/teams/", "GET"): _rule("teams:read"), + ("/api/teams//members", "GET"): _rule("teams:read"), + ("/api/teams//grants", "GET"): _rule("teams:read"), + ("/api/resource_shares", "GET"): _rule("teams:read"), + # Chat + ("/api/answer", "POST"): _rule("chat:run", **_CHAT), + ("/stream", "POST"): _rule("chat:run", **_CHAT), + ("/api/search", "POST"): _rule("chat:run", **_CHAT), + ("/api/store_attachment", "POST"): _rule("chat:run", family=None), + ("/api/sources//search", "POST"): _rule( + "chat:run", family=None, refs=(("sources", (VIEW, "source_id")),) + ), +} + +#: Routes a PAT may never call, by exact rule string ("*" = every method) or prefix. +#: Token management, admin, login flows and interactive OAuth handshakes need a +#: signed-in session; the rest have no scope yet. Listing them keeps the +#: classification test honest: a new route must land here or in ``RULES``. +DENIED: dict[str, tuple[str, ...]] = { + "/": ("*",), + "/api/user/tokens": ("*",), + "/api/user/tokens/": ("*",), + "/api/generate_token": ("*",), + "/api/combine": ("*",), + "/api/download": ("*",), + "/api/upload_index": ("*",), + "/api/share": ("*",), + "/api/shared_agent": ("*",), + "/api/shared_conversation/": ("*",), + "/api/webhooks/agents/": ("*",), + "/api/images//": ("*",), + "/api/mcp_server/callback": ("*",), + "/api/mcp_server/auth_status": ("*",), + "/api/artifact/": ("*",), + "/api/artifacts": ("*",), + "/api/artifacts/": ("*",), + "/api/artifacts//restore": ("*",), + "/api/artifacts//versions/": ("*",), + "/api/stt": ("*",), + "/api/stt/live/start": ("*",), + "/api/stt/live/chunk": ("*",), + "/api/stt/live/finish": ("*",), + "/api/tts": ("*",), + "/api/teams": ("POST",), + "/api/teams/": ("PUT", "DELETE"), + "/api/teams//members": ("POST",), + "/api/teams//members/": ("*",), + "/api/teams//grants": ("POST", "DELETE"), + "/api/teams//transfer_owner": ("*",), + "/swagger.json": ("*",), +} +DENIED_PREFIXES = ( + "/api/admin/", + "/api/auth/oidc/", + "/api/connectors/", + "/api/devices", + "/scim/", + "/static/", + "/swaggerui/", + "/v1/", +) + + +def is_denied(rule: str, method: str) -> bool: + if rule.startswith(DENIED_PREFIXES): + return True + methods = DENIED.get(rule) + return bool(methods) and ("*" in methods or method in methods) + + +def _json_body(request) -> dict: + body = request.get_json(silent=True) + return body if isinstance(body, dict) else {} + + +def _as_ids(value: Any) -> list[str]: + """Flatten whatever a route accepts as ids: a string, a JSON-encoded or plain list, or an ``{id}`` dict.""" + if value is None or value == "": + return [] + if isinstance(value, dict): + return _as_ids(value.get("id") or value.get("_id") or value.get("workflow_id")) + if isinstance(value, (list, tuple)): + out: list[str] = [] + for item in value: + out.extend(_as_ids(item)) + return out + text = str(value).strip() + if text[:1] in "[{": + try: + return _as_ids(json.loads(text)) + except ValueError: + return [text] + return [text] + + +def _read(request, locator: Locator) -> list[str]: + where, key = locator + if where == VIEW: + return _as_ids((request.view_args or {}).get(key)) + if where == QUERY: + return _as_ids(request.args.get(key)) + if where == JSON: + return _as_ids(_json_body(request).get(key)) + if where == FORM: + return _as_ids(request.form.get(key)) + # BODY: routes that accept JSON or a multipart form interchangeably. + if request.is_json: + return _as_ids(_json_body(request).get(key)) + return _as_ids(request.form.get(key)) + + +def _canonical(value: str) -> str: + try: + return str(uuid.UUID(value)) + except (ValueError, AttributeError, TypeError): + return value + + +def _all_allowed(ids: Iterable[str], allowed: Iterable[str]) -> bool: + allowlist = {_canonical(a) for a in allowed} + return all(_canonical(i) in allowlist for i in ids) + + +def authorize(request, decoded_token: dict) -> Optional[tuple[dict, int]]: + """Check a PAT request against the table. ``None`` allows; otherwise ``(body, status)``.""" + url_rule = getattr(request, "url_rule", None) + rule = RULES.get((url_rule.rule, request.method)) if url_rule is not None else None + if rule is None: + return ( + { + "success": False, + "error": "not_available_to_tokens", + "message": "This endpoint cannot be called with a personal access token", + }, + 403, + ) + granted = set(decoded_token.get("scopes") or []) + if rule.scopes and not granted.intersection(rule.scopes): + return ( + { + "success": False, + "error": "insufficient_scope", + "message": f"Token lacks the required scope: {' or '.join(rule.scopes)}", + "required_scope": rule.scopes[0], + }, + 403, + ) + resource_filter = decoded_token.get("resource_filter") or {} + if not resource_filter: + return None + reason = _check_resources(request, rule, resource_filter) + if reason is None: + return None + return ({"success": False, "error": "resource_not_allowed", "message": reason}, 403) + + +def _check_resources(request, rule: Rule, resource_filter: dict) -> Optional[str]: + for family in rule.blocked_by: + if family in resource_filter: + return f"This endpoint is not available to a token restricted to specific {family}" + if rule.check is not None: + reason = rule.check(request, resource_filter) + if reason: + return reason + for family, locator in rule.refs: + if family not in resource_filter: + continue + ids = _read(request, locator) + if ids and not _all_allowed(ids, resource_filter[family]): + return f"Token is not allowed to use one of the referenced {family}" + family = rule.family + if family is None or family not in resource_filter or rule.open or rule.listing or rule.in_route: + return None + ids = [i for locator in rule.ids for i in _read(request, locator)] + if not ids: + return f"This token is restricted to specific {family} and cannot use this endpoint" + if not _all_allowed(ids, resource_filter[family]): + return f"Token is not allowed to access this resource ({family})" + return None + + +def allowed_ids(request, family: str) -> Optional[set[str]]: + """The caller's allowlist for ``family``, or ``None`` when unrestricted (or not a PAT). + + Used by listing routes (``listing=True``) and by ``in_route`` handlers. + """ + decoded = getattr(request, "decoded_token", None) or {} + if decoded.get("auth_method") != "pat": + return None + ids = (decoded.get("resource_filter") or {}).get(family) + if ids is None: + return None + return {_canonical(str(i)) for i in ids} + + +def filter_listing(request, family: str, items: list, key: str = "id") -> list: + """Drop rows outside the caller's allowlist. Rows without a UUID id (built-in presets) are kept.""" + allowed = allowed_ids(request, family) + if allowed is None: + return items + kept = [] + for item in items: + value = str(item.get(key, "")) + if not _is_uuid(value) or _canonical(value) in allowed: + kept.append(item) + return kept + + +def _is_uuid(value: str) -> bool: + try: + uuid.UUID(value) + except (ValueError, AttributeError, TypeError): + return False + return True diff --git a/docsgpt/api/user/agents/portability.py b/docsgpt/api/user/agents/portability.py index 621d0c2d..a9121532 100644 --- a/docsgpt/api/user/agents/portability.py +++ b/docsgpt/api/user/agents/portability.py @@ -36,6 +36,7 @@ from docsgpt.agents.default_tools import ( synthesized_tool_name_for_id, ) from docsgpt.api import api +from docsgpt.api.pat.rules import allowed_ids from docsgpt.core.model_utils import validate_model_id from docsgpt.core.url_validation import SSRFError, validate_url from docsgpt.security.safe_url import UnsafeUserUrlError, validate_user_base_url @@ -1717,6 +1718,32 @@ def _read_import_payload(req): return raw.decode("utf-8", "replace"), {} +def _restricted_token_denial(conn, user: str, doc: dict) -> Optional[str]: + """Why a resource-restricted personal access token may not import ``doc``, if it may not. + + An import resolves sources, prompts and tools by name and may create them, + so a token restricted on any of those families cannot be held to its + allowlist here. A token restricted to specific agents may update exactly + those; it can never create one. + """ + for family in ("sources", "prompts", "tools", "workflows"): + if allowed_ids(request, family) is not None: + return f"A token restricted to specific {family} cannot import agents" + allowed_agents = allowed_ids(request, "agents") + if allowed_agents is None: + return None + target = _resolve_target(conn, user, doc.get("metadata") or {}) + if target["action"] != "update" or target["agent_id"] not in allowed_agents: + return "This token is restricted to specific agents and may only update those" + return None + + +def _token_denied_response(reason: str): + return make_response( + jsonify({"success": False, "error": "resource_not_allowed", "message": reason}), 403 + ) + + @agents_portability_ns.route("/export_agent") class ExportAgent(Resource): @api.doc(params={"id": "Agent ID"}, description="Export an agent as YAML") @@ -1763,6 +1790,8 @@ class ImportAgentPlan(Resource): return make_response(jsonify({"success": False, "message": str(exc)}), 400) try: with db_readonly() as conn: + if reason := _restricted_token_denial(conn, user, doc): + return _token_denied_response(reason) plan = plan_import(conn, user, doc) except Exception: current_app.logger.error("Agent import plan failed", exc_info=True) @@ -1786,6 +1815,8 @@ class ImportAgent(Resource): return make_response(jsonify({"success": False, "message": str(exc)}), 400) try: with db_session() as conn: + if reason := _restricted_token_denial(conn, user, doc): + return _token_denied_response(reason) result = apply_import(conn, user, doc, resolution) except AgentImportError as exc: # Apply-time rejection of the user's document (e.g. the workflow diff --git a/docsgpt/api/user/agents/routes.py b/docsgpt/api/user/agents/routes.py index 2ddcda5d..757dd933 100644 --- a/docsgpt/api/user/agents/routes.py +++ b/docsgpt/api/user/agents/routes.py @@ -9,6 +9,7 @@ from flask_restx import fields, Namespace, Resource from pydantic import ValidationError as PydanticValidationError from docsgpt.api import api +from docsgpt.api.pat.rules import filter_listing from docsgpt.guardrails.config import AgentConfig from docsgpt.api.user.base import ( copy_agent_image_for_user, @@ -498,7 +499,7 @@ class GetAgents(Resource): except Exception as err: current_app.logger.error(f"Error retrieving agents: {err}", exc_info=True) return make_response(jsonify({"success": False}), 400) - return make_response(jsonify(list_agents), 200) + return make_response(jsonify(filter_listing(request, "agents", list_agents)), 200) @agents_ns.route("/create_agent") diff --git a/docsgpt/api/user/prompts/routes.py b/docsgpt/api/user/prompts/routes.py index d798d01b..af5b9867 100644 --- a/docsgpt/api/user/prompts/routes.py +++ b/docsgpt/api/user/prompts/routes.py @@ -5,6 +5,7 @@ from flask import current_app, jsonify, make_response, request from flask_restx import fields, Namespace, Resource from docsgpt.api import api +from docsgpt.api.pat.rules import filter_listing from docsgpt.api.user.team_sharing import team_access_for, visible_with_access from docsgpt.storage.db.repositories.prompts import PromptsRepository from docsgpt.prompts.composer import compose_preset, is_composed_preset @@ -91,7 +92,8 @@ class GetPrompts(Resource): except Exception as err: current_app.logger.error(f"Error retrieving prompts: {err}", exc_info=True) return make_response(jsonify({"success": False}), 400) - return make_response(jsonify(list_prompts), 200) + # Presets (default/creative/strict) have no row id and stay visible to a restricted token. + return make_response(jsonify(filter_listing(request, "prompts", list_prompts)), 200) @prompts_ns.route("/get_single_prompt") diff --git a/docsgpt/api/user/sources/routes.py b/docsgpt/api/user/sources/routes.py index dc675063..2708b387 100644 --- a/docsgpt/api/user/sources/routes.py +++ b/docsgpt/api/user/sources/routes.py @@ -10,6 +10,7 @@ from pydantic import ValidationError from docsgpt.agents.tools.path_utils import validate_tool_path from docsgpt.api import api +from docsgpt.api.pat.rules import filter_listing from docsgpt.api.user.tasks import ( convert_source_to_wiki, extract_graph, @@ -130,7 +131,7 @@ class CombinedJson(Resource): except Exception as err: current_app.logger.error(f"Error retrieving sources: {err}", exc_info=True) return make_response(jsonify({"success": False}), 400) - return make_response(jsonify(data), 200) + return make_response(jsonify(filter_listing(request, "sources", data)), 200) @sources_ns.route("/sources/paginated") diff --git a/docsgpt/api/user/tools/routes.py b/docsgpt/api/user/tools/routes.py index ecd6f719..8b70abac 100644 --- a/docsgpt/api/user/tools/routes.py +++ b/docsgpt/api/user/tools/routes.py @@ -16,6 +16,7 @@ from docsgpt.agents.default_tools import ( from docsgpt.agents.tools.spec_parser import parse_spec from docsgpt.agents.tools.tool_manager import ToolManager from docsgpt.api import api +from docsgpt.api.pat.rules import filter_listing from docsgpt.api.user.artifacts.authz import Principal, authorize_artifact from docsgpt.api.user.team_sharing import effective_write_owner, visible_with_access from docsgpt.core.settings import settings @@ -265,6 +266,9 @@ class GetTools(Resource): shaped = _shape_tool(row, ownership="team", force_strip_secret=True) shaped["team_access"] = team_shared.get(str(row["id"])) user_tools.append(shaped) + # A resource-restricted token sees only its allowed tools; the + # default and builtin rows appended below belong to no user. + user_tools = filter_listing(request, "tools", user_tools) # ``scheduler`` is dual-registered (default chat tool + agent- # selectable builtin) and resolves to the same synthetic uuid5 id. diff --git a/docsgpt/app.py b/docsgpt/app.py index 4a11937e..88e2d5cd 100644 --- a/docsgpt/app.py +++ b/docsgpt/app.py @@ -22,8 +22,11 @@ from docsgpt.api.devices import devices_bp # noqa: E402 from docsgpt.api.internal.routes import internal # noqa: E402 from docsgpt.api.oidc import oidc_bp # noqa: E402 from docsgpt.api.oidc.denylist import is_denied as oidc_session_denied # noqa: E402 +from docsgpt.api.pat.routes import pat_ns # noqa: E402 +from docsgpt.api.pat.rules import authorize as authorize_pat # noqa: E402 +from docsgpt.api.pat.tokens import is_pat # noqa: E402 from docsgpt.api.scim import scim_bp # noqa: E402 -from docsgpt.api.user.authz import resolve_roles # noqa: E402 +from docsgpt.api.user.authz import ROLE_USER, resolve_roles # noqa: E402 from docsgpt.api.user.routes import user # noqa: E402 from docsgpt.api.connector.routes import connector # noqa: E402 from docsgpt.api.v1 import v1_bp # noqa: E402 @@ -109,6 +112,8 @@ app.register_blueprint(v1_bp) # first app and raise "add_url_rule can no longer be called". if admin_ns not in api.namespaces: api.add_namespace(admin_ns) +if pat_ns not in api.namespaces: + api.add_namespace(pat_ns) app.config.update( UPLOAD_FOLDER="inputs", CELERY_BROKER_URL=settings.CELERY_BROKER_URL, @@ -317,6 +322,18 @@ def authenticate_request(): request.decoded_token = None elif "error" in decoded_token: return jsonify(decoded_token), 401 + elif is_pat(decoded_token): + # Scopes and resource restrictions are enforced here, centrally and + # deny by default (docsgpt/api/pat/rules.py). A token never carries + # admin, whatever its owner holds, and the session denylist does not + # apply: the token lookup already excludes revoked tokens and + # deactivated users. + denied = authorize_pat(request, decoded_token) + if denied is not None: + body, status = denied + return jsonify(body), status + decoded_token["roles"] = [ROLE_USER] + request.decoded_token = decoded_token elif settings.AUTH_TYPE == "oidc" and oidc_session_denied(decoded_token): # Back-channel logout / SCIM deactivation revoked this session. return ( diff --git a/tests/api/test_agent_portability.py b/tests/api/test_agent_portability.py index 4f64cbd9..a3cc85f6 100644 --- a/tests/api/test_agent_portability.py +++ b/tests/api/test_agent_portability.py @@ -822,3 +822,72 @@ def test_api_tool_without_actions_imports_with_warning(pg_conn, monkeypatch): agent = AgentsRepository(pg_conn).get(result["agent_id"], user) row = UserToolsRepository(pg_conn).get_any(agent["tools"][0], user) assert row["config"]["actions"] == {} + + +# --- resource-restricted personal access tokens ------------------------------ + + +def _pat_request(app, resource_filter): + from flask import request + + ctx = app.test_request_context("/api/import_agent", method="POST") + ctx.push() + request.decoded_token = { + "sub": "u_pat_import", + "auth_method": "pat", + "scopes": ["agents:read", "agents:write"], + "resource_filter": resource_filter, + } + return ctx + + +@pytest.fixture +def flask_ctx_app(): + from flask import Flask + + return Flask(__name__) + + +def test_restricted_token_may_update_only_its_agents(pg_conn, flask_ctx_app): + from docsgpt.api.user.agents.portability import _restricted_token_denial + + user = "u_pat_import" + allowed = _make_agent(pg_conn, user, slug="allowed") + other = _make_agent(pg_conn, user, slug="other") + + ctx = _pat_request(flask_ctx_app, {"agents": [str(allowed["id"])]}) + try: + assert _restricted_token_denial(pg_conn, user, _doc(_slug="allowed")) is None + assert "only update" in _restricted_token_denial(pg_conn, user, _doc(_slug="other")) + assert "only update" in _restricted_token_denial( + pg_conn, user, {"metadata": {"id": str(other["id"])}} + ) + # A slug that matches nothing would create a new agent: never for a restricted token. + assert "only update" in _restricted_token_denial(pg_conn, user, _doc(_slug="brand-new")) + finally: + ctx.pop() + + +@pytest.mark.parametrize("family", ["sources", "prompts", "tools", "workflows"]) +def test_token_restricted_on_referenced_families_cannot_import(pg_conn, flask_ctx_app, family): + from docsgpt.api.user.agents.portability import _restricted_token_denial + + ctx = _pat_request(flask_ctx_app, {family: ["aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"]}) + try: + assert family in _restricted_token_denial(pg_conn, "u_pat_import", _doc()) + finally: + ctx.pop() + + +def test_unrestricted_token_and_sessions_import_freely(pg_conn, flask_ctx_app): + from flask import request + + from docsgpt.api.user.agents.portability import _restricted_token_denial + + ctx = _pat_request(flask_ctx_app, {}) + try: + assert _restricted_token_denial(pg_conn, "u_pat_import", _doc()) is None + request.decoded_token = {"sub": "u_pat_import"} + assert _restricted_token_denial(pg_conn, "u_pat_import", _doc()) is None + finally: + ctx.pop() diff --git a/tests/api/test_pat_routes.py b/tests/api/test_pat_routes.py new file mode 100644 index 00000000..cd157284 --- /dev/null +++ b/tests/api/test_pat_routes.py @@ -0,0 +1,244 @@ +"""Endpoint tests for personal access token management (/api/user/tokens, admin).""" + +from __future__ import annotations + +import json +from contextlib import contextmanager +from unittest.mock import patch + +import pytest +from sqlalchemy import text + +from docsgpt.api.pat import routes as pat_routes +from docsgpt.api.pat import tokens as pat_tokens + +AGENT_A = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" + + +@pytest.fixture +def client(): + from docsgpt.app import app + + app.config["TESTING"] = True + return app.test_client() + + +@pytest.fixture(autouse=True) +def _policy(monkeypatch): + monkeypatch.setattr(pat_tokens.settings, "AUTH_TYPE", "oidc") + monkeypatch.setattr(pat_tokens.settings, "PAT_ENABLED", True) + monkeypatch.setattr(pat_tokens.settings, "PAT_DEFAULT_LIFETIME_DAYS", 90) + monkeypatch.setattr(pat_tokens.settings, "PAT_MAX_LIFETIME_DAYS", 365) + monkeypatch.setattr(pat_tokens.settings, "PAT_ALLOW_NON_EXPIRING", False) + monkeypatch.setattr(pat_tokens.settings, "PAT_MAX_PER_USER", 25) + + +@pytest.fixture +def db(pg_conn): + """Route every session the token code opens onto the test's rolled-back connection.""" + + @contextmanager + def _yield_conn(): + yield pg_conn + + with patch.object(pat_routes, "db_session", _yield_conn), patch.object( + pat_routes, "db_readonly", _yield_conn + ), patch.object(pat_tokens, "db_session", _yield_conn), patch.object( + pat_tokens, "db_readonly", _yield_conn + ): + yield pg_conn + + +@contextmanager +def _session(sub="alice", roles=("user",)): + with patch("docsgpt.app.handle_auth", return_value={"sub": sub}), patch( + "docsgpt.app.resolve_roles", return_value=list(roles) + ), patch("docsgpt.app.oidc_session_denied", return_value=False): + yield + + +def _create(client, **body): + body.setdefault("name", "ci") + body.setdefault("scopes", ["agents:write"]) + with _session(): + return client.post("/api/user/tokens", json=body) + + +class TestCreate: + def test_returns_the_secret_once_and_stores_only_its_hash(self, client, db): + response = _create(client) + assert response.status_code == 201 + body = json.loads(response.data) + token = body["token"] + assert token.startswith("dgpt_pat_") + public = body["personal_access_token"] + assert public["token_prefix"] == token[:15] + assert "token" not in public and "token_hash" not in public + stored = db.execute(text("SELECT token_hash FROM personal_access_tokens")).scalar_one() + assert stored == pat_tokens.hash_token(token) + assert token not in stored + + with _session(): + listed = json.loads(client.get("/api/user/tokens").data) + assert [t["name"] for t in listed["tokens"]] == ["ci"] + assert token not in json.dumps(listed) + + def test_default_expiry_applies(self, client, db): + body = json.loads(_create(client).data) + assert body["personal_access_token"]["expires_at"] is not None + + def test_non_expiring_needs_operator_opt_in(self, client, db, monkeypatch): + assert _create(client, expires_in_days=0).status_code == 400 + monkeypatch.setattr(pat_tokens.settings, "PAT_ALLOW_NON_EXPIRING", True) + response = _create(client, expires_in_days=0) + assert response.status_code == 201 + assert json.loads(response.data)["personal_access_token"]["expires_at"] is None + + def test_lifetime_cap(self, client, db): + assert _create(client, expires_in_days=366).status_code == 400 + + def test_resource_filter_is_stored(self, client, db): + response = _create(client, resource_filter={"agents": [AGENT_A]}) + assert json.loads(response.data)["personal_access_token"]["resource_filter"] == { + "agents": [AGENT_A] + } + + @pytest.mark.parametrize( + "body", + [ + {"name": ""}, + {"name": "x" * 101}, + {"scopes": []}, + {"scopes": ["admin:all"]}, + {"resource_filter": {"agents": ["nope"]}}, + {"resource_filter": {"sources": [AGENT_A]}}, + {"expires_in_days": "soon"}, + ], + ) + def test_validation(self, client, db, body): + assert _create(client, **body).status_code == 400 + + def test_duplicate_name_conflicts(self, client, db): + assert _create(client).status_code == 201 + assert _create(client).status_code == 409 + + def test_per_user_cap(self, client, db, monkeypatch): + monkeypatch.setattr(pat_tokens.settings, "PAT_MAX_PER_USER", 1) + assert _create(client, name="one").status_code == 201 + assert _create(client, name="two").status_code == 409 + + @pytest.mark.parametrize("auth_type", ["simple_jwt", "session_jwt"]) + def test_unavailable_without_a_stable_identity(self, client, db, monkeypatch, auth_type): + monkeypatch.setattr(pat_tokens.settings, "AUTH_TYPE", auth_type) + assert _create(client).status_code == 403 + + def test_disabled_by_operator(self, client, db, monkeypatch): + monkeypatch.setattr(pat_tokens.settings, "PAT_ENABLED", False) + assert _create(client).status_code == 403 + with _session(): + assert json.loads(client.get("/api/user/tokens").data)["policy"]["enabled"] is False + + def test_requires_a_session(self, client, db): + with patch("docsgpt.app.handle_auth", return_value=None): + assert client.post("/api/user/tokens", json={"name": "x", "scopes": ["agents:read"]}).status_code == 401 + + def test_audited(self, client, db): + _create(client) + event, metadata = db.execute( + text("SELECT event, metadata FROM auth_events WHERE user_id = 'alice'") + ).one() + assert event == "pat_created" + assert metadata["scopes"] == ["agents:write"] + assert "dgpt_pat_" not in json.dumps(metadata) + + +class TestList: + def test_includes_scope_catalog_and_policy(self, client, db): + with _session(): + body = json.loads(client.get("/api/user/tokens").data) + assert {s["name"] for s in body["scopes"]} == set(pat_tokens.SCOPES) + assert body["policy"] == { + "enabled": True, + "default_lifetime_days": 90, + "max_lifetime_days": 365, + "allow_non_expiring": False, + "max_per_user": 25, + "filterable_families": list(pat_tokens.FILTERABLE_FAMILIES), + } + + def test_is_owner_scoped(self, client, db): + _create(client) + with _session(sub="bob"): + assert json.loads(client.get("/api/user/tokens").data)["tokens"] == [] + + +class TestEndToEnd: + def test_created_token_authenticates_and_revocation_is_immediate(self, client, db): + created = json.loads(_create(client, scopes=["prompts:read"]).data) + headers = {"Authorization": f"Bearer {created['token']}"} + + me = client.get("/api/user/me", headers=headers) + assert me.status_code == 200 + body = json.loads(me.data) + assert body["user_id"] == "alice" + assert body["token"]["scopes"] == ["prompts:read"] + + # Scoped: cannot list agents, cannot manage tokens. + assert client.get("/api/get_agents", headers=headers).status_code == 403 + assert client.get("/api/user/tokens", headers=headers).status_code == 403 + assert client.post("/api/user/tokens", headers=headers, json={}).status_code == 403 + + with _session(): + token_id = created["personal_access_token"]["id"] + assert client.delete(f"/api/user/tokens/{token_id}").status_code == 200 + assert client.get("/api/user/me", headers=headers).status_code == 401 + + def test_tampered_token_is_rejected(self, client, db): + token = json.loads(_create(client).data)["token"] + response = client.get("/api/user/me", headers={"Authorization": f"Bearer {token}x"}) + assert response.status_code == 401 + + +class TestRevoke: + def test_cannot_revoke_someone_elses_token(self, client, db): + token_id = json.loads(_create(client).data)["personal_access_token"]["id"] + with _session(sub="bob"): + assert client.delete(f"/api/user/tokens/{token_id}").status_code == 404 + + def test_unknown_and_malformed_ids(self, client, db): + with _session(): + assert client.delete(f"/api/user/tokens/{AGENT_A}").status_code == 404 + assert client.delete("/api/user/tokens/not-a-uuid").status_code == 404 + + +class TestAdmin: + def test_requires_admin(self, client, db): + with _session(): + assert client.get("/api/admin/users/alice/tokens").status_code == 403 + assert client.delete(f"/api/admin/tokens/{AGENT_A}").status_code == 403 + + def test_admin_can_list_and_revoke_any_token(self, client, db): + created = json.loads(_create(client).data) + token_id = created["personal_access_token"]["id"] + with _session(sub="root", roles=("admin", "user")): + listed = json.loads(client.get("/api/admin/users/alice/tokens").data) + assert [t["id"] for t in listed["tokens"]] == [token_id] + assert client.delete(f"/api/admin/tokens/{token_id}").status_code == 200 + assert client.delete(f"/api/admin/tokens/{token_id}").status_code == 404 + headers = {"Authorization": f"Bearer {created['token']}"} + assert client.get("/api/user/me", headers=headers).status_code == 401 + + def test_revoke_sessions_also_revokes_tokens(self, client, db): + from docsgpt.api.admin import routes as admin_routes + + @contextmanager + def _yield_conn(): + yield db + + created = json.loads(_create(client).data) + with _session(sub="root", roles=("admin", "user")), patch.object( + admin_routes, "db_session", _yield_conn + ), patch.object(admin_routes.denylist, "deny_user", return_value=True): + assert client.post("/api/admin/users/alice/revoke-sessions").status_code == 200 + headers = {"Authorization": f"Bearer {created['token']}"} + assert client.get("/api/user/me", headers=headers).status_code == 401 diff --git a/tests/api/test_pat_rules.py b/tests/api/test_pat_rules.py new file mode 100644 index 00000000..3a3cd75a --- /dev/null +++ b/tests/api/test_pat_rules.py @@ -0,0 +1,294 @@ +"""The PAT rule table: route classification and enforcement at the Flask chokepoint.""" + +from __future__ import annotations + +import json +from unittest.mock import patch + +import pytest + +from docsgpt.api.pat import rules +from docsgpt.api.pat.tokens import SCOPES, expand_scopes + +AGENT_A = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" +AGENT_B = "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb" +SOURCE_A = "cccccccc-cccc-cccc-cccc-cccccccccccc" +SOURCE_B = "dddddddd-dddd-dddd-dddd-dddddddddddd" + + +@pytest.fixture(scope="module") +def flask_app(): + from docsgpt.app import app + + app.config["TESTING"] = True + return app + + +@pytest.fixture +def client(flask_app): + return flask_app.test_client() + + +def _claims(scopes, resource_filter=None): + return { + "sub": "alice", + "auth_method": "pat", + "pat_id": "t1", + "pat_name": "ci", + "scopes": sorted(expand_scopes(scopes)), + "resource_filter": resource_filter or {}, + } + + +def _call(client, method, path, claims, **kwargs): + """Send a request as a PAT; returns the response, short-circuiting the view.""" + with patch("docsgpt.app.handle_auth", return_value=claims): + return client.open(path, method=method, **kwargs) + + +def _denied(response): + """The rule table refused (as opposed to the view answering 4xx itself).""" + if response.status_code != 403: + return None + return (json.loads(response.data) or {}).get("error") + + +@pytest.mark.unit +class TestClassification: + def test_every_route_is_classified(self, flask_app): + unclassified = [] + for url_rule in flask_app.url_map.iter_rules(): + for method in sorted(url_rule.methods - {"HEAD", "OPTIONS"}): + if (url_rule.rule, method) in rules.RULES: + continue + if rules.is_denied(url_rule.rule, method): + continue + unclassified.append(f"{method} {url_rule.rule}") + assert not unclassified, ( + "New routes must be given a scope in docsgpt/api/pat/rules.py RULES, " + f"or listed in DENIED: {unclassified}" + ) + + def test_no_stale_rules(self, flask_app): + registered = { + (r.rule, m) for r in flask_app.url_map.iter_rules() for m in r.methods + } + assert not [key for key in rules.RULES if key not in registered] + known_rules = {r.rule for r in flask_app.url_map.iter_rules()} + assert not [rule for rule in rules.DENIED if rule not in known_rules] + + def test_a_route_is_never_both_allowed_and_denied(self): + assert not [key for key in rules.RULES if rules.is_denied(*key)] + + def test_rules_only_name_real_scopes(self): + for key, rule in rules.RULES.items(): + for scope in rule.scopes: + assert scope in SCOPES, (key, scope) + + @pytest.mark.parametrize( + "rule,method", + [ + ("/api/user/tokens", "POST"), + ("/api/user/tokens", "GET"), + ("/api/user/tokens/", "DELETE"), + ("/api/admin/users", "GET"), + ("/api/admin/tokens/", "DELETE"), + ("/api/generate_token", "GET"), + ("/api/devices/pairings", "POST"), + ("/api/connectors/auth", "GET"), + ("/api/mcp_server/callback", "GET"), + ], + ) + def test_sensitive_routes_are_never_token_reachable(self, rule, method): + assert (rule, method) not in rules.RULES + assert rules.is_denied(rule, method) + + +@pytest.mark.unit +class TestScopeEnforcement: + def test_unlisted_route_is_refused(self, client): + response = _call(client, "GET", "/api/user/tokens", _claims(list(SCOPES))) + assert _denied(response) == "not_available_to_tokens" + + def test_admin_routes_are_refused_even_with_every_scope(self, client): + response = _call(client, "GET", "/api/admin/users", _claims(list(SCOPES))) + assert _denied(response) == "not_available_to_tokens" + + def test_missing_scope_is_refused_and_names_the_scope(self, client): + response = _call(client, "GET", "/api/get_agents", _claims(["sources:read"])) + assert _denied(response) == "insufficient_scope" + assert json.loads(response.data)["required_scope"] == "agents:read" + + def test_read_scope_cannot_write(self, client): + response = _call( + client, "DELETE", f"/api/delete_agent?id={AGENT_A}", _claims(["agents:read"]) + ) + assert _denied(response) == "insufficient_scope" + + def test_write_scope_can_read(self, client): + response = _call(client, "GET", "/api/get_agents", _claims(["agents:write"])) + assert _denied(response) is None + + def test_agent_key_regeneration_needs_its_own_scope(self, client): + response = _call( + client, "POST", f"/api/regenerate_agent_key/{AGENT_A}", _claims(["agents:write"]) + ) + assert _denied(response) == "insufficient_scope" + + def test_any_token_can_identify_itself(self, client): + response = _call(client, "GET", "/api/user/me", _claims(["prompts:read"])) + assert response.status_code == 200 + body = json.loads(response.data) + assert body["auth_method"] == "pat" + assert body["token"]["scopes"] == ["prompts:read"] + assert body["roles"] == ["user"] + + def test_token_never_carries_admin(self, client): + with patch("docsgpt.app.resolve_roles", return_value=["admin", "user"]) as resolver: + response = _call(client, "GET", "/api/user/me", _claims(["agents:read"])) + resolver.assert_not_called() + assert json.loads(response.data)["roles"] == ["user"] + + def test_session_callers_bypass_the_table(self, client): + with patch("docsgpt.app.handle_auth", return_value={"sub": "alice"}), patch( + "docsgpt.app.resolve_roles", return_value=["user"] + ), patch("docsgpt.app.authorize_pat") as authorize: + client.get("/api/user/me") + authorize.assert_not_called() + + def test_invalid_token_is_401(self, client): + with patch( + "docsgpt.app.handle_auth", + return_value={"error": "invalid_token", "message": "Authentication error: invalid token"}, + ): + assert client.get("/api/get_agents").status_code == 401 + + +@pytest.mark.unit +class TestResourceRestrictions: + def _restricted(self, scopes, **families): + return _claims(scopes, resource_filter=families) + + def test_allowed_id_passes_and_other_id_is_refused(self, client): + claims = self._restricted(["agents:read"], agents=[AGENT_A]) + assert _denied(_call(client, "GET", f"/api/get_agent?id={AGENT_A}", claims)) is None + assert ( + _denied(_call(client, "GET", f"/api/get_agent?id={AGENT_B}", claims)) + == "resource_not_allowed" + ) + + def test_id_comparison_ignores_case(self, client): + claims = self._restricted(["agents:read"], agents=[AGENT_A]) + assert _denied(_call(client, "GET", f"/api/get_agent?id={AGENT_A.upper()}", claims)) is None + + def test_view_arg_ids(self, client): + claims = self._restricted(["agents:write"], agents=[AGENT_A]) + ok = _call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, json={"name": "x"}) + bad = _call(client, "PUT", f"/api/update_agent/{AGENT_B}", claims, json={"name": "x"}) + assert _denied(ok) is None + assert _denied(bad) == "resource_not_allowed" + + def test_json_body_ids_including_lists(self, client): + claims = self._restricted(["agents:write"], agents=[AGENT_A]) + ok = _call(client, "POST", "/api/agents/folders/bulk_move", claims, json={"agent_ids": [AGENT_A]}) + bad = _call( + client, "POST", "/api/agents/folders/bulk_move", claims, json={"agent_ids": [AGENT_A, AGENT_B]} + ) + assert _denied(ok) is None + assert _denied(bad) == "resource_not_allowed" + + def test_missing_id_is_refused_for_a_restricted_token(self, client): + claims = self._restricted(["agents:read"], agents=[AGENT_A]) + assert _denied(_call(client, "GET", "/api/get_agent", claims)) == "resource_not_allowed" + + def test_restricted_token_cannot_create(self, client): + claims = self._restricted(["agents:write"], agents=[AGENT_A]) + response = _call(client, "POST", "/api/create_agent", claims, json={"name": "new"}) + assert _denied(response) == "resource_not_allowed" + + def test_unrestricted_family_is_untouched(self, client): + claims = self._restricted(["agents:write", "prompts:write"], agents=[AGENT_A]) + response = _call(client, "POST", "/api/create_prompt", claims, json={}) + assert _denied(response) is None + + def test_cross_family_references_are_checked(self, client): + claims = self._restricted(["agents:write", "sources:read"], sources=[SOURCE_A]) + ok = _call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, json={"source": SOURCE_A}) + bad = _call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, json={"sources": [SOURCE_A, SOURCE_B]}) + form = _call( + client, "PUT", f"/api/update_agent/{AGENT_A}", claims, + data={"sources": json.dumps([SOURCE_B])}, + ) + assert _denied(ok) is None + assert _denied(bad) == "resource_not_allowed" + assert _denied(form) == "resource_not_allowed" + + def test_routes_hanging_off_a_restricted_family_are_blocked(self, client): + claims = self._restricted(["agents:read", "schedules:read", "analytics:read"], agents=[AGENT_A]) + assert _denied(_call(client, "GET", "/api/schedules/s1", claims)) == "resource_not_allowed" + assert _denied(_call(client, "POST", "/api/get_token_analytics", claims, json={})) == "resource_not_allowed" + assert _denied(_call(client, "GET", f"/api/agents/{AGENT_A}/schedules", claims)) is None + assert _denied(_call(client, "GET", f"/api/agents/{AGENT_B}/schedules", claims)) == "resource_not_allowed" + + def test_sql_paged_listing_is_closed_to_restricted_tokens(self, client): + claims = self._restricted(["sources:read"], sources=[SOURCE_A]) + assert _denied(_call(client, "GET", "/api/sources/paginated", claims)) == "resource_not_allowed" + + +@pytest.mark.unit +class TestChatRestrictions: + def _chat(self, client, claims, body): + return _denied(_call(client, "POST", "/api/answer", claims, json=body)) + + def test_unrestricted_token_can_chat_any_way(self, client): + claims = _claims(["chat:run"]) + assert self._chat(client, claims, {"question": "hi", "api_key": "k"}) is None + assert self._chat(client, claims, {"question": "hi", "agent_id": AGENT_B}) is None + + def test_agent_restricted_token_must_name_an_allowed_agent(self, client): + claims = _claims(["chat:run"], {"agents": [AGENT_A]}) + assert self._chat(client, claims, {"question": "hi", "agent_id": AGENT_A}) is None + assert self._chat(client, claims, {"question": "hi", "agent_id": AGENT_B}) == "resource_not_allowed" + assert self._chat(client, claims, {"question": "hi"}) == "resource_not_allowed" + assert self._chat(client, claims, {"question": "hi", "api_key": "k"}) == "resource_not_allowed" + + def test_restricted_token_cannot_run_an_inline_workflow(self, client): + claims = _claims(["chat:run"], {"agents": [AGENT_A]}) + body = {"question": "hi", "agent_id": AGENT_A, "workflow": {"nodes": []}} + assert self._chat(client, claims, body) == "resource_not_allowed" + + def test_source_restricted_token_cannot_reach_other_sources_through_an_agent(self, client): + claims = _claims(["chat:run"], {"sources": [SOURCE_A]}) + assert self._chat(client, claims, {"question": "hi", "active_docs": SOURCE_A}) is None + assert self._chat(client, claims, {"question": "hi", "active_docs": [SOURCE_A, SOURCE_B]}) == "resource_not_allowed" + assert self._chat(client, claims, {"question": "hi", "agent_id": AGENT_A}) == "resource_not_allowed" + + def test_retrieval_test_honours_the_source_allowlist(self, client): + claims = _claims(["chat:run"], {"sources": [SOURCE_A]}) + ok = _call(client, "POST", f"/api/sources/{SOURCE_A}/search", claims, json={"query": "q"}) + bad = _call(client, "POST", f"/api/sources/{SOURCE_B}/search", claims, json={"query": "q"}) + assert _denied(ok) is None + assert _denied(bad) == "resource_not_allowed" + + +@pytest.mark.unit +class TestListingFilter: + def test_unrestricted_and_session_callers_see_everything(self, flask_app): + from flask import request + + items = [{"id": AGENT_A}, {"id": AGENT_B}] + with flask_app.test_request_context("/"): + request.decoded_token = {"sub": "alice"} + assert rules.filter_listing(request, "agents", items) == items + request.decoded_token = _claims(["agents:read"]) + assert rules.filter_listing(request, "agents", items) == items + + def test_restricted_token_sees_only_its_rows_plus_builtin_presets(self, flask_app): + from flask import request + + items = [{"id": AGENT_A}, {"id": AGENT_B}, {"id": "default"}] + with flask_app.test_request_context("/"): + request.decoded_token = _claims(["prompts:read"], {"prompts": [AGENT_A]}) + assert rules.filter_listing(request, "prompts", items) == [{"id": AGENT_A}, {"id": "default"}] + assert rules.allowed_ids(request, "agents") is None + assert rules.allowed_ids(request, "prompts") == {AGENT_A} From 225d9d3065f2723aaadb829064da2e07ac45b550 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sat, 19 Sep 2026 23:39:58 +0100 Subject: [PATCH 100/130] feat(frontend): access tokens settings tab Settings > Access Tokens lists a user's personal access tokens and lets them create and revoke tokens: scopes grouped by family, optional restriction to specific resources, and an expiry bounded by the server's policy. The secret is shown once after creation and never stored client side. Removes the unused legacy api-key endpoints, wrappers, type and locale block. --- frontend/src/api/endpoints.ts | 5 +- frontend/src/api/services/patService.test.ts | 156 ++++++ frontend/src/api/services/patService.ts | 110 +++++ frontend/src/api/services/userService.ts | 6 - frontend/src/locale/accessTokens.test.ts | 59 +++ frontend/src/locale/de.json | 108 +++- frontend/src/locale/en.json | 108 +++- frontend/src/locale/es.json | 108 +++- frontend/src/locale/jp.json | 108 +++- frontend/src/locale/ru.json | 128 ++++- frontend/src/locale/zh-TW.json | 108 +++- frontend/src/locale/zh.json | 108 +++- .../src/modals/AccessTokenCreatedModal.tsx | 104 ++++ .../src/modals/CreateAccessTokenModal.tsx | 462 ++++++++++++++++++ .../src/settings/PersonalAccessTokens.tsx | 443 +++++++++++++++++ .../src/settings/accessTokenUtils.test.ts | 289 +++++++++++ frontend/src/settings/accessTokenUtils.ts | 224 +++++++++ frontend/src/settings/index.tsx | 7 + frontend/src/settings/types/index.ts | 15 +- 19 files changed, 2563 insertions(+), 93 deletions(-) create mode 100644 frontend/src/api/services/patService.test.ts create mode 100644 frontend/src/api/services/patService.ts create mode 100644 frontend/src/locale/accessTokens.test.ts create mode 100644 frontend/src/modals/AccessTokenCreatedModal.tsx create mode 100644 frontend/src/modals/CreateAccessTokenModal.tsx create mode 100644 frontend/src/settings/PersonalAccessTokens.tsx create mode 100644 frontend/src/settings/accessTokenUtils.test.ts create mode 100644 frontend/src/settings/accessTokenUtils.ts diff --git a/frontend/src/api/endpoints.ts b/frontend/src/api/endpoints.ts index b6c00973..cd19899c 100644 --- a/frontend/src/api/endpoints.ts +++ b/frontend/src/api/endpoints.ts @@ -10,9 +10,6 @@ const endpoints = { MODELS: '/api/models', DOCS: '/api/sources', DOCS_PAGINATED: '/api/sources/paginated', - API_KEYS: '/api/get_api_keys', - CREATE_API_KEY: '/api/create_api_key', - DELETE_API_KEY: '/api/delete_api_key', AGENT: (id: string) => `/api/get_agent?id=${id}`, AGENTS: '/api/get_agents', GUARDRAIL_CATALOG: '/api/guardrails/catalog', @@ -155,6 +152,8 @@ const endpoints = { DEVICE_PAIRINGS: '/api/devices/pairings', DEVICE_PAIRING: (deviceCode: string) => `/api/devices/pairings/${deviceCode}`, + ACCESS_TOKENS: '/api/user/tokens', + ACCESS_TOKEN: (id: string) => `/api/user/tokens/${id}`, }, V1: { CHAT_COMPLETIONS: '/v1/chat/completions', diff --git a/frontend/src/api/services/patService.test.ts b/frontend/src/api/services/patService.test.ts new file mode 100644 index 00000000..63b7a88e --- /dev/null +++ b/frontend/src/api/services/patService.test.ts @@ -0,0 +1,156 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; + +import apiClient from '../client'; +import patService, { AccessTokenApiError } from './patService'; + +afterEach(() => vi.restoreAllMocks()); + +const response = (body: unknown, status = 200) => + ({ + ok: status >= 200 && status < 300, + status, + json: async () => body, + }) as unknown as Response; + +const TOKEN_ROW = { + id: '7f1c2c1e-0f59-4d0a-9d55-0f6a3d1d7a11', + name: 'ci', + token_prefix: 'dgpt_pat_ab12cd', + scopes: ['agents:read'], + resource_filter: {}, + status: 'active', + expires_at: null, + last_used_at: null, + last_used_ip: null, + created_at: '2026-09-01T10:00:00+00:00', + revoked_at: null, +}; + +describe('patService.list', () => { + it('GETs /api/user/tokens with the session token and returns the body', async () => { + const body = { + success: true, + tokens: [TOKEN_ROW], + scopes: [{ name: 'agents:read', description: 'View agents' }], + policy: { enabled: true }, + }; + const spy = vi.spyOn(apiClient, 'get').mockResolvedValue(response(body)); + + const result = await patService.list('session-jwt'); + + expect(spy).toHaveBeenCalledWith('/api/user/tokens', 'session-jwt'); + expect(result.tokens).toEqual([TOKEN_ROW]); + expect(result.scopes[0].name).toBe('agents:read'); + }); + + it('throws with the server message on a non-2xx response', async () => { + vi.spyOn(apiClient, 'get').mockResolvedValue( + response({ success: false, message: 'Authentication required' }, 401), + ); + + await expect(patService.list(null)).rejects.toMatchObject({ + name: 'AccessTokenApiError', + message: 'Authentication required', + status: 401, + }); + }); +}); + +describe('patService.create', () => { + it('POSTs the payload and returns the one-time plaintext token', async () => { + const spy = vi.spyOn(apiClient, 'post').mockResolvedValue( + response( + { + success: true, + token: 'dgpt_pat_secret', + personal_access_token: TOKEN_ROW, + }, + 201, + ), + ); + const payload = { + name: 'ci', + scopes: ['agents:read'], + resource_filter: { agents: [TOKEN_ROW.id] }, + expires_in_days: 30, + }; + + const result = await patService.create(payload, 'session-jwt'); + + expect(spy).toHaveBeenCalledWith( + '/api/user/tokens', + payload, + 'session-jwt', + ); + expect(result.token).toBe('dgpt_pat_secret'); + expect(result.personal_access_token.id).toBe(TOKEN_ROW.id); + }); + + it.each([ + [400, 'Unknown scopes: nope:read'], + [403, 'Personal access tokens are not available on this server'], + [409, 'A token with this name already exists'], + ])('surfaces the %i error message', async (status, message) => { + vi.spyOn(apiClient, 'post').mockResolvedValue( + response({ success: false, message }, status), + ); + + const error = await patService + .create({ name: 'ci', scopes: ['agents:read'] }, 'session-jwt') + .catch((e) => e); + + expect(error).toBeInstanceOf(AccessTokenApiError); + expect(error.message).toBe(message); + expect(error.status).toBe(status); + }); + + it('throws an empty message when the error body is not JSON', async () => { + vi.spyOn(apiClient, 'post').mockResolvedValue({ + ok: false, + status: 502, + json: async () => { + throw new SyntaxError('Unexpected token <'); + }, + } as unknown as Response); + + await expect( + patService.create({ name: 'ci', scopes: ['agents:read'] }, null), + ).rejects.toMatchObject({ message: '', status: 502 }); + }); + + it('treats success:false on a 2xx as a failure', async () => { + vi.spyOn(apiClient, 'post').mockResolvedValue( + response({ success: false, message: 'nope' }, 200), + ); + + await expect( + patService.create({ name: 'ci', scopes: ['agents:read'] }, null), + ).rejects.toThrow('nope'); + }); +}); + +describe('patService.revoke', () => { + it('DELETEs the token by id', async () => { + const spy = vi + .spyOn(apiClient, 'delete') + .mockResolvedValue(response({ success: true })); + + await expect( + patService.revoke(TOKEN_ROW.id, 'session-jwt'), + ).resolves.toEqual({ success: true }); + expect(spy).toHaveBeenCalledWith( + `/api/user/tokens/${TOKEN_ROW.id}`, + 'session-jwt', + ); + }); + + it('throws when the token is not found', async () => { + vi.spyOn(apiClient, 'delete').mockResolvedValue( + response({ success: false, message: 'Token not found' }, 404), + ); + + await expect(patService.revoke('missing', null)).rejects.toThrow( + 'Token not found', + ); + }); +}); diff --git a/frontend/src/api/services/patService.ts b/frontend/src/api/services/patService.ts new file mode 100644 index 00000000..09eb197a --- /dev/null +++ b/frontend/src/api/services/patService.ts @@ -0,0 +1,110 @@ +import apiClient from '../client'; +import endpoints from '../endpoints'; + +export interface PersonalAccessToken { + id: string; + name: string; + /** Display prefix only (e.g. `dgpt_pat_ab12cd`); never the full secret. */ + token_prefix: string; + scopes: string[]; + /** `{family: [resource ids]}`; a missing family means "all resources". */ + resource_filter: Record; + status: string; + expires_at: string | null; + last_used_at: string | null; + last_used_ip: string | null; + created_at: string | null; + revoked_at: string | null; +} + +export interface AccessTokenScope { + name: string; + description: string; +} + +export interface AccessTokenPolicy { + enabled: boolean; + default_lifetime_days: number; + max_lifetime_days: number; + allow_non_expiring: boolean; + max_per_user: number; + filterable_families: string[]; +} + +export interface AccessTokenListResponse { + tokens: PersonalAccessToken[]; + scopes: AccessTokenScope[]; + policy: AccessTokenPolicy; +} + +export interface CreateAccessTokenPayload { + name: string; + scopes: string[]; + resource_filter?: Record; + /** `null`/omitted = server default, `0` = never expires. */ + expires_in_days?: number | null; +} + +export interface CreateAccessTokenResponse { + /** Plaintext secret. The server returns it exactly once. */ + token: string; + personal_access_token: PersonalAccessToken; +} + +/** Error carrying the server's user-facing `message` and the HTTP status. */ +export class AccessTokenApiError extends Error { + status: number; + + constructor(message: string, status: number) { + super(message); + this.name = 'AccessTokenApiError'; + this.status = status; + } +} + +// apiClient resolves to the raw fetch Response (the app convention). Parse it +// here and turn `{success:false, message}` / non-2xx into a thrown error so +// callers can show the server's message inline. +const parse = async (response: Response): Promise => { + let body: { success?: boolean; message?: unknown } | null = null; + try { + body = await response.json(); + } catch { + body = null; + } + if (!response.ok || body?.success === false) { + throw new AccessTokenApiError( + typeof body?.message === 'string' ? body.message : '', + response.status, + ); + } + return body as T; +}; + +const patService = { + list: async (token: string | null): Promise => + parse( + await apiClient.get(endpoints.USER.ACCESS_TOKENS, token), + ), + + create: async ( + payload: CreateAccessTokenPayload, + token: string | null, + ): Promise => + parse( + await apiClient.post(endpoints.USER.ACCESS_TOKENS, payload, token), + ), + + revoke: async ( + id: string, + token: string | null, + ): Promise<{ success: boolean }> => + parse<{ success: boolean }>( + await apiClient.delete( + endpoints.USER.ACCESS_TOKEN(encodeURIComponent(id)), + token, + ), + ), +}; + +export default patService; diff --git a/frontend/src/api/services/userService.ts b/frontend/src/api/services/userService.ts index 585a1ef3..b9de3bc0 100644 --- a/frontend/src/api/services/userService.ts +++ b/frontend/src/api/services/userService.ts @@ -19,12 +19,6 @@ const userService = { apiClient.get(`${endpoints.USER.DOCS}`, token), getDocsWithPagination: (query: string, token: string | null): Promise => apiClient.get(`${endpoints.USER.DOCS_PAGINATED}?${query}`, token), - getAPIKeys: (token: string | null): Promise => - apiClient.get(endpoints.USER.API_KEYS, token), - createAPIKey: (data: any, token: string | null): Promise => - apiClient.post(endpoints.USER.CREATE_API_KEY, data, token), - deleteAPIKey: (data: any, token: string | null): Promise => - apiClient.post(endpoints.USER.DELETE_API_KEY, data, token), getAgent: (id: string, token: string | null): Promise => throttledApiClient.get(endpoints.USER.AGENT(id), token), getAgents: (token: string | null): Promise => diff --git a/frontend/src/locale/accessTokens.test.ts b/frontend/src/locale/accessTokens.test.ts new file mode 100644 index 00000000..5415f4ce --- /dev/null +++ b/frontend/src/locale/accessTokens.test.ts @@ -0,0 +1,59 @@ +import { describe, expect, it } from 'vitest'; + +import de from './de.json'; +import en from './en.json'; +import es from './es.json'; +import jp from './jp.json'; +import ru from './ru.json'; +import zhTW from './zh-TW.json'; +import zh from './zh.json'; + +type Tree = { [key: string]: string | Tree }; + +const PLURAL_SUFFIX = /_(zero|one|two|few|many|other)$/; + +// Plural categories differ per language (ru adds few/many), so compare the +// keys with the suffix stripped. +const flatten = (tree: Tree, prefix = ''): string[] => + Object.entries(tree).flatMap(([key, value]) => + typeof value === 'string' + ? [prefix + key.replace(PLURAL_SUFFIX, '')] + : flatten(value, `${prefix}${key}.`), + ); + +const keysOf = (locale: { settings: object }): string[] => { + const block = (locale.settings as Tree).accessTokens as Tree; + return Array.from(new Set(flatten(block))).sort(); +}; + +const LOCALES = { es, de, jp, ru, zh, zhTW }; + +describe('settings.accessTokens locale block', () => { + it.each(Object.entries(LOCALES))( + '%s has the same keys as en', + (_name, locale) => { + expect(keysOf(locale)).toEqual(keysOf(en)); + }, + ); + + it.each(Object.entries(LOCALES))( + '%s is translated, not an English copy', + (_name, locale) => { + const block = (locale.settings as Tree).accessTokens as Tree; + const source = (en.settings as Tree).accessTokens as Tree; + expect(block.label).not.toBe(source.label); + expect((block.created as Tree).warning).not.toBe( + (source.created as Tree).warning, + ); + }, + ); + + it('every locale provides both plural forms i18next falls back through', () => { + [en, ...Object.values(LOCALES)].forEach((locale) => { + const relative = ((locale.settings as Tree).accessTokens as Tree) + .relative as Tree; + expect(relative.days_one).toBeTruthy(); + expect(relative.days_other).toBeTruthy(); + }); + }); +}); diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json index e64ed52f..7291a332 100644 --- a/frontend/src/locale/de.json +++ b/frontend/src/locale/de.json @@ -321,16 +321,6 @@ } } }, - "apiKeys": { - "label": "Chatbots", - "name": "Name", - "key": "API-Schlüssel", - "sourceDoc": "Quelldokument", - "createNew": "Neu erstellen", - "noData": "Keine vorhandenen Chatbots", - "deleteConfirmation": "Bist du sicher, dass du den API-Schlüssel '{{name}}' löschen möchtest?", - "description": "Hier kannst du deine Chatbots erstellen und verwalten. Chatbots können als Widgets auf Websites eingebunden oder in deinen Anwendungen verwendet werden." - }, "analytics": { "label": "Analytik", "subtitle": "Nachrichtenvolumen, Token-Verbrauch und Nutzerfeedback in deinem Konto verfolgen", @@ -747,6 +737,104 @@ "testFailed": "Verbindungstest fehlgeschlagen" } }, + "accessTokens": { + "label": "Zugriffstoken", + "subtitle": "Mit persönlichen Zugriffstoken können Skripte, die CLI und CI-Pipelines die DocsGPT-API in deinem Namen aufrufen. Gib jedem Token nur die Berechtigungen, die es braucht.", + "createToken": "Token erstellen", + "disabledNotice": "Persönliche Zugriffstoken sind auf diesem Server nicht verfügbar. Sie setzen eine kontobasierte Anmeldung voraus und müssen von deinem Administrator aktiviert werden.", + "limitReached_one": "Du hast das Limit von {{count}} aktiven Token erreicht. Widerrufe eines, um ein neues zu erstellen.", + "limitReached_other": "Du hast das Limit von {{count}} aktiven Token erreicht. Widerrufe eines, um ein neues zu erstellen.", + "loadError": "Zugriffstoken konnten nicht geladen werden. Bitte versuche es erneut.", + "revokeError": "Das Token konnte nicht widerrufen werden. Bitte versuche es erneut.", + "empty": "Noch keine Zugriffstoken", + "emptyHint": "Erstelle ein Token, um die DocsGPT-API aus Skripten und Kommandozeilen-Tools zu nutzen.", + "name": "Name", + "scopes": "Berechtigungen", + "resources": "Ressourcen", + "createdAt": "Erstellt", + "lastUsed": "Zuletzt verwendet", + "expires": "Läuft ab", + "actions": "Aktionen", + "allResources": "Alle Ressourcen", + "never": "Nie", + "expired": "Abgelaufen am {{date}}", + "expiresSoon": "Läuft ab am {{date}}", + "moreScopes": "+{{count}} weitere", + "showLess": "Weniger anzeigen", + "revoke": "Widerrufen", + "revokeAria": "Token {{name}} widerrufen", + "revokeWarning": "Token „{{name}}“ widerrufen? Alles, was es noch verwendet, funktioniert sofort nicht mehr. Das kann nicht rückgängig gemacht werden.", + "restrictionCount": { + "agents_one": "{{count}} Agent", + "agents_other": "{{count}} Agenten", + "sources_one": "{{count}} Quelle", + "sources_other": "{{count}} Quellen", + "prompts_one": "{{count}} Prompt", + "prompts_other": "{{count}} Prompts", + "tools_one": "{{count}} Tool", + "tools_other": "{{count}} Tools", + "workflows_one": "{{count}} Workflow", + "workflows_other": "{{count}} Workflows" + }, + "relative": { + "now": "Gerade eben", + "minutes_one": "Vor {{count}} Minute", + "minutes_other": "Vor {{count}} Minuten", + "hours_one": "Vor {{count}} Stunde", + "hours_other": "Vor {{count}} Stunden", + "days_one": "Vor {{count}} Tag", + "days_other": "Vor {{count}} Tagen" + }, + "families": { + "agents": "Agenten", + "sources": "Quellen", + "prompts": "Prompts", + "tools": "Tools", + "models": "Modelle", + "workflows": "Workflows", + "schedules": "Zeitpläne", + "conversations": "Unterhaltungen", + "analytics": "Analysen", + "teams": "Teams", + "chat": "Chat" + }, + "create": { + "title": "Zugriffstoken erstellen", + "subtitle": "Das Token handelt in deinem Namen, beschränkt auf die Berechtigungen und Ressourcen, die du hier auswählst.", + "name": "Name", + "namePlaceholder": "z. B. CI-Pipeline", + "expiration": "Ablauf", + "expiryDays_one": "{{count}} Tag", + "expiryDays_other": "{{count}} Tage", + "noExpiration": "Kein Ablaufdatum", + "noExpirationHint": "Token ohne Ablaufdatum sind riskanter, falls sie in falsche Hände geraten. Ein Ablaufdatum ist die bessere Wahl.", + "expiresOn": "Läuft ab am {{date}}", + "scopes": "Berechtigungen", + "scopesHint": "Lege fest, was dieses Token darf. Schreibzugriff schließt Lesezugriff ein.", + "impliedByWrite": "Im Schreibzugriff enthalten", + "restrict": "Auf bestimmte Ressourcen beschränken", + "restrictHint": "Optional. Beschränke das Token auf die ausgewählten Elemente. Lass eine Liste leer, um alle zuzulassen.", + "restrictNoFamilies": "Wähle zuerst eine Berechtigung für Agenten, Quellen, Prompts, Tools oder Chat, um deren Ressourcen zu beschränken.", + "allSelected": "Alle (keine Beschränkung)", + "noResources": "Nichts gefunden.", + "searchResources": "Suchen...", + "resourcesError": "Die Liste konnte nicht geladen werden.", + "retry": "Erneut versuchen", + "cancel": "Abbrechen", + "submit": "Token erstellen", + "submitting": "Wird erstellt", + "error": "Das Token konnte nicht erstellt werden. Bitte versuche es erneut." + }, + "created": { + "title": "Token erstellt", + "subtitle": "Dein neues Token „{{name}}“ ist einsatzbereit.", + "warning": "Kopiere das Token jetzt und bewahre es an einem sicheren Ort auf. Aus Sicherheitsgründen wird es nicht noch einmal angezeigt.", + "tokenLabel": "Dein Token", + "usageTitle": "So verwendest du es", + "usageHint": "Exportiere es als Umgebungsvariable und sende es als Bearer-Token im Authorization-Header.", + "done": "Fertig" + } + }, "scrollTabsLeft": "Tabs nach links scrollen", "tabsAriaLabel": "Einstellungs-Tabs", "scrollTabsRight": "Tabs nach rechts scrollen" diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index 03f0f0dd..a5118073 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -325,16 +325,6 @@ } } }, - "apiKeys": { - "label": "Chatbots", - "name": "Name", - "key": "API Key", - "sourceDoc": "Source Document", - "createNew": "Create New", - "noData": "No existing Chatbots", - "deleteConfirmation": "Are you sure you want to delete the API key '{{name}}'?", - "description": "Here you can create and manage your chatbots. Chatbots can be deployed to websites as widgets or used inside your applications." - }, "analytics": { "label": "Analytics", "subtitle": "Track message volume, token usage, and user feedback across your account", @@ -752,6 +742,104 @@ "testFailed": "Connection test failed" } }, + "accessTokens": { + "label": "Access Tokens", + "subtitle": "Personal access tokens let scripts, the CLI and CI pipelines call the DocsGPT API on your behalf. Give each token only the scopes it needs.", + "createToken": "Create token", + "disabledNotice": "Personal access tokens are not available on this server. They require account-based sign-in and have to be enabled by your administrator.", + "limitReached_one": "You have reached the limit of {{count}} active token. Revoke one to create another.", + "limitReached_other": "You have reached the limit of {{count}} active tokens. Revoke one to create another.", + "loadError": "Failed to load access tokens. Please try again.", + "revokeError": "Failed to revoke the token. Please try again.", + "empty": "No access tokens yet", + "emptyHint": "Create a token to use the DocsGPT API from scripts and command-line tools.", + "name": "Name", + "scopes": "Scopes", + "resources": "Resources", + "createdAt": "Created", + "lastUsed": "Last used", + "expires": "Expires", + "actions": "Actions", + "allResources": "All resources", + "never": "Never", + "expired": "Expired {{date}}", + "expiresSoon": "Expires {{date}}", + "moreScopes": "+{{count}} more", + "showLess": "Show less", + "revoke": "Revoke", + "revokeAria": "Revoke token {{name}}", + "revokeWarning": "Revoke the token \"{{name}}\"? Anything still using it will stop working immediately. This cannot be undone.", + "restrictionCount": { + "agents_one": "{{count}} agent", + "agents_other": "{{count}} agents", + "sources_one": "{{count}} source", + "sources_other": "{{count}} sources", + "prompts_one": "{{count}} prompt", + "prompts_other": "{{count}} prompts", + "tools_one": "{{count}} tool", + "tools_other": "{{count}} tools", + "workflows_one": "{{count}} workflow", + "workflows_other": "{{count}} workflows" + }, + "relative": { + "now": "Just now", + "minutes_one": "{{count}} minute ago", + "minutes_other": "{{count}} minutes ago", + "hours_one": "{{count}} hour ago", + "hours_other": "{{count}} hours ago", + "days_one": "{{count}} day ago", + "days_other": "{{count}} days ago" + }, + "families": { + "agents": "Agents", + "sources": "Sources", + "prompts": "Prompts", + "tools": "Tools", + "models": "Models", + "workflows": "Workflows", + "schedules": "Schedules", + "conversations": "Conversations", + "analytics": "Analytics", + "teams": "Teams", + "chat": "Chat" + }, + "create": { + "title": "Create access token", + "subtitle": "The token acts as you, limited to the scopes and resources you pick here.", + "name": "Name", + "namePlaceholder": "e.g. CI pipeline", + "expiration": "Expiration", + "expiryDays_one": "{{count}} day", + "expiryDays_other": "{{count}} days", + "noExpiration": "No expiration", + "noExpirationHint": "Tokens that never expire are riskier if leaked. Prefer an expiration date.", + "expiresOn": "Expires on {{date}}", + "scopes": "Scopes", + "scopesHint": "Choose what this token may do. Write access includes read access.", + "impliedByWrite": "Included with write access", + "restrict": "Restrict to specific resources", + "restrictHint": "Optional. Limit the token to the items you select. Leave a list empty to allow all of them.", + "restrictNoFamilies": "Select a scope for agents, sources, prompts, tools or chat first to restrict its resources.", + "allSelected": "All (no restriction)", + "noResources": "Nothing found.", + "searchResources": "Search...", + "resourcesError": "Failed to load the list.", + "retry": "Retry", + "cancel": "Cancel", + "submit": "Create token", + "submitting": "Creating", + "error": "Failed to create the token. Please try again." + }, + "created": { + "title": "Token created", + "subtitle": "Your new token \"{{name}}\" is ready.", + "warning": "Copy the token now and store it somewhere safe. For security reasons it won't be shown again.", + "tokenLabel": "Your token", + "usageTitle": "How to use it", + "usageHint": "Export it as an environment variable, then send it as a Bearer token in the Authorization header.", + "done": "Done" + } + }, "scrollTabsLeft": "Scroll tabs left", "tabsAriaLabel": "Settings tabs", "scrollTabsRight": "Scroll tabs right" diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json index c7b4f794..c6f7fc92 100644 --- a/frontend/src/locale/es.json +++ b/frontend/src/locale/es.json @@ -321,16 +321,6 @@ } } }, - "apiKeys": { - "label": "Chatbots", - "name": "Nombre", - "key": "Clave de API", - "sourceDoc": "Documento Fuente", - "createNew": "Crear Nuevo", - "noData": "No hay chatbots existentes", - "deleteConfirmation": "¿Estás seguro de que quieres eliminar la clave API '{{name}}'?", - "description": "Aquí puede crear y gestionar sus chatbots. Los chatbots se pueden implementar en sitios web como widgets o utilizarse dentro de sus aplicaciones." - }, "analytics": { "label": "Analítica", "subtitle": "Realiza un seguimiento del volumen de mensajes, uso de tokens y comentarios de usuarios en tu cuenta", @@ -747,6 +737,104 @@ "testFailed": "La prueba de conexión falló" } }, + "accessTokens": { + "label": "Tokens de acceso", + "subtitle": "Los tokens de acceso personal permiten que scripts, la CLI y pipelines de CI llamen a la API de DocsGPT en tu nombre. Da a cada token solo los permisos que necesite.", + "createToken": "Crear token", + "disabledNotice": "Los tokens de acceso personal no están disponibles en este servidor. Requieren inicio de sesión con cuenta y deben ser habilitados por tu administrador.", + "limitReached_one": "Has alcanzado el límite de {{count}} token activo. Revoca uno para crear otro.", + "limitReached_other": "Has alcanzado el límite de {{count}} tokens activos. Revoca uno para crear otro.", + "loadError": "No se pudieron cargar los tokens de acceso. Inténtalo de nuevo.", + "revokeError": "No se pudo revocar el token. Inténtalo de nuevo.", + "empty": "Aún no hay tokens de acceso", + "emptyHint": "Crea un token para usar la API de DocsGPT desde scripts y herramientas de línea de comandos.", + "name": "Nombre", + "scopes": "Permisos", + "resources": "Recursos", + "createdAt": "Creado", + "lastUsed": "Último uso", + "expires": "Caduca", + "actions": "Acciones", + "allResources": "Todos los recursos", + "never": "Nunca", + "expired": "Caducó el {{date}}", + "expiresSoon": "Caduca el {{date}}", + "moreScopes": "+{{count}} más", + "showLess": "Mostrar menos", + "revoke": "Revocar", + "revokeAria": "Revocar el token {{name}}", + "revokeWarning": "¿Revocar el token \"{{name}}\"? Todo lo que aún lo use dejará de funcionar de inmediato. Esta acción no se puede deshacer.", + "restrictionCount": { + "agents_one": "{{count}} agente", + "agents_other": "{{count}} agentes", + "sources_one": "{{count}} fuente", + "sources_other": "{{count}} fuentes", + "prompts_one": "{{count}} prompt", + "prompts_other": "{{count}} prompts", + "tools_one": "{{count}} herramienta", + "tools_other": "{{count}} herramientas", + "workflows_one": "{{count}} flujo de trabajo", + "workflows_other": "{{count}} flujos de trabajo" + }, + "relative": { + "now": "Ahora mismo", + "minutes_one": "Hace {{count}} minuto", + "minutes_other": "Hace {{count}} minutos", + "hours_one": "Hace {{count}} hora", + "hours_other": "Hace {{count}} horas", + "days_one": "Hace {{count}} día", + "days_other": "Hace {{count}} días" + }, + "families": { + "agents": "Agentes", + "sources": "Fuentes", + "prompts": "Prompts", + "tools": "Herramientas", + "models": "Modelos", + "workflows": "Flujos de trabajo", + "schedules": "Programaciones", + "conversations": "Conversaciones", + "analytics": "Analíticas", + "teams": "Equipos", + "chat": "Chat" + }, + "create": { + "title": "Crear token de acceso", + "subtitle": "El token actúa en tu nombre, limitado a los permisos y recursos que elijas aquí.", + "name": "Nombre", + "namePlaceholder": "p. ej. Pipeline de CI", + "expiration": "Caducidad", + "expiryDays_one": "{{count}} día", + "expiryDays_other": "{{count}} días", + "noExpiration": "Sin caducidad", + "noExpirationHint": "Los tokens que nunca caducan son más peligrosos si se filtran. Es preferible fijar una fecha de caducidad.", + "expiresOn": "Caduca el {{date}}", + "scopes": "Permisos", + "scopesHint": "Elige qué puede hacer este token. El acceso de escritura incluye el de lectura.", + "impliedByWrite": "Incluido con el acceso de escritura", + "restrict": "Restringir a recursos específicos", + "restrictHint": "Opcional. Limita el token a los elementos que selecciones. Deja una lista vacía para permitirlos todos.", + "restrictNoFamilies": "Selecciona primero un permiso de agentes, fuentes, prompts, herramientas o chat para restringir sus recursos.", + "allSelected": "Todos (sin restricción)", + "noResources": "No se encontró nada.", + "searchResources": "Buscar...", + "resourcesError": "No se pudo cargar la lista.", + "retry": "Reintentar", + "cancel": "Cancelar", + "submit": "Crear token", + "submitting": "Creando", + "error": "No se pudo crear el token. Inténtalo de nuevo." + }, + "created": { + "title": "Token creado", + "subtitle": "Tu nuevo token \"{{name}}\" está listo.", + "warning": "Copia el token ahora y guárdalo en un lugar seguro. Por motivos de seguridad no se volverá a mostrar.", + "tokenLabel": "Tu token", + "usageTitle": "Cómo usarlo", + "usageHint": "Expórtalo como variable de entorno y envíalo como token Bearer en la cabecera Authorization.", + "done": "Listo" + } + }, "scrollTabsLeft": "Desplazar pestañas a la izquierda", "tabsAriaLabel": "Pestañas de configuración", "scrollTabsRight": "Desplazar pestañas a la derecha" diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json index dc022fe8..8c0bfc73 100644 --- a/frontend/src/locale/jp.json +++ b/frontend/src/locale/jp.json @@ -321,16 +321,6 @@ } } }, - "apiKeys": { - "label": "チャットボット", - "name": "名前", - "key": "APIキー", - "sourceDoc": "ソースドキュメント", - "createNew": "新規作成", - "noData": "既存のチャットボットはありません", - "deleteConfirmation": "APIキー '{{name}}' を削除してもよろしいですか?", - "description": "ここでチャットボットを作成・管理できます。チャットボットはウィジェットとしてウェブサイトに導入したり、アプリケーション内で使用したりすることができます。" - }, "analytics": { "label": "分析", "subtitle": "アカウント全体のメッセージ量、トークン使用量、ユーザーフィードバックを追跡", @@ -747,6 +737,104 @@ "testFailed": "接続テストに失敗しました" } }, + "accessTokens": { + "label": "アクセストークン", + "subtitle": "個人用アクセストークンを使うと、スクリプト、CLI、CIパイプラインがあなたに代わってDocsGPT APIを呼び出せます。各トークンには必要なスコープだけを付与してください。", + "createToken": "トークンを作成", + "disabledNotice": "このサーバーでは個人用アクセストークンを利用できません。アカウントでのサインインが必要で、管理者が有効にする必要があります。", + "limitReached_one": "有効なトークンが上限の{{count}}個に達しました。新しく作成するには、いずれかを取り消してください。", + "limitReached_other": "有効なトークンが上限の{{count}}個に達しました。新しく作成するには、いずれかを取り消してください。", + "loadError": "アクセストークンを読み込めませんでした。もう一度お試しください。", + "revokeError": "トークンを取り消せませんでした。もう一度お試しください。", + "empty": "アクセストークンはまだありません", + "emptyHint": "トークンを作成すると、スクリプトやコマンドラインツールからDocsGPT APIを利用できます。", + "name": "名前", + "scopes": "スコープ", + "resources": "リソース", + "createdAt": "作成日", + "lastUsed": "最終使用", + "expires": "有効期限", + "actions": "操作", + "allResources": "すべてのリソース", + "never": "なし", + "expired": "{{date}}に期限切れ", + "expiresSoon": "{{date}}に期限切れ予定", + "moreScopes": "他{{count}}件", + "showLess": "折りたたむ", + "revoke": "取り消し", + "revokeAria": "トークン{{name}}を取り消す", + "revokeWarning": "トークン「{{name}}」を取り消しますか?このトークンを使用しているものはすぐに動作しなくなります。この操作は元に戻せません。", + "restrictionCount": { + "agents_one": "エージェント{{count}}件", + "agents_other": "エージェント{{count}}件", + "sources_one": "ソース{{count}}件", + "sources_other": "ソース{{count}}件", + "prompts_one": "プロンプト{{count}}件", + "prompts_other": "プロンプト{{count}}件", + "tools_one": "ツール{{count}}件", + "tools_other": "ツール{{count}}件", + "workflows_one": "ワークフロー{{count}}件", + "workflows_other": "ワークフロー{{count}}件" + }, + "relative": { + "now": "たった今", + "minutes_one": "{{count}}分前", + "minutes_other": "{{count}}分前", + "hours_one": "{{count}}時間前", + "hours_other": "{{count}}時間前", + "days_one": "{{count}}日前", + "days_other": "{{count}}日前" + }, + "families": { + "agents": "エージェント", + "sources": "ソース", + "prompts": "プロンプト", + "tools": "ツール", + "models": "モデル", + "workflows": "ワークフロー", + "schedules": "スケジュール", + "conversations": "会話", + "analytics": "分析", + "teams": "チーム", + "chat": "チャット" + }, + "create": { + "title": "アクセストークンを作成", + "subtitle": "トークンはあなたとして動作し、ここで選択したスコープとリソースに制限されます。", + "name": "名前", + "namePlaceholder": "例: CIパイプライン", + "expiration": "有効期限", + "expiryDays_one": "{{count}}日", + "expiryDays_other": "{{count}}日", + "noExpiration": "無期限", + "noExpirationHint": "無期限のトークンは漏洩した場合のリスクが高くなります。有効期限の設定をおすすめします。", + "expiresOn": "{{date}}に期限切れ", + "scopes": "スコープ", + "scopesHint": "このトークンに許可する操作を選択します。書き込み権限には読み取り権限が含まれます。", + "impliedByWrite": "書き込み権限に含まれます", + "restrict": "特定のリソースに制限する", + "restrictHint": "任意。選択した項目だけにトークンを制限します。リストを空のままにすると、すべて許可されます。", + "restrictNoFamilies": "リソースを制限するには、まずエージェント、ソース、プロンプト、ツール、またはチャットのスコープを選択してください。", + "allSelected": "すべて(制限なし)", + "noResources": "見つかりませんでした。", + "searchResources": "検索...", + "resourcesError": "リストを読み込めませんでした。", + "retry": "再試行", + "cancel": "キャンセル", + "submit": "トークンを作成", + "submitting": "作成中", + "error": "トークンを作成できませんでした。もう一度お試しください。" + }, + "created": { + "title": "トークンを作成しました", + "subtitle": "新しいトークン「{{name}}」の準備ができました。", + "warning": "今すぐトークンをコピーして、安全な場所に保管してください。セキュリティ上の理由から、再度表示されることはありません。", + "tokenLabel": "あなたのトークン", + "usageTitle": "使い方", + "usageHint": "環境変数としてエクスポートし、AuthorizationヘッダーのBearerトークンとして送信します。", + "done": "完了" + } + }, "scrollTabsLeft": "タブを左にスクロール", "tabsAriaLabel": "設定タブ", "scrollTabsRight": "タブを右にスクロール" diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json index dd7aef1e..ff29d056 100644 --- a/frontend/src/locale/ru.json +++ b/frontend/src/locale/ru.json @@ -321,16 +321,6 @@ } } }, - "apiKeys": { - "label": "API ключи", - "name": "Название", - "key": "API ключ", - "sourceDoc": "Источник документа", - "createNew": "Создать новый", - "noData": "Нет существующих чатботов", - "deleteConfirmation": "Вы уверены, что хотите удалить API ключ '{{name}}'?", - "description": "Здесь вы можете создавать и управлять чат-ботами. Чат-боты могут быть развернуты на веб-сайтах в виде виджетов или использоваться внутри ваших приложений." - }, "analytics": { "label": "Аналитика", "subtitle": "Отслеживайте объём сообщений, использование токенов и отзывы пользователей в вашем аккаунте", @@ -747,6 +737,124 @@ "testFailed": "Проверка соединения не удалась" } }, + "accessTokens": { + "label": "Токены доступа", + "subtitle": "Персональные токены доступа позволяют скриптам, CLI и CI-конвейерам обращаться к API DocsGPT от вашего имени. Выдавайте каждому токену только необходимые права.", + "createToken": "Создать токен", + "disabledNotice": "Персональные токены доступа недоступны на этом сервере. Для них нужен вход через учётную запись, и их должен включить администратор.", + "limitReached_one": "Достигнут лимит активных токенов: {{count}}. Отзовите один, чтобы создать новый.", + "limitReached_few": "Достигнут лимит активных токенов: {{count}}. Отзовите один, чтобы создать новый.", + "limitReached_many": "Достигнут лимит активных токенов: {{count}}. Отзовите один, чтобы создать новый.", + "limitReached_other": "Достигнут лимит активных токенов: {{count}}. Отзовите один, чтобы создать новый.", + "loadError": "Не удалось загрузить токены доступа. Попробуйте ещё раз.", + "revokeError": "Не удалось отозвать токен. Попробуйте ещё раз.", + "empty": "Токенов доступа пока нет", + "emptyHint": "Создайте токен, чтобы использовать API DocsGPT из скриптов и инструментов командной строки.", + "name": "Название", + "scopes": "Права", + "resources": "Ресурсы", + "createdAt": "Создан", + "lastUsed": "Последнее использование", + "expires": "Истекает", + "actions": "Действия", + "allResources": "Все ресурсы", + "never": "Никогда", + "expired": "Истёк {{date}}", + "expiresSoon": "Истекает {{date}}", + "moreScopes": "ещё {{count}}", + "showLess": "Свернуть", + "revoke": "Отозвать", + "revokeAria": "Отозвать токен {{name}}", + "revokeWarning": "Отозвать токен «{{name}}»? Всё, что его использует, сразу перестанет работать. Это действие нельзя отменить.", + "restrictionCount": { + "agents_one": "{{count}} агент", + "agents_few": "{{count}} агента", + "agents_many": "{{count}} агентов", + "agents_other": "{{count}} агента", + "sources_one": "{{count}} источник", + "sources_few": "{{count}} источника", + "sources_many": "{{count}} источников", + "sources_other": "{{count}} источника", + "prompts_one": "{{count}} промпт", + "prompts_few": "{{count}} промпта", + "prompts_many": "{{count}} промптов", + "prompts_other": "{{count}} промпта", + "tools_one": "{{count}} инструмент", + "tools_few": "{{count}} инструмента", + "tools_many": "{{count}} инструментов", + "tools_other": "{{count}} инструмента", + "workflows_one": "{{count}} рабочий процесс", + "workflows_few": "{{count}} рабочих процесса", + "workflows_many": "{{count}} рабочих процессов", + "workflows_other": "{{count}} рабочего процесса" + }, + "relative": { + "now": "Только что", + "minutes_one": "{{count}} минуту назад", + "minutes_few": "{{count}} минуты назад", + "minutes_many": "{{count}} минут назад", + "minutes_other": "{{count}} минуты назад", + "hours_one": "{{count}} час назад", + "hours_few": "{{count}} часа назад", + "hours_many": "{{count}} часов назад", + "hours_other": "{{count}} часа назад", + "days_one": "{{count}} день назад", + "days_few": "{{count}} дня назад", + "days_many": "{{count}} дней назад", + "days_other": "{{count}} дня назад" + }, + "families": { + "agents": "Агенты", + "sources": "Источники", + "prompts": "Промпты", + "tools": "Инструменты", + "models": "Модели", + "workflows": "Рабочие процессы", + "schedules": "Расписания", + "conversations": "Диалоги", + "analytics": "Аналитика", + "teams": "Команды", + "chat": "Чат" + }, + "create": { + "title": "Создать токен доступа", + "subtitle": "Токен действует от вашего имени в пределах прав и ресурсов, которые вы выберете здесь.", + "name": "Название", + "namePlaceholder": "например, CI-конвейер", + "expiration": "Срок действия", + "expiryDays_one": "{{count}} день", + "expiryDays_few": "{{count}} дня", + "expiryDays_many": "{{count}} дней", + "expiryDays_other": "{{count}} дня", + "noExpiration": "Бессрочно", + "noExpirationHint": "Бессрочные токены опаснее при утечке. Лучше задать срок действия.", + "expiresOn": "Истекает {{date}}", + "scopes": "Права", + "scopesHint": "Выберите, что разрешено этому токену. Право на запись включает право на чтение.", + "impliedByWrite": "Входит в право на запись", + "restrict": "Ограничить определёнными ресурсами", + "restrictHint": "Необязательно. Ограничьте токен выбранными элементами. Оставьте список пустым, чтобы разрешить все.", + "restrictNoFamilies": "Сначала выберите право для агентов, источников, промптов, инструментов или чата, чтобы ограничить их ресурсы.", + "allSelected": "Все (без ограничений)", + "noResources": "Ничего не найдено.", + "searchResources": "Поиск...", + "resourcesError": "Не удалось загрузить список.", + "retry": "Повторить", + "cancel": "Отмена", + "submit": "Создать токен", + "submitting": "Создание", + "error": "Не удалось создать токен. Попробуйте ещё раз." + }, + "created": { + "title": "Токен создан", + "subtitle": "Ваш новый токен «{{name}}» готов.", + "warning": "Скопируйте токен сейчас и сохраните его в надёжном месте. В целях безопасности он больше не будет показан.", + "tokenLabel": "Ваш токен", + "usageTitle": "Как использовать", + "usageHint": "Экспортируйте его как переменную окружения и передавайте как Bearer-токен в заголовке Authorization.", + "done": "Готово" + } + }, "scrollTabsLeft": "Прокрутить вкладки влево", "tabsAriaLabel": "Вкладки настроек", "scrollTabsRight": "Прокрутить вкладки вправо" diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json index 16236cec..1952b6e5 100644 --- a/frontend/src/locale/zh-TW.json +++ b/frontend/src/locale/zh-TW.json @@ -321,16 +321,6 @@ } } }, - "apiKeys": { - "label": "聊天機器人", - "name": "名稱", - "key": "API 金鑰", - "sourceDoc": "來源文件", - "createNew": "建立新的", - "noData": "沒有現有的聊天機器人", - "deleteConfirmation": "您確定要刪除 API 金鑰 '{{name}}' 嗎?", - "description": "在這裡,您可以創建和管理您的聊天機器人。聊天機器人可以作為小部件部署到網站上,或在您的應用程序中使用。" - }, "analytics": { "label": "分析", "subtitle": "追蹤您賃戶的訊息量、Token 使用量和使用者回饋", @@ -747,6 +737,104 @@ "testFailed": "連線測試失敗" } }, + "accessTokens": { + "label": "存取權杖", + "subtitle": "個人存取權杖可讓指令碼、CLI 和 CI 管線以您的身分呼叫 DocsGPT API。請只授予每個權杖所需的權限範圍。", + "createToken": "建立權杖", + "disabledNotice": "此伺服器不支援個人存取權杖。此功能需要使用帳戶登入,並須由管理員啟用。", + "limitReached_one": "您的有效權杖已達 {{count}} 個的上限。請先撤銷一個再建立新的權杖。", + "limitReached_other": "您的有效權杖已達 {{count}} 個的上限。請先撤銷一個再建立新的權杖。", + "loadError": "載入存取權杖失敗,請再試一次。", + "revokeError": "撤銷權杖失敗,請再試一次。", + "empty": "尚無存取權杖", + "emptyHint": "建立權杖後,即可在指令碼和命令列工具中使用 DocsGPT API。", + "name": "名稱", + "scopes": "權限範圍", + "resources": "資源", + "createdAt": "建立時間", + "lastUsed": "上次使用", + "expires": "到期時間", + "actions": "操作", + "allResources": "所有資源", + "never": "從不", + "expired": "已於 {{date}} 到期", + "expiresSoon": "將於 {{date}} 到期", + "moreScopes": "另外 {{count}} 項", + "showLess": "收合", + "revoke": "撤銷", + "revokeAria": "撤銷權杖 {{name}}", + "revokeWarning": "確定要撤銷權杖「{{name}}」嗎?所有仍在使用它的程式將立即失效,且此操作無法復原。", + "restrictionCount": { + "agents_one": "{{count}} 個代理", + "agents_other": "{{count}} 個代理", + "sources_one": "{{count}} 個來源", + "sources_other": "{{count}} 個來源", + "prompts_one": "{{count}} 個提示詞", + "prompts_other": "{{count}} 個提示詞", + "tools_one": "{{count}} 個工具", + "tools_other": "{{count}} 個工具", + "workflows_one": "{{count}} 個工作流程", + "workflows_other": "{{count}} 個工作流程" + }, + "relative": { + "now": "剛剛", + "minutes_one": "{{count}} 分鐘前", + "minutes_other": "{{count}} 分鐘前", + "hours_one": "{{count}} 小時前", + "hours_other": "{{count}} 小時前", + "days_one": "{{count}} 天前", + "days_other": "{{count}} 天前" + }, + "families": { + "agents": "代理", + "sources": "來源", + "prompts": "提示詞", + "tools": "工具", + "models": "模型", + "workflows": "工作流程", + "schedules": "排程", + "conversations": "對話", + "analytics": "分析", + "teams": "團隊", + "chat": "聊天" + }, + "create": { + "title": "建立存取權杖", + "subtitle": "權杖將以您的身分操作,但僅限於您在此選擇的權限範圍和資源。", + "name": "名稱", + "namePlaceholder": "例如:CI 管線", + "expiration": "有效期限", + "expiryDays_one": "{{count}} 天", + "expiryDays_other": "{{count}} 天", + "noExpiration": "永不到期", + "noExpirationHint": "永不到期的權杖一旦外洩風險更高,建議設定到期時間。", + "expiresOn": "將於 {{date}} 到期", + "scopes": "權限範圍", + "scopesHint": "選擇此權杖可以執行的操作。寫入權限包含讀取權限。", + "impliedByWrite": "已包含在寫入權限中", + "restrict": "限制為特定資源", + "restrictHint": "選填。將權杖限制為您選擇的項目。清單留空表示允許全部。", + "restrictNoFamilies": "請先選擇代理、來源、提示詞、工具或聊天的權限範圍,才能限制其資源。", + "allSelected": "全部(不限制)", + "noResources": "找不到任何內容。", + "searchResources": "搜尋...", + "resourcesError": "載入清單失敗。", + "retry": "重試", + "cancel": "取消", + "submit": "建立權杖", + "submitting": "建立中", + "error": "建立權杖失敗,請再試一次。" + }, + "created": { + "title": "權杖已建立", + "subtitle": "您的新權杖「{{name}}」已就緒。", + "warning": "請立即複製權杖並妥善保存。基於安全考量,它不會再次顯示。", + "tokenLabel": "您的權杖", + "usageTitle": "使用方式", + "usageHint": "將其匯出為環境變數,然後在 Authorization 標頭中作為 Bearer 權杖傳送。", + "done": "完成" + } + }, "scrollTabsLeft": "向左捲動標籤", "tabsAriaLabel": "設定標籤", "scrollTabsRight": "向右捲動標籤" diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json index 5b4cf5a6..74316ec5 100644 --- a/frontend/src/locale/zh.json +++ b/frontend/src/locale/zh.json @@ -321,16 +321,6 @@ } } }, - "apiKeys": { - "label": "聊天机器人", - "name": "名称", - "key": "API 密钥", - "sourceDoc": "源文档", - "createNew": "创建新的", - "noData": "没有现有的聊天机器人", - "deleteConfirmation": "您确定要删除 API 密钥 '{{name}}' 吗?", - "description": "在这里,您可以创建和管理您的聊天机器人。聊天机器人可以作为小部件部署到网站上,或在您的应用程序中使用。" - }, "analytics": { "label": "分析", "subtitle": "跟踪您账户的消息量、令牌使用情况和用户反馈", @@ -747,6 +737,104 @@ "testFailed": "连接测试失败" } }, + "accessTokens": { + "label": "访问令牌", + "subtitle": "个人访问令牌可让脚本、CLI 和 CI 流水线以您的身份调用 DocsGPT API。请只为每个令牌授予其所需的权限范围。", + "createToken": "创建令牌", + "disabledNotice": "此服务器不支持个人访问令牌。该功能需要使用账户登录,并须由管理员启用。", + "limitReached_one": "您的有效令牌已达到 {{count}} 个的上限。请先撤销一个再创建新的令牌。", + "limitReached_other": "您的有效令牌已达到 {{count}} 个的上限。请先撤销一个再创建新的令牌。", + "loadError": "加载访问令牌失败,请重试。", + "revokeError": "撤销令牌失败,请重试。", + "empty": "暂无访问令牌", + "emptyHint": "创建令牌后,即可在脚本和命令行工具中使用 DocsGPT API。", + "name": "名称", + "scopes": "权限范围", + "resources": "资源", + "createdAt": "创建时间", + "lastUsed": "上次使用", + "expires": "过期时间", + "actions": "操作", + "allResources": "所有资源", + "never": "从不", + "expired": "已于 {{date}} 过期", + "expiresSoon": "将于 {{date}} 过期", + "moreScopes": "另外 {{count}} 项", + "showLess": "收起", + "revoke": "撤销", + "revokeAria": "撤销令牌 {{name}}", + "revokeWarning": "确定要撤销令牌“{{name}}”吗?所有仍在使用它的程序将立即失效,且此操作无法撤回。", + "restrictionCount": { + "agents_one": "{{count}} 个代理", + "agents_other": "{{count}} 个代理", + "sources_one": "{{count}} 个来源", + "sources_other": "{{count}} 个来源", + "prompts_one": "{{count}} 个提示词", + "prompts_other": "{{count}} 个提示词", + "tools_one": "{{count}} 个工具", + "tools_other": "{{count}} 个工具", + "workflows_one": "{{count}} 个工作流", + "workflows_other": "{{count}} 个工作流" + }, + "relative": { + "now": "刚刚", + "minutes_one": "{{count}} 分钟前", + "minutes_other": "{{count}} 分钟前", + "hours_one": "{{count}} 小时前", + "hours_other": "{{count}} 小时前", + "days_one": "{{count}} 天前", + "days_other": "{{count}} 天前" + }, + "families": { + "agents": "代理", + "sources": "来源", + "prompts": "提示词", + "tools": "工具", + "models": "模型", + "workflows": "工作流", + "schedules": "计划任务", + "conversations": "对话", + "analytics": "分析", + "teams": "团队", + "chat": "聊天" + }, + "create": { + "title": "创建访问令牌", + "subtitle": "令牌将以您的身份操作,但仅限于您在此选择的权限范围和资源。", + "name": "名称", + "namePlaceholder": "例如:CI 流水线", + "expiration": "有效期", + "expiryDays_one": "{{count}} 天", + "expiryDays_other": "{{count}} 天", + "noExpiration": "永不过期", + "noExpirationHint": "永不过期的令牌一旦泄露风险更高,建议设置过期时间。", + "expiresOn": "将于 {{date}} 过期", + "scopes": "权限范围", + "scopesHint": "选择此令牌可以执行的操作。写入权限包含读取权限。", + "impliedByWrite": "已包含在写入权限中", + "restrict": "限制为特定资源", + "restrictHint": "可选。将令牌限制为您选择的项目。列表留空表示允许全部。", + "restrictNoFamilies": "请先选择代理、来源、提示词、工具或聊天的权限范围,才能限制其资源。", + "allSelected": "全部(不限制)", + "noResources": "未找到任何内容。", + "searchResources": "搜索...", + "resourcesError": "加载列表失败。", + "retry": "重试", + "cancel": "取消", + "submit": "创建令牌", + "submitting": "正在创建", + "error": "创建令牌失败,请重试。" + }, + "created": { + "title": "令牌已创建", + "subtitle": "您的新令牌“{{name}}”已就绪。", + "warning": "请立即复制令牌并妥善保存。出于安全考虑,它不会再次显示。", + "tokenLabel": "您的令牌", + "usageTitle": "使用方法", + "usageHint": "将其导出为环境变量,然后在 Authorization 请求头中作为 Bearer 令牌发送。", + "done": "完成" + } + }, "scrollTabsLeft": "向左滚动标签", "tabsAriaLabel": "设置标签", "scrollTabsRight": "向右滚动标签" diff --git a/frontend/src/modals/AccessTokenCreatedModal.tsx b/frontend/src/modals/AccessTokenCreatedModal.tsx new file mode 100644 index 00000000..66026edc --- /dev/null +++ b/frontend/src/modals/AccessTokenCreatedModal.tsx @@ -0,0 +1,104 @@ +import { TriangleAlert } from 'lucide-react'; +import { useTranslation } from 'react-i18next'; + +import { baseURL } from '../api/client'; +import CopyButton from '../components/CopyButton'; +import { Button } from '../components/ui/button'; +import { Modal } from '../components/ui/modal'; + +interface AccessTokenCreatedModalProps { + /** Plaintext secret; `null` keeps the modal closed. Held only in the parent's component state. */ + token: string | null; + name: string; + onClose: () => void; +} + +export default function AccessTokenCreatedModal({ + token, + name, + onClose, +}: AccessTokenCreatedModalProps) { + const { t } = useTranslation(); + const exportSnippet = `export DOCSGPT_TOKEN=${token ?? ''}`; + const curlSnippet = `curl -H "Authorization: Bearer $DOCSGPT_TOKEN" ${baseURL}/api/get_agents`; + + return ( + !o && onClose()} + hideTitle + title={t('settings.accessTokens.created.title')} + size="lg" + mobileVariant="sheet" + // The secret cannot be shown again, so a stray click outside must not + // dismiss it; closing takes the explicit button (or Esc). + isPerformingTask + footer={ + + } + > +
+
+

+ {t('settings.accessTokens.created.title')} +

+

+ {t('settings.accessTokens.created.subtitle', { name })} +

+
+ +
+
+ +
+

+ {t('settings.accessTokens.created.tokenLabel')} +

+
+ + {token} + + {token && } +
+
+ +
+

+ {t('settings.accessTokens.created.usageTitle')} +

+

+ {t('settings.accessTokens.created.usageHint')} +

+ {[exportSnippet, curlSnippet].map((snippet, index) => ( +
+
+                {snippet}
+              
+ +
+ ))} +
+
+
+ ); +} diff --git a/frontend/src/modals/CreateAccessTokenModal.tsx b/frontend/src/modals/CreateAccessTokenModal.tsx new file mode 100644 index 00000000..438403a8 --- /dev/null +++ b/frontend/src/modals/CreateAccessTokenModal.tsx @@ -0,0 +1,462 @@ +import React from 'react'; +import { useTranslation } from 'react-i18next'; +import { useSelector } from 'react-redux'; + +import patService, { + AccessTokenPolicy, + AccessTokenScope, + CreateAccessTokenResponse, +} from '../api/services/patService'; +import userService from '../api/services/userService'; +import Spinner from '../components/Spinner'; +import { Button } from '../components/ui/button'; +import { Input } from '../components/ui/input'; +import { Label } from '../components/ui/label'; +import { Modal } from '../components/ui/modal'; +import { MultiSelect } from '../components/ui/multi-select'; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from '../components/ui/select'; +import { Switch } from '../components/ui/switch'; +import { selectToken } from '../preferences/preferenceSlice'; +import { + buildResourceFilter, + defaultExpiry, + eligibleFilterFamilies, + expiryOptions, + groupScopesByFamily, + isScopeImplied, + NO_EXPIRY, + PICKER_FAMILIES, + ResourceOption, + scopesToSubmit, + toResourceOptions, +} from '../settings/accessTokenUtils'; +import { formatDateOnly } from '../utils/dateTimeUtils'; + +const MAX_NAME_LENGTH = 100; +const DAY_MS = 24 * 60 * 60 * 1000; + +type ResourceState = + | { status: 'loading' } + | { status: 'error' } + | { status: 'ready'; options: ResourceOption[] }; + +const RESOURCE_FETCHERS: Record< + string, + (token: string | null) => Promise +> = { + agents: (token) => userService.getAgents(token), + sources: (token) => userService.getDocs(token), + prompts: (token) => userService.getPrompts(token), + tools: (token) => userService.getUserTools(token), +}; + +interface CreateAccessTokenModalProps { + open: boolean; + onClose: () => void; + scopes: AccessTokenScope[]; + policy: AccessTokenPolicy; + onCreated: (created: CreateAccessTokenResponse) => void; +} + +export default function CreateAccessTokenModal({ + open, + onClose, + scopes, + policy, + onCreated, +}: CreateAccessTokenModalProps) { + const { t } = useTranslation(); + const token = useSelector(selectToken); + + const [name, setName] = React.useState(''); + const [selectedScopes, setSelectedScopes] = React.useState([]); + const [expiry, setExpiry] = React.useState(() => + defaultExpiry(policy), + ); + const [restrict, setRestrict] = React.useState(false); + const [resourceSelection, setResourceSelection] = React.useState< + Record + >({}); + const [resources, setResources] = React.useState< + Record + >({}); + const [submitting, setSubmitting] = React.useState(false); + const [error, setError] = React.useState(null); + + // Start from a clean form on every open; resource lists are refetched so a + // freshly created agent/source shows up without a page reload. + React.useEffect(() => { + if (!open) return; + setName(''); + setSelectedScopes([]); + setExpiry(defaultExpiry(policy)); + setRestrict(false); + setResourceSelection({}); + setResources({}); + setSubmitting(false); + setError(null); + }, [open, policy]); + + const scopeGroups = React.useMemo( + () => groupScopesByFamily(scopes), + [scopes], + ); + const expiryChoices = React.useMemo(() => expiryOptions(policy), [policy]); + const pickerFamilies = React.useMemo( + () => + eligibleFilterFamilies(selectedScopes, policy.filterable_families).filter( + (family) => PICKER_FAMILIES.includes(family), + ), + [selectedScopes, policy.filterable_families], + ); + + const loadResources = React.useCallback( + (family: string) => { + const fetcher = RESOURCE_FETCHERS[family]; + if (!fetcher) return; + setResources((prev) => ({ ...prev, [family]: { status: 'loading' } })); + fetcher(token) + .then(async (response) => { + if (!response.ok) throw new Error(`HTTP ${response.status}`); + const options = toResourceOptions(family, await response.json()); + setResources((prev) => ({ + ...prev, + [family]: { status: 'ready', options }, + })); + }) + .catch((err) => { + console.error(`Failed to load ${family}:`, err); + setResources((prev) => ({ ...prev, [family]: { status: 'error' } })); + }); + }, + [token], + ); + + // Lazy: a family's list is only fetched once its picker is actually shown. + React.useEffect(() => { + if (!open || !restrict) return; + pickerFamilies.forEach((family) => { + if (!resources[family]) loadResources(family); + }); + }, [open, restrict, pickerFamilies, resources, loadResources]); + + const toggleScope = (scope: string) => { + setError(null); + setSelectedScopes((prev) => + prev.includes(scope) ? prev.filter((s) => s !== scope) : [...prev, scope], + ); + }; + + const familyLabel = (family: string) => + t(`settings.accessTokens.families.${family}`, { defaultValue: family }); + + const trimmedName = name.trim(); + const canSubmit = + trimmedName.length > 0 && selectedScopes.length > 0 && !submitting; + + const handleClose = () => { + if (submitting) return; + onClose(); + }; + + const handleSubmit = async () => { + if (!canSubmit) return; + setSubmitting(true); + setError(null); + try { + const created = await patService.create( + { + name: trimmedName, + scopes: scopesToSubmit(selectedScopes, scopes), + resource_filter: restrict + ? buildResourceFilter(resourceSelection, pickerFamilies) + : undefined, + expires_in_days: expiry, + }, + token, + ); + onCreated(created); + } catch (err) { + const message = err instanceof Error ? err.message : ''; + setError(message || t('settings.accessTokens.create.error')); + } finally { + setSubmitting(false); + } + }; + + const expiryLabel = (days: number) => + days === NO_EXPIRY + ? t('settings.accessTokens.create.noExpiration') + : t('settings.accessTokens.create.expiryDays', { count: days }); + + return ( + !o && handleClose()} + hideTitle + title={t('settings.accessTokens.create.title')} + size="lg" + mobileVariant="sheet" + isPerformingTask={submitting} + contentClassName="max-h-[65vh]" + footer={ + <> + + + + } + > +
{ + e.preventDefault(); + handleSubmit(); + }} + > +
+

+ {t('settings.accessTokens.create.title')} +

+

+ {t('settings.accessTokens.create.subtitle')} +

+
+ +
+
+ + { + setName(e.target.value); + setError(null); + }} + placeholder={t('settings.accessTokens.create.namePlaceholder')} + className="rounded-xl" + autoComplete="off" + /> +
+
+ + +

+ {expiry === NO_EXPIRY + ? t('settings.accessTokens.create.noExpirationHint') + : t('settings.accessTokens.create.expiresOn', { + date: formatDateOnly( + new Date(Date.now() + expiry * DAY_MS).toISOString(), + ), + })} +

+
+
+ +
+ + {t('settings.accessTokens.create.scopes')} + * + +

+ {t('settings.accessTokens.create.scopesHint')} +

+
+ {scopeGroups.map((group) => ( +
+

+ {familyLabel(group.family)} +

+
+ {group.scopes.map((scope) => { + const implied = isScopeImplied(scope.name, selectedScopes); + const checked = + implied || selectedScopes.includes(scope.name); + const id = `pat-scope-${scope.name}`; + return ( + + ); + })} +
+
+ ))} +
+
+ + {policy.filterable_families.length > 0 && ( +
+
+
+ +

+ {t('settings.accessTokens.create.restrictHint')} +

+
+ +
+ {restrict && pickerFamilies.length === 0 && ( +

+ {t('settings.accessTokens.create.restrictNoFamilies')} +

+ )} + {restrict && + pickerFamilies.map((family) => { + const state = resources[family]; + return ( +
+ + {!state || state.status === 'loading' ? ( +
+ +
+ ) : state.status === 'error' ? ( +
+ + {t('settings.accessTokens.create.resourcesError')} + + +
+ ) : ( + + setResourceSelection((prev) => ({ + ...prev, + [family]: ids, + })) + } + placeholder={t( + 'settings.accessTokens.create.allSelected', + )} + emptyText={t( + 'settings.accessTokens.create.noResources', + )} + searchPlaceholder={t( + 'settings.accessTokens.create.searchResources', + )} + className="rounded-xl" + /> + )} +
+ ); + })} +
+ )} + + {error && ( +
+ {error} +
+ )} +
+
+ ); +} diff --git a/frontend/src/settings/PersonalAccessTokens.tsx b/frontend/src/settings/PersonalAccessTokens.tsx new file mode 100644 index 00000000..d81416b0 --- /dev/null +++ b/frontend/src/settings/PersonalAccessTokens.tsx @@ -0,0 +1,443 @@ +import { TriangleAlert } from 'lucide-react'; +import React from 'react'; +import { useTranslation } from 'react-i18next'; +import { useSelector } from 'react-redux'; + +import patService, { + AccessTokenPolicy, + AccessTokenScope, + CreateAccessTokenResponse, + PersonalAccessToken, +} from '../api/services/patService'; +import NoFilesDarkIcon from '../assets/no-files-dark.svg'; +import NoFilesIcon from '../assets/no-files.svg'; +import SkeletonLoader from '../components/SkeletonLoader'; +import { Alert, AlertDescription } from '../components/ui/alert'; +import { Button } from '../components/ui/button'; +import { + Table, + TableBody, + TableCell, + TableContainer, + TableHead, + TableHeader, + TableRow, +} from '../components/ui/table'; +import { useDarkTheme } from '../hooks'; +import AccessTokenCreatedModal from '../modals/AccessTokenCreatedModal'; +import ConfirmationModal from '../modals/ConfirmationModal'; +import CreateAccessTokenModal from '../modals/CreateAccessTokenModal'; +import { ActiveState } from '../models/misc'; +import { selectToken } from '../preferences/preferenceSlice'; +import { formatDateOnly, formatDateTime } from '../utils/dateTimeUtils'; +import { + countLiveTokens, + expiryStatus, + relativeTime, + restrictionCounts, +} from './accessTokenUtils'; + +const VISIBLE_SCOPES = 3; + +function ScopeChips({ scopes }: { scopes: string[] }) { + const { t } = useTranslation(); + const [expanded, setExpanded] = React.useState(false); + const visible = expanded ? scopes : scopes.slice(0, VISIBLE_SCOPES); + const hidden = scopes.length - visible.length; + + return ( +
+ {visible.map((scope) => ( + + {scope} + + ))} + {scopes.length > VISIBLE_SCOPES && ( + + )} +
+ ); +} + +export default function PersonalAccessTokens() { + const { t } = useTranslation(); + const token = useSelector(selectToken); + const [isDarkTheme] = useDarkTheme(); + + const [tokens, setTokens] = React.useState([]); + const [scopes, setScopes] = React.useState([]); + const [policy, setPolicy] = React.useState(null); + const [loading, setLoading] = React.useState(true); + const [error, setError] = React.useState(null); + + const [createOpen, setCreateOpen] = React.useState(false); + // The plaintext secret lives only here, and only until the modal closes. + const [created, setCreated] = + React.useState(null); + const [revokeState, setRevokeState] = React.useState('INACTIVE'); + const [tokenToRevoke, setTokenToRevoke] = + React.useState(null); + + const loadTokens = React.useCallback( + async (showLoader: boolean) => { + if (showLoader) setLoading(true); + try { + const data = await patService.list(token); + setTokens(data.tokens ?? []); + setScopes(data.scopes ?? []); + setPolicy(data.policy ?? null); + setError(null); + } catch (err) { + console.error('Failed to load access tokens:', err); + setError(t('settings.accessTokens.loadError')); + } finally { + setLoading(false); + } + }, + [token, t], + ); + + React.useEffect(() => { + loadTokens(true); + }, [loadTokens]); + + const handleCreated = (response: CreateAccessTokenResponse) => { + setCreateOpen(false); + setCreated(response); + setTokens((prev) => [response.personal_access_token, ...prev]); + }; + + const requestRevoke = (item: PersonalAccessToken) => { + setTokenToRevoke(item); + setRevokeState('ACTIVE'); + }; + + const confirmRevoke = async () => { + if (!tokenToRevoke) return; + const target = tokenToRevoke; + try { + await patService.revoke(target.id, token); + setTokens((prev) => prev.filter((item) => item.id !== target.id)); + setError(null); + } catch (err) { + console.error('Failed to revoke access token:', err); + setError( + (err instanceof Error && err.message) || + t('settings.accessTokens.revokeError'), + ); + } finally { + setTokenToRevoke(null); + } + }; + + const limitReached = + !!policy && countLiveTokens(tokens) >= policy.max_per_user; + + const renderRestrictions = (item: PersonalAccessToken) => { + const counts = restrictionCounts(item.resource_filter); + if (counts.length === 0) return t('settings.accessTokens.allResources'); + return counts + .map(({ family, count }) => + t(`settings.accessTokens.restrictionCount.${family}`, { + count, + defaultValue: `${count} ${family}`, + }), + ) + .join(', '); + }; + + const renderLastUsed = (item: PersonalAccessToken) => { + const relative = relativeTime(item.last_used_at); + if (!relative || !item.last_used_at) { + return t('settings.accessTokens.never'); + } + const label = + relative.unit === 'date' + ? formatDateOnly(item.last_used_at) + : relative.unit === 'now' + ? t('settings.accessTokens.relative.now') + : t(`settings.accessTokens.relative.${relative.unit}`, { + count: relative.count, + }); + const details = [formatDateTime(item.last_used_at), item.last_used_ip] + .filter(Boolean) + .join(' · '); + return {label}; + }; + + const renderExpiry = (item: PersonalAccessToken) => { + const status = expiryStatus(item.expires_at); + if (status === 'never' || !item.expires_at) { + return t('settings.accessTokens.never'); + } + const date = formatDateOnly(item.expires_at); + if (status === 'ok') return date; + const expired = status === 'expired'; + return ( + + + ); + }; + + const renderRevokeButton = (item: PersonalAccessToken) => ( + + ); + + const renderPrefix = (item: PersonalAccessToken) => ( + + {item.token_prefix}… + + ); + + const renderEmptyState = () => ( +
+ +

+ {t('settings.accessTokens.empty')} +

+ {policy?.enabled && ( +

+ {t('settings.accessTokens.emptyHint')} +

+ )} +
+ ); + + const mobileField = (label: string, value: React.ReactNode) => ( +
+ {label} + + {value} + +
+ ); + + return ( +
+
+
+

+ {t('settings.accessTokens.subtitle')} +

+ {policy?.enabled && ( + + )} +
+ + {policy && !policy.enabled && ( + + + {t('settings.accessTokens.disabledNotice')} + + + )} + {policy?.enabled && limitReached && ( +

+ {t('settings.accessTokens.limitReached', { + count: policy.max_per_user, + })} +

+ )} + {error && ( + + {error} + + )} + +
+ + {loading ? ( + + ) : tokens.length === 0 ? ( + !error && renderEmptyState() + ) : ( + <> + {/* Desktop: table */} + + + + + {t('settings.accessTokens.name')} + + {t('settings.accessTokens.scopes')} + + + {t('settings.accessTokens.resources')} + + + {t('settings.accessTokens.createdAt')} + + + {t('settings.accessTokens.lastUsed')} + + + {t('settings.accessTokens.expires')} + + + + {t('settings.accessTokens.actions')} + + + + + + {tokens.map((item) => ( + + +

+ {item.name} +

+ {renderPrefix(item)} +
+ + + + + {renderRestrictions(item)} + + + {item.created_at + ? formatDateOnly(item.created_at) + : '-'} + + + {renderLastUsed(item)} + + + {renderExpiry(item)} + + + {renderRevokeButton(item)} + +
+ ))} +
+
+
+ + {/* Mobile / tablet: cards */} +
    + {tokens.map((item) => ( +
  • +
    +
    +

    + {item.name} +

    + {renderPrefix(item)} +
    + {renderRevokeButton(item)} +
    + +
    + {mobileField( + t('settings.accessTokens.resources'), + renderRestrictions(item), + )} + {mobileField( + t('settings.accessTokens.createdAt'), + item.created_at ? formatDateOnly(item.created_at) : '-', + )} + {mobileField( + t('settings.accessTokens.lastUsed'), + renderLastUsed(item), + )} + {mobileField( + t('settings.accessTokens.expires'), + renderExpiry(item), + )} +
    +
  • + ))} +
+ + )} +
+ + {policy && ( + setCreateOpen(false)} + scopes={scopes} + policy={policy} + onCreated={handleCreated} + /> + )} + setCreated(null)} + /> + +
+ ); +} diff --git a/frontend/src/settings/accessTokenUtils.test.ts b/frontend/src/settings/accessTokenUtils.test.ts new file mode 100644 index 00000000..f58f95e7 --- /dev/null +++ b/frontend/src/settings/accessTokenUtils.test.ts @@ -0,0 +1,289 @@ +import { describe, expect, it } from 'vitest'; + +import { + buildResourceFilter, + countLiveTokens, + defaultExpiry, + eligibleFilterFamilies, + expiryOptions, + expiryStatus, + groupScopesByFamily, + isScopeImplied, + isUuid, + NO_EXPIRY, + relativeTime, + restrictionCounts, + scopesToSubmit, + toResourceOptions, +} from './accessTokenUtils'; + +const CATALOG = [ + { name: 'agents:read', description: 'View agents' }, + { name: 'agents:write', description: 'Edit agents' }, + { name: 'agents:keys', description: 'Agent keys' }, + { name: 'sources:read', description: 'View sources' }, + { name: 'analytics:read', description: 'View analytics' }, + { name: 'chat:run', description: 'Ask agents' }, +]; +const FAMILIES = ['agents', 'sources', 'prompts', 'tools', 'workflows']; +const POLICY = { + default_lifetime_days: 90, + max_lifetime_days: 365, + allow_non_expiring: false, +}; +const NOW = Date.parse('2026-09-19T12:00:00Z'); +const DAY = 24 * 60 * 60 * 1000; +const ID_A = '11111111-1111-4111-8111-111111111111'; +const ID_B = '22222222-2222-4222-8222-222222222222'; + +describe('groupScopesByFamily', () => { + it('groups by family in catalog order', () => { + const groups = groupScopesByFamily(CATALOG); + expect(groups.map((g) => g.family)).toEqual([ + 'agents', + 'sources', + 'analytics', + 'chat', + ]); + expect(groups[0].scopes.map((s) => s.name)).toEqual([ + 'agents:read', + 'agents:write', + 'agents:keys', + ]); + }); +}); + +describe('scope implication', () => { + it('write implies read of the same family only', () => { + expect(isScopeImplied('agents:read', ['agents:write'])).toBe(true); + expect(isScopeImplied('sources:read', ['agents:write'])).toBe(false); + expect(isScopeImplied('agents:keys', ['agents:write'])).toBe(false); + expect(isScopeImplied('agents:write', ['agents:write'])).toBe(false); + expect(isScopeImplied('agents:read', ['agents:read'])).toBe(false); + }); + + it('scopesToSubmit drops implied reads and keeps catalog order', () => { + expect( + scopesToSubmit( + ['chat:run', 'agents:write', 'agents:read', 'sources:read'], + CATALOG, + ), + ).toEqual(['agents:write', 'sources:read', 'chat:run']); + }); + + it('scopesToSubmit ignores scopes missing from the catalog', () => { + expect(scopesToSubmit(['bogus:read', 'agents:read'], CATALOG)).toEqual([ + 'agents:read', + ]); + }); +}); + +describe('eligibleFilterFamilies', () => { + it('only offers families that have a selected scope', () => { + expect(eligibleFilterFamilies(['sources:read'], FAMILIES)).toEqual([ + 'sources', + ]); + expect(eligibleFilterFamilies(['analytics:read'], FAMILIES)).toEqual([]); + expect(eligibleFilterFamilies([], FAMILIES)).toEqual([]); + }); + + it('chat:run additionally permits agents and sources', () => { + expect(eligibleFilterFamilies(['chat:run'], FAMILIES)).toEqual([ + 'agents', + 'sources', + ]); + }); + + it('respects the families the server says are filterable', () => { + expect(eligibleFilterFamilies(['chat:run'], ['sources'])).toEqual([ + 'sources', + ]); + }); +}); + +describe('buildResourceFilter', () => { + it('drops empty and ineligible families', () => { + expect( + buildResourceFilter({ agents: [ID_A], sources: [], tools: [ID_B] }, [ + 'agents', + 'sources', + ]), + ).toEqual({ agents: [ID_A] }); + }); + + it('returns undefined when nothing is restricted', () => { + expect(buildResourceFilter({ agents: [] }, ['agents'])).toBeUndefined(); + expect(buildResourceFilter({}, [])).toBeUndefined(); + }); +}); + +describe('restrictionCounts', () => { + it('counts ids per family', () => { + expect( + restrictionCounts({ agents: [ID_A, ID_B], sources: [ID_A] }), + ).toEqual([ + { family: 'agents', count: 2 }, + { family: 'sources', count: 1 }, + ]); + }); + + it('is empty ("all resources") for no filter', () => { + expect(restrictionCounts({})).toEqual([]); + expect(restrictionCounts(null)).toEqual([]); + expect(restrictionCounts({ agents: [] })).toEqual([]); + }); +}); + +describe('expiryOptions / defaultExpiry', () => { + it('offers every preset up to the maximum', () => { + expect(expiryOptions(POLICY)).toEqual([7, 30, 60, 90, 180, 365]); + expect(defaultExpiry(POLICY)).toBe(90); + }); + + it('filters presets above max_lifetime_days', () => { + expect(expiryOptions({ ...POLICY, max_lifetime_days: 90 })).toEqual([ + 7, 30, 60, 90, + ]); + }); + + it('adds a non-preset default in sorted position', () => { + const policy = { ...POLICY, default_lifetime_days: 45 }; + expect(expiryOptions(policy)).toEqual([7, 30, 45, 60, 90, 180, 365]); + expect(defaultExpiry(policy)).toBe(45); + }); + + it('appends "no expiration" only when allowed', () => { + const options = expiryOptions({ ...POLICY, allow_non_expiring: true }); + expect(options[options.length - 1]).toBe(NO_EXPIRY); + expect(expiryOptions(POLICY)).not.toContain(NO_EXPIRY); + }); + + it('falls back to the longest allowed lifetime for an out-of-range default', () => { + const policy = { + ...POLICY, + default_lifetime_days: 400, + max_lifetime_days: 60, + }; + expect(expiryOptions(policy)).toEqual([7, 30, 60]); + expect(defaultExpiry(policy)).toBe(60); + }); + + it('never returns an empty list when the maximum is below every preset', () => { + const policy = { + ...POLICY, + default_lifetime_days: 90, + max_lifetime_days: 3, + }; + expect(expiryOptions(policy)).toEqual([3]); + expect(defaultExpiry(policy)).toBe(3); + }); +}); + +describe('expiryStatus', () => { + const at = (offset: number) => new Date(NOW + offset).toISOString(); + + it('classifies expiry relative to now', () => { + expect(expiryStatus(null, NOW)).toBe('never'); + expect(expiryStatus(at(-1000), NOW)).toBe('expired'); + expect(expiryStatus(at(3 * DAY), NOW)).toBe('expiringSoon'); + expect(expiryStatus(at(7 * DAY), NOW)).toBe('expiringSoon'); + expect(expiryStatus(at(7 * DAY + 60_000), NOW)).toBe('ok'); + }); + + it('does not flag an unparseable date', () => { + expect(expiryStatus('not-a-date', NOW)).toBe('ok'); + }); + + it('countLiveTokens skips expired tokens, like the server cap', () => { + expect( + countLiveTokens( + [ + { expires_at: null }, + { expires_at: at(DAY) }, + { expires_at: at(-DAY) }, + ], + NOW, + ), + ).toBe(2); + }); +}); + +describe('relativeTime', () => { + const ago = (ms: number) => new Date(NOW - ms).toISOString(); + + it('buckets past timestamps', () => { + expect(relativeTime(null, NOW)).toBeNull(); + expect(relativeTime('garbage', NOW)).toBeNull(); + expect(relativeTime(ago(20_000), NOW)).toEqual({ unit: 'now' }); + expect(relativeTime(ago(5 * 60_000), NOW)).toEqual({ + unit: 'minutes', + count: 5, + }); + expect(relativeTime(ago(3 * 3_600_000), NOW)).toEqual({ + unit: 'hours', + count: 3, + }); + expect(relativeTime(ago(2 * DAY), NOW)).toEqual({ unit: 'days', count: 2 }); + expect(relativeTime(ago(45 * DAY), NOW)).toEqual({ unit: 'date' }); + }); + + it('treats clock skew into the future as "now"', () => { + expect(relativeTime(ago(-30_000), NOW)).toEqual({ unit: 'now' }); + }); +}); + +describe('toResourceOptions', () => { + it('accepts only UUID ids', () => { + expect(isUuid(ID_A)).toBe(true); + expect(isUuid('default')).toBe(false); + expect(isUuid(undefined)).toBe(false); + }); + + it('keeps own items with UUID ids, sorted by label', () => { + expect( + toResourceOptions('sources', [ + { id: ID_B, name: 'Zeta' }, + { id: 'default', name: 'Default' }, + { id: ID_A, name: 'Alpha', ownership: 'user' }, + { + id: '33333333-3333-4333-8333-333333333333', + name: 'Team', + ownership: 'team', + }, + ]), + ).toEqual([ + { value: ID_A, label: 'Alpha' }, + { value: ID_B, label: 'Zeta' }, + ]); + }); + + it('skips built-in and team prompts', () => { + expect( + toResourceOptions('prompts', [ + { id: 'default', name: 'default', type: 'public' }, + { id: ID_A, name: 'Mine', type: 'private' }, + { id: ID_B, name: 'Shared', type: 'team' }, + ]), + ).toEqual([{ value: ID_A, label: 'Mine' }]); + }); + + it('reads tools from the {tools: []} envelope and prefers the custom name', () => { + expect( + toResourceOptions('tools', { + success: true, + tools: [ + { + id: ID_A, + name: 'brave', + displayName: 'Brave', + customName: 'My search', + }, + ], + }), + ).toEqual([{ value: ID_A, label: 'My search' }]); + }); + + it('returns nothing for an error body', () => { + expect(toResourceOptions('agents', { success: false })).toEqual([]); + }); +}); diff --git a/frontend/src/settings/accessTokenUtils.ts b/frontend/src/settings/accessTokenUtils.ts new file mode 100644 index 00000000..134d347d --- /dev/null +++ b/frontend/src/settings/accessTokenUtils.ts @@ -0,0 +1,224 @@ +import type { + AccessTokenPolicy, + AccessTokenScope, +} from '../api/services/patService'; + +const DAY_MS = 24 * 60 * 60 * 1000; + +/** Lifetimes (in days) offered in the create form, before policy filtering. */ +export const EXPIRY_PRESETS = [7, 30, 60, 90, 180, 365]; +/** `expires_in_days` value the server reads as "never expires". */ +export const NO_EXPIRY = 0; +/** Tokens expiring within this many days get a warning style. */ +export const EXPIRY_WARNING_DAYS = 7; +/** Families `chat:run` acts on, so they may be restricted alongside it. */ +const CHAT_FAMILIES = ['agents', 'sources']; + +const UUID_REGEX = + /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i; + +/** The server only accepts UUIDs in `resource_filter` (built-ins like `default` are not). */ +export const isUuid = (value: unknown): value is string => + typeof value === 'string' && UUID_REGEX.test(value); + +export const scopeFamily = (scope: string): string => scope.split(':')[0]; +export const scopeAction = (scope: string): string => scope.split(':')[1] ?? ''; + +export interface ScopeGroup { + family: string; + scopes: AccessTokenScope[]; +} + +/** Group the server's scope catalog by family, keeping the server's order. */ +export function groupScopesByFamily(catalog: AccessTokenScope[]): ScopeGroup[] { + const groups: ScopeGroup[] = []; + const byFamily = new Map(); + catalog.forEach((scope) => { + const family = scopeFamily(scope.name); + let group = byFamily.get(family); + if (!group) { + group = { family, scopes: [] }; + byFamily.set(family, group); + groups.push(group); + } + group.scopes.push(scope); + }); + return groups; +} + +/** `x:read` is implied (granted server-side) whenever `x:write` is selected. */ +export function isScopeImplied(scope: string, selected: string[]): boolean { + return ( + scopeAction(scope) === 'read' && + selected.includes(`${scopeFamily(scope)}:write`) + ); +} + +/** + * Scopes to send: the selection minus reads that a selected write already + * implies, in catalog order. The raw selection is kept separately so + * unticking write restores whatever read was before. + */ +export function scopesToSubmit( + selected: string[], + catalog: AccessTokenScope[], +): string[] { + const chosen = new Set(selected); + return catalog + .map((scope) => scope.name) + .filter((name) => chosen.has(name) && !isScopeImplied(name, selected)); +} + +/** Filterable families the selected scopes allow restricting, in policy order. */ +export function eligibleFilterFamilies( + selected: string[], + filterableFamilies: string[], +): string[] { + const families = new Set(selected.map(scopeFamily)); + if (families.has('chat')) CHAT_FAMILIES.forEach((f) => families.add(f)); + return filterableFamilies.filter((family) => families.has(family)); +} + +/** Drop empty and ineligible families; `undefined` when nothing is restricted. */ +export function buildResourceFilter( + selection: Record, + eligibleFamilies: string[], +): Record | undefined { + const out: Record = {}; + eligibleFamilies.forEach((family) => { + const ids = selection[family]; + if (ids && ids.length > 0) out[family] = [...ids]; + }); + return Object.keys(out).length > 0 ? out : undefined; +} + +export interface RestrictionCount { + family: string; + count: number; +} + +/** Per-family counts of a token's `resource_filter`; empty means "all resources". */ +export function restrictionCounts( + resourceFilter: Record | null | undefined, +): RestrictionCount[] { + return Object.entries(resourceFilter ?? {}) + .filter(([, ids]) => Array.isArray(ids) && ids.length > 0) + .map(([family, ids]) => ({ family, count: ids.length })); +} + +type ExpiryPolicy = Pick< + AccessTokenPolicy, + 'default_lifetime_days' | 'max_lifetime_days' | 'allow_non_expiring' +>; + +/** + * Lifetimes (days) to offer: presets within the server maximum, the server + * default when it isn't a preset, and `NO_EXPIRY` last when allowed. + */ +export function expiryOptions(policy: ExpiryPolicy): number[] { + const max = policy.max_lifetime_days; + const days = new Set(EXPIRY_PRESETS.filter((d) => d <= max)); + const fallback = policy.default_lifetime_days; + if (fallback > 0 && fallback <= max) days.add(fallback); + // Never leave the select empty (e.g. max below the smallest preset). + if (days.size === 0 && max > 0) days.add(max); + const options = Array.from(days).sort((a, b) => a - b); + if (policy.allow_non_expiring) options.push(NO_EXPIRY); + return options; +} + +/** The option preselected in the form: the server default, clamped to what's offered. */ +export function defaultExpiry(policy: ExpiryPolicy): number { + const options = expiryOptions(policy); + if (options.includes(policy.default_lifetime_days)) { + return policy.default_lifetime_days; + } + const finite = options.filter((d) => d !== NO_EXPIRY); + return finite.length > 0 ? finite[finite.length - 1] : NO_EXPIRY; +} + +export type ExpiryStatus = 'never' | 'expired' | 'expiringSoon' | 'ok'; + +export function expiryStatus( + expiresAt: string | null | undefined, + now: number = Date.now(), +): ExpiryStatus { + if (!expiresAt) return 'never'; + const at = Date.parse(expiresAt); + if (Number.isNaN(at)) return 'ok'; + if (at <= now) return 'expired'; + if (at - now <= EXPIRY_WARNING_DAYS * DAY_MS) return 'expiringSoon'; + return 'ok'; +} + +/** Tokens that count against `max_per_user`: the server ignores expired ones. */ +export function countLiveTokens( + tokens: { expires_at: string | null }[], + now: number = Date.now(), +): number { + return tokens.filter((tk) => expiryStatus(tk.expires_at, now) !== 'expired') + .length; +} + +export type RelativeTime = + | { unit: 'now' } + | { unit: 'minutes' | 'hours' | 'days'; count: number } + | { unit: 'date' }; + +/** Bucket a past timestamp for display; older than 30 days falls back to a date. */ +export function relativeTime( + value: string | null | undefined, + now: number = Date.now(), +): RelativeTime | null { + if (!value) return null; + const at = Date.parse(value); + if (Number.isNaN(at)) return null; + const minutes = Math.floor(Math.max(0, now - at) / 60_000); + if (minutes < 1) return { unit: 'now' }; + if (minutes < 60) return { unit: 'minutes', count: minutes }; + const hours = Math.floor(minutes / 60); + if (hours < 24) return { unit: 'hours', count: hours }; + const days = Math.floor(hours / 24); + if (days <= 30) return { unit: 'days', count: days }; + return { unit: 'date' }; +} + +export interface ResourceOption { + value: string; + label: string; +} + +/** Families the create form can offer a picker for (workflows have no list endpoint). */ +export const PICKER_FAMILIES = ['agents', 'sources', 'prompts', 'tools']; + +/** + * Turn a list-endpoint payload into picker options: only the caller's own + * items (not team-shared or built-in ones) with ids the server will accept. + */ +export function toResourceOptions( + family: string, + payload: unknown, +): ResourceOption[] { + const envelope = payload as { tools?: unknown } | null; + const rows: unknown[] = Array.isArray(payload) + ? payload + : family === 'tools' && Array.isArray(envelope?.tools) + ? envelope.tools + : []; + return (rows as (Record | null)[]) + .filter( + (row): row is Record => + !!row && + isUuid(row.id) && + row.ownership !== 'team' && + !( + family === 'prompts' && + (row.type === 'public' || row.type === 'team') + ), + ) + .map((row) => ({ + value: row.id as string, + label: String(row.customName || row.displayName || row.name || row.id), + })) + .sort((a, b) => a.label.localeCompare(b.label)); +} diff --git a/frontend/src/settings/index.tsx b/frontend/src/settings/index.tsx index 1087ce30..f7ccf5f0 100644 --- a/frontend/src/settings/index.tsx +++ b/frontend/src/settings/index.tsx @@ -28,6 +28,7 @@ import CustomModels from './CustomModels'; import Sources from './Sources'; import General from './General'; import Logs from './Logs'; +import PersonalAccessTokens from './PersonalAccessTokens'; import Tools from './Tools'; type HiddenGradientType = 'left' | 'right' | undefined; @@ -47,6 +48,8 @@ export default function Settings() { if (path.includes('/settings/tools')) return t('settings.tools.label'); if (path.includes('/settings/custom-models')) return t('settings.customModels.label'); + if (path.includes('/settings/access-tokens')) + return t('settings.accessTokens.label'); return t('settings.general.label'); }; @@ -58,6 +61,7 @@ export default function Settings() { t('settings.logs.label'), t('settings.tools.label'), t('settings.customModels.label'), + t('settings.accessTokens.label'), ]; const [hiddenGradient, setHiddenGradient] = useState('left'); @@ -93,6 +97,8 @@ export default function Settings() { else if (tab === t('settings.tools.label')) navigate('/settings/tools'); else if (tab === t('settings.customModels.label')) navigate('/settings/custom-models'); + else if (tab === t('settings.accessTokens.label')) + navigate('/settings/access-tokens'); }; React.useEffect(() => { @@ -198,6 +204,7 @@ export default function Settings() { element={} /> } /> + } /> } />
diff --git a/frontend/src/settings/types/index.ts b/frontend/src/settings/types/index.ts index 3b1139e2..fe662e7b 100644 --- a/frontend/src/settings/types/index.ts +++ b/frontend/src/settings/types/index.ts @@ -6,21 +6,8 @@ export type ChunkType = { metadata: { [key: string]: string }; }; -export type APIKeyData = { - id: string; - name: string; - key: string; - source: string; - prompt_id: string; - chunks: string; -}; - export type LogEventType = - | 'chat' - | 'schedule' - | 'webhook' - | 'workflow' - | 'system'; + 'chat' | 'schedule' | 'webhook' | 'workflow' | 'system'; export type LogData = { id: string; From 17d66c061d2815bfdb7558ec32a88c63d2a96485 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sat, 19 Sep 2026 23:39:58 +0100 Subject: [PATCH 101/130] docs: personal access tokens guide --- docs/content/Extensions/_meta.js | 4 + .../Extensions/personal-access-tokens.mdx | 181 ++++++++++++++++++ 2 files changed, 185 insertions(+) create mode 100644 docs/content/Extensions/personal-access-tokens.mdx diff --git a/docs/content/Extensions/_meta.js b/docs/content/Extensions/_meta.js index c265957e..af42438c 100644 --- a/docs/content/Extensions/_meta.js +++ b/docs/content/Extensions/_meta.js @@ -3,6 +3,10 @@ export default { "title": "🔑 Getting API key", "href": "/Extensions/api-key-guide" }, + "personal-access-tokens": { + "title": "🎟️ Personal Access Tokens", + "href": "/Extensions/personal-access-tokens" + }, "chat-widget": { "title": "💬️ Chat Widget", "href": "/Extensions/chat-widget" diff --git a/docs/content/Extensions/personal-access-tokens.mdx b/docs/content/Extensions/personal-access-tokens.mdx new file mode 100644 index 00000000..24ab642d --- /dev/null +++ b/docs/content/Extensions/personal-access-tokens.mdx @@ -0,0 +1,181 @@ +--- +title: Personal Access Tokens +description: Scoped, revocable API tokens for managing agents, sources and other DocsGPT resources from the CLI, scripts and CI/CD pipelines. +--- + +# Personal Access Tokens + +A personal access token (PAT) lets a script, the [DocsGPT CLI](https://github.com/arc53/DocsGPT-cli) or a CI/CD pipeline act on your account without a browser session. Unlike an [agent API key](/Extensions/api-key-guide), which can only talk to one agent, a PAT manages resources: it can create and update agents, upload sources, edit prompts and tools, and run agents for benchmarking. + +Every token is limited in three ways: + +- **Scopes** decide which parts of the API the token may call. +- **Resource restrictions** (optional) narrow a token to specific agents, sources, prompts, tools or workflows. +- **Expiry** ends the token's life automatically. + +## Creating a token + +1. Open **Settings → Access Tokens** in the DocsGPT web app. +2. Choose **Create token**, give it a name, and select the scopes it needs. +3. Optionally restrict it to specific resources and pick an expiry. +4. Copy the token. It starts with `dgpt_pat_` and is shown **once**. DocsGPT stores only a hash of it, so a lost token cannot be recovered. Revoke it and create a new one. + +Tokens can only be created and revoked from a signed-in session. A token cannot create, list or revoke tokens, so a leaked token cannot mint a replacement for itself. + +## Using a token + +Send the token as a bearer credential: + +```bash +export DOCSGPT_URL=https://docsgpt.example.com +export DOCSGPT_TOKEN=dgpt_pat_... + +curl -H "Authorization: Bearer $DOCSGPT_TOKEN" "$DOCSGPT_URL/api/user/me" +``` + +`GET /api/user/me` works with any valid token and reports what the token may do, which makes it a convenient first step in a pipeline: + +```json +{ + "success": true, + "user_id": "alice@example.com", + "roles": ["user"], + "auth_method": "pat", + "token": { + "id": "0b6c...", + "name": "ci-deploy", + "scopes": ["agents:read", "agents:write"], + "resource_filter": {} + } +} +``` + +### Applying agent definitions + +Agents can be exported to YAML and applied back, which makes them reviewable and deployable like any other configuration. With the CLI: + +```bash +docsgpt-cli agents export -o support-bot.agent.yaml +docsgpt-cli agents apply -f support-bot.agent.yaml --dry-run +docsgpt-cli agents apply -f support-bot.agent.yaml +``` + +Or with the API directly (`agents:write`): + +```bash +curl -X POST "$DOCSGPT_URL/api/import_agent/plan" \ + -H "Authorization: Bearer $DOCSGPT_TOKEN" \ + -H "Content-Type: application/json" \ + -d "$(jq -Rs '{yaml: .}' support-bot.agent.yaml)" +``` + +`/api/import_agent/plan` is a dry run that reports whether the agent would be created or updated and how each referenced source, tool and prompt resolves. `/api/import_agent` applies it. An agent is matched by `metadata.id`, then `metadata.slug`; when nothing matches, a new draft agent is created. + +### GitHub Actions example + +```yaml +jobs: + deploy-agents: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - name: Apply agent definitions + env: + DOCSGPT_URL: ${{ vars.DOCSGPT_URL }} + DOCSGPT_TOKEN: ${{ secrets.DOCSGPT_TOKEN }} + run: docsgpt-cli agents apply -f agents/ +``` + +## Scopes + +A `write` scope includes the matching `read` scope. + +| Scope | Allows | +| --- | --- | +| `agents:read` | View agents, folders, guardrail events and export agent definitions | +| `agents:write` | Create, update, delete, share and import (apply) agents and folders | +| `agents:keys` | Regenerate agent API keys and read incoming webhook URLs | +| `sources:read` | View sources, their files, chunks and ingestion task status | +| `sources:write` | Upload, ingest, sync, edit and delete sources and chunks | +| `prompts:read` / `prompts:write` | View / create, update and delete prompts | +| `tools:read` / `tools:write` | View / create, update and delete tools and MCP servers | +| `models:read` / `models:write` | View models / manage custom models | +| `workflows:read` / `workflows:write` | View / create, update and delete workflows | +| `schedules:read` / `schedules:write` | View / create, update, run and delete agent schedules | +| `conversations:read` / `conversations:write` | View / rename, delete and rate conversations | +| `analytics:read` | View usage analytics and logs | +| `teams:read` | View teams, members and resource shares | +| `chat:run` | Ask agents and search sources (`/api/answer`, `/stream`, `/api/search`); used for benchmarking | + +`agents:keys` is separate from `agents:write` on purpose: a deployment token that updates agents does not need to be able to read or rotate the secrets other systems use to call them. + +Some parts of the API are never available to a token, whatever its scopes: token management, the admin API, team management, sign-in flows, device pairing, and the interactive OAuth handshakes used by connectors and MCP servers. A token also never carries the `admin` role, even when its owner is an admin. + +Authorization is deny by default. An endpoint that is not explicitly mapped to a scope cannot be called with a token, and answers `403` with `"error": "not_available_to_tokens"`. A mapped endpoint called without the scope answers `403` with `"error": "insufficient_scope"` and names the `required_scope`. + +## Resource restrictions + +A token can be narrowed to specific resources in any of these families: `agents`, `sources`, `prompts`, `tools`, `workflows`. A family that is not listed stays unrestricted within the token's scopes. + +```json +{ + "name": "support-bot-deploy", + "scopes": ["agents:write", "chat:run"], + "resource_filter": { "agents": ["3f0e8f0c-5a53-4f0e-9a39-0e5f4f8d2c11"] }, + "expires_in_days": 30 +} +``` + +For a restricted family the token: + +- can read, update and delete only the listed resources, and listings show only those; +- **cannot create** new resources of that family, since a new resource would be outside the list; +- cannot attach a resource outside the list to something else, for example set an agent's source to a source the token may not use; +- is refused (`403`, `"error": "resource_not_allowed"`) wherever DocsGPT cannot prove the request stays inside the list. Schedules, conversations and analytics are closed to agent-restricted tokens for that reason, and `/api/sources/paginated` is closed to source-restricted tokens (use `/api/sources`). + +Restrictions and chat (`chat:run`): + +- A token restricted to specific **agents** must pass `agent_id` in the request body, and it must be one of the listed agents. An agent `api_key` in the body is refused. +- A token restricted to specific **sources** (but not agents) may chat against those sources with `active_docs`. It cannot run agents, because an agent brings its own sources. Restrict the token to agents instead to allow that. + +Restrictions and `agents apply`: a token restricted to specific agents can apply a definition only when it updates one of those agents. A token restricted on sources, prompts, tools or workflows cannot import agents at all, because an import resolves those references by name and may create them. + +## Expiry and revocation + +- A token created without an explicit lifetime expires after `PAT_DEFAULT_LIFETIME_DAYS` (90 by default). Users can choose any lifetime up to `PAT_MAX_LIFETIME_DAYS` (365 by default). +- Non-expiring tokens are available only when the operator sets `PAT_ALLOW_NON_EXPIRING=true`. +- Revoking a token in **Settings → Access Tokens** takes effect on the next request. +- Admins can list a user's tokens with `GET /api/admin/users//tokens` and revoke any token with `DELETE /api/admin/tokens/`. The admin **revoke sessions** action also revokes all of that user's tokens. +- Tokens of a deactivated user (through the admin API or SCIM) stop working immediately and work again if the user is reactivated. +- Token creation and revocation are recorded in the authentication audit log (`pat_created`, `pat_revoked`), visible to admins. + +Each user may hold up to `PAT_MAX_PER_USER` live tokens (25 by default). The token list shows when and from which IP address each token was last used. + +## Operator settings + +| Setting | Default | Purpose | +| --- | --- | --- | +| `PAT_ENABLED` | `true` | Allow users to create and use personal access tokens | +| `PAT_DEFAULT_LIFETIME_DAYS` | `90` | Lifetime of a token created without an explicit expiry | +| `PAT_MAX_LIFETIME_DAYS` | `365` | Longest lifetime a user may request | +| `PAT_ALLOW_NON_EXPIRING` | `false` | Let users create tokens that never expire | +| `PAT_MAX_PER_USER` | `25` | Maximum number of live tokens per user | + +Personal access tokens need a stable user identity, so they are available with `AUTH_TYPE=oidc` and with authentication disabled (single-user self-hosting). They are not available with `simple_jwt` or `session_jwt`. See the [Settings Reference](/Deploying/Settings-Reference) for details. + +## Management API + +These endpoints need a signed-in session and cannot be called with a token. + +| Endpoint | Purpose | +| --- | --- | +| `GET /api/user/tokens` | List your tokens, the scope catalog and the server's token policy | +| `POST /api/user/tokens` | Create a token. Body: `name`, `scopes`, optional `resource_filter`, optional `expires_in_days` (`0` = never, when allowed). The response carries the plaintext `token` once | +| `DELETE /api/user/tokens/` | Revoke a token | + +## Good practice + +- Give each pipeline its own token with the narrowest scopes that work, and name it after where it is used. +- Store tokens in your CI system's secret store. Never commit them. The `dgpt_pat_` prefix lets secret scanners recognise them. +- Prefer short lifetimes for tokens used by automation you can easily re-provision. +- Revoke a token as soon as it is no longer needed or may have been exposed. From 934ae1afb31df5a44a43ecd4f5df81ab6d467260 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sat, 19 Sep 2026 23:54:35 +0100 Subject: [PATCH 102/130] fix(pat): close message replay to restricted tokens, filter builtin tools, reject non-object bodies Message tail and the ASGI reconnect stream cannot tie a message to an allowlist, so any token with a resource filter is refused there. The tools listing now applies the allowlist to default and builtin rows as well. Token creation answers 400 instead of 500 for a JSON body that is not an object. --- docs/content/Extensions/personal-access-tokens.mdx | 2 +- docsgpt/api/asgi_auth.py | 14 +++++++++++++- docsgpt/api/pat/routes.py | 6 +++++- docsgpt/api/pat/rules.py | 5 ++++- docsgpt/api/user/tools/routes.py | 6 +++--- tests/api/test_asgi_auth.py | 8 ++++++++ tests/api/test_pat_routes.py | 5 +++++ tests/api/test_pat_rules.py | 6 ++++++ 8 files changed, 45 insertions(+), 7 deletions(-) diff --git a/docs/content/Extensions/personal-access-tokens.mdx b/docs/content/Extensions/personal-access-tokens.mdx index 24ab642d..c533b528 100644 --- a/docs/content/Extensions/personal-access-tokens.mdx +++ b/docs/content/Extensions/personal-access-tokens.mdx @@ -131,7 +131,7 @@ For a restricted family the token: - can read, update and delete only the listed resources, and listings show only those; - **cannot create** new resources of that family, since a new resource would be outside the list; - cannot attach a resource outside the list to something else, for example set an agent's source to a source the token may not use; -- is refused (`403`, `"error": "resource_not_allowed"`) wherever DocsGPT cannot prove the request stays inside the list. Schedules, conversations and analytics are closed to agent-restricted tokens for that reason, and `/api/sources/paginated` is closed to source-restricted tokens (use `/api/sources`). +- is refused (`403`, `"error": "resource_not_allowed"`) wherever DocsGPT cannot prove the request stays inside the list. Schedules, conversations and analytics are closed to agent-restricted tokens for that reason, the message replay endpoints (`/api/messages//tail` and `/api/messages//events`) are closed to every restricted token, and `/api/sources/paginated` is closed to source-restricted tokens (use `/api/sources`). Restrictions and chat (`chat:run`): diff --git a/docsgpt/api/asgi_auth.py b/docsgpt/api/asgi_auth.py index aa2ac5bd..a65e83ab 100644 --- a/docsgpt/api/asgi_auth.py +++ b/docsgpt/api/asgi_auth.py @@ -31,7 +31,8 @@ async def authenticate( request: The incoming Starlette request. pat_scope: Scope a personal access token needs for this route. Left unset, the route rejects PATs outright (deny by default, matching - the Flask rule table in ``docsgpt/api/pat/rules.py``). + the Flask rule table in ``docsgpt/api/pat/rules.py``). A token with + a resource filter is always rejected. Returns: tuple: ``(claims, None)`` for an authenticated caller, ``(None, None)`` @@ -52,6 +53,17 @@ async def authenticate( {"success": False, "message": "Token lacks the required scope", "error": "insufficient_scope"}, status_code=403, ) + if decoded.get("resource_filter"): + # These routes sit outside the Flask rule table and cannot tie what + # they serve to an allowlist, so a restricted token is kept out. + return None, JSONResponse( + { + "success": False, + "message": "This endpoint is not available to a resource-restricted token", + "error": "resource_not_allowed", + }, + status_code=403, + ) return decoded, None # The denylist is a sync Redis read; keep it off the event loop. if settings.AUTH_TYPE == "oidc" and await anyio.to_thread.run_sync(oidc_session_denied, decoded): diff --git a/docsgpt/api/pat/routes.py b/docsgpt/api/pat/routes.py index 84360a1b..ad22cf6d 100644 --- a/docsgpt/api/pat/routes.py +++ b/docsgpt/api/pat/routes.py @@ -113,7 +113,11 @@ class PersonalAccessTokens(Resource): if not auth_type_supports_pats(): return _error("Personal access tokens are not available on this server", 403) - body = request.get_json(silent=True) or {} + body = request.get_json(silent=True) + if body is None: + body = {} + if not isinstance(body, dict): + return _error("Request body must be a JSON object", 400) name = body.get("name") if not isinstance(name, str) or not name.strip(): return _error("name is required", 400) diff --git a/docsgpt/api/pat/rules.py b/docsgpt/api/pat/rules.py index f340d698..6794390b 100644 --- a/docsgpt/api/pat/rules.py +++ b/docsgpt/api/pat/rules.py @@ -53,6 +53,8 @@ def _rule(scope: Optional[str] = None, *ids: Locator, any_of: tuple[str, ...] = return Rule(scopes=scopes, family=family, ids=tuple(ids), **kwargs) +_ALL_FAMILIES = ("agents", "sources", "prompts", "tools", "workflows") + # Ids of other families that agent create/update accept in their JSON-or-form body. _AGENT_BODY_REFS = ( ("sources", (BODY, "source")), @@ -216,8 +218,9 @@ RULES: dict[tuple[str, str], Rule] = { ("/api/get_conversations", "GET"): _rule("conversations:read", blocked_by=("agents",)), ("/api/search_conversations", "GET"): _rule("conversations:read", blocked_by=("agents",)), ("/api/get_single_conversation", "GET"): _rule("conversations:read", blocked_by=("agents",)), + # A message cannot be tied to an allowlist from here, so any restricted token is kept out. ("/api/messages//tail", "GET"): _rule( - any_of=("conversations:read", "chat:run"), family=None + any_of=("conversations:read", "chat:run"), family=None, blocked_by=_ALL_FAMILIES ), ("/api/delete_conversation", "POST"): _rule("conversations:write", blocked_by=("agents",)), ("/api/delete_all_conversations", "GET"): _rule("conversations:write", blocked_by=("agents",)), diff --git a/docsgpt/api/user/tools/routes.py b/docsgpt/api/user/tools/routes.py index 8b70abac..ddc56916 100644 --- a/docsgpt/api/user/tools/routes.py +++ b/docsgpt/api/user/tools/routes.py @@ -266,9 +266,6 @@ class GetTools(Resource): shaped = _shape_tool(row, ownership="team", force_strip_secret=True) shaped["team_access"] = team_shared.get(str(row["id"])) user_tools.append(shaped) - # A resource-restricted token sees only its allowed tools; the - # default and builtin rows appended below belong to no user. - user_tools = filter_listing(request, "tools", user_tools) # ``scheduler`` is dual-registered (default chat tool + agent- # selectable builtin) and resolves to the same synthetic uuid5 id. @@ -298,6 +295,9 @@ class GetTools(Resource): builtin_copy.get("name") in WORKFLOW_ONLY_BUILTINS ) user_tools.append(builtin_copy) + # A resource-restricted token sees only its allowed tools. Default + # and builtin rows have ids too, so they follow the same allowlist. + user_tools = filter_listing(request, "tools", user_tools) except Exception as err: current_app.logger.error(f"Error getting user tools: {err}", exc_info=True) return make_response(jsonify({"success": False}), 400) diff --git a/tests/api/test_asgi_auth.py b/tests/api/test_asgi_auth.py index 341ded0d..d40fae0e 100644 --- a/tests/api/test_asgi_auth.py +++ b/tests/api/test_asgi_auth.py @@ -156,3 +156,11 @@ class TestPersonalAccessTokens: decoded, error = await asgi_auth.authenticate(_request(), pat_scope="chat:run") denied.assert_not_called() assert error is None and decoded is not None + + async def test_restricted_token_is_refused_even_with_the_scope(self): + claims = dict(_PAT_CLAIMS, resource_filter={"agents": ["a1"]}) + with patch.object(asgi_auth, "handle_auth", return_value=claims): + decoded, error = await asgi_auth.authenticate(_request(), pat_scope="chat:run") + assert decoded is None + assert error.status_code == 403 + assert json.loads(error.body)["error"] == "resource_not_allowed" diff --git a/tests/api/test_pat_routes.py b/tests/api/test_pat_routes.py index cd157284..be07b7d5 100644 --- a/tests/api/test_pat_routes.py +++ b/tests/api/test_pat_routes.py @@ -118,6 +118,11 @@ class TestCreate: def test_validation(self, client, db, body): assert _create(client, **body).status_code == 400 + @pytest.mark.parametrize("body", [[1], "text", 5]) + def test_non_object_body_is_a_client_error(self, client, db, body): + with _session(): + assert client.post("/api/user/tokens", json=body).status_code == 400 + def test_duplicate_name_conflicts(self, client, db): assert _create(client).status_code == 201 assert _create(client).status_code == 409 diff --git a/tests/api/test_pat_rules.py b/tests/api/test_pat_rules.py index 3a3cd75a..c6c9f0b8 100644 --- a/tests/api/test_pat_rules.py +++ b/tests/api/test_pat_rules.py @@ -230,6 +230,12 @@ class TestResourceRestrictions: assert _denied(_call(client, "GET", f"/api/agents/{AGENT_A}/schedules", claims)) is None assert _denied(_call(client, "GET", f"/api/agents/{AGENT_B}/schedules", claims)) == "resource_not_allowed" + @pytest.mark.parametrize("family", ["agents", "sources", "prompts", "tools", "workflows"]) + def test_message_tail_is_closed_to_any_restricted_token(self, client, family): + claims = _claims(["chat:run"], {family: [AGENT_A]}) + assert _denied(_call(client, "GET", "/api/messages/m1/tail", claims)) == "resource_not_allowed" + assert _denied(_call(client, "GET", "/api/messages/m1/tail", _claims(["chat:run"]))) is None + def test_sql_paged_listing_is_closed_to_restricted_tokens(self, client): claims = self._restricted(["sources:read"], sources=[SOURCE_A]) assert _denied(_call(client, "GET", "/api/sources/paginated", claims)) == "resource_not_allowed" From 8c5a5190d9da32ce5d59686fd93fe4273a928455 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:03 +0100 Subject: [PATCH 103/130] fix(retriever): carry a graph source's own options to the retriever The three per-source graph options are read from the retrieval config the Dispatcher hands over, and it only hands one over for a source it considers overridden -- which it decided from chunks, score_threshold, rephrase_query and prescreen alone. A source that changed only its graph options was not "overridden", so nothing was carried and every graph source ran the defaults: the UI toggles did nothing at all. They count as an override now, for graphrag sources only. They mean nothing to any other retriever, and an override also hands the source its own chunk budget, which a classic source must not pick up from a graph setting. --- docsgpt/retriever/dispatcher.py | 11 +++++++++++ tests/test_dispatcher.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/docsgpt/retriever/dispatcher.py b/docsgpt/retriever/dispatcher.py index a14fc02f..31cc75e0 100644 --- a/docsgpt/retriever/dispatcher.py +++ b/docsgpt/retriever/dispatcher.py @@ -185,12 +185,23 @@ class Dispatcher(BaseRetriever): score_threshold / rephrase_query) plus an opted-in prescreen config; a source left at defaults takes the global path so all-classic retrieval stays byte-identical with zero extra LLM calls. + + A graph source's ``graph`` options count too: they are read from the + per-source config this records, so a source that changes only those + would otherwise run the defaults and the options would do nothing. + They mean nothing to any other retriever, so they only count for + ``graphrag`` -- an override hands the source its own chunk budget as + well, which a classic source must not pick up from a graph setting. """ return ( retrieval.chunks != _DEFAULT_RETRIEVAL.chunks or retrieval.score_threshold != _DEFAULT_RETRIEVAL.score_threshold or retrieval.rephrase_query != _DEFAULT_RETRIEVAL.rephrase_query or retrieval.prescreen is not None + or ( + (retrieval.retriever or "").lower() == "graphrag" + and retrieval.graph != _DEFAULT_RETRIEVAL.graph + ) ) @staticmethod diff --git a/tests/test_dispatcher.py b/tests/test_dispatcher.py index 18816aa1..050664c1 100644 --- a/tests/test_dispatcher.py +++ b/tests/test_dispatcher.py @@ -74,6 +74,35 @@ class TestDispatcherGrouping: assert "b" not in retrievals + def test_graph_options_count_as_an_override(self, _patch_llm_creator): + """A graph source that changes only its graph options still needs its + config carried over: those options live on the per-source retrieval the + Dispatcher hands the retriever, so without this the UI toggles are + no-ops and every source runs the defaults.""" + sources = [ + { + "id": "a", + "retrieval": RetrievalConfig( + retriever="graphrag", graph={"seed_strategy": "relationships"} + ), + }, + {"id": "b", "retrieval": RetrievalConfig(retriever="graphrag")}, + ] + d = Dispatcher(source={"question": "q", "active_docs": ["a", "b"]}, sources=sources) + retrievals = d._groups[0]["retrievals"] + assert "a" in retrievals + assert retrievals["a"].graph.seed_strategy == "relationships" + # A source on the defaults still takes the shared path. + assert "b" not in retrievals + + def test_graph_options_on_a_classic_source_are_not_an_override(self, _patch_llm_creator): + # They only mean anything to the graph retriever; treating them as an + # override would hand a classic source its own chunk budget. + sources = [{"id": "a", "retrieval": RetrievalConfig(graph={"blend_vector": False})}] + d = Dispatcher(source={"question": "q", "active_docs": ["a"]}, sources=sources) + assert d._groups[0]["retrievals"] == {} + + @pytest.mark.unit class TestDispatcherSharedBudget: def test_single_group_full_budget(self, _patch_llm_creator): From bc0ef9f3b01f34aaac5188fe2561b96ce88f2aa0 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:03 +0100 Subject: [PATCH 104/130] fix(retriever): search a graph source classically when its graph answers nothing Only a raise routed a source to the ClassicRAG fallback, and every graph read logs its own failure and returns empty. So a query that broke, a half-built graph and a walk that genuinely found nothing were indistinguishable, and each made the source contribute nothing at all to the answer -- no fallback, no vector blend, which is skipped by the same early return. A graph source that produces no documents now joins the classic batch, exactly as a source with no graph already does. The passage stage also rescanned every node's chunk list once per candidate. It inverts the links once instead: 7.3 ms to 0.16 ms on a 400-node subgraph at the candidate cap, with identical output. --- docsgpt/retriever/graph_rag.py | 56 +++++++++++++++++------ tests/graphrag/test_retriever_passages.py | 19 ++++++++ tests/retriever/test_graph_rag.py | 17 +++++-- 3 files changed, 76 insertions(+), 16 deletions(-) diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index 818f4f07..8d191f7b 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -56,6 +56,19 @@ def _idf(doc_freq: Any) -> float: return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0) +def _nodes_by_chunk(chunk_links: Dict[str, List[str]]) -> Dict[str, List[str]]: + """Invert ``node -> chunk ids`` into ``chunk id -> node ids``. + + Node order within a chunk follows the node order of ``chunk_links``, so the + passage edges are added in the same order as before. + """ + inverted: Dict[str, List[str]] = {} + for node, chunks in chunk_links.items(): + for chunk_id in chunks or (): + inverted.setdefault(chunk_id, []).append(node) + return inverted + + def _damping(passage_nodes: bool) -> float: """PageRank damping for a ranking mode: the value that mode was measured at.""" return DAMPING_WITH_PASSAGES if passage_nodes else DAMPING_ENTITIES_ONLY @@ -320,9 +333,12 @@ class GraphRAGRetriever(BaseRetriever): personalization = dict(seeds) passage_of: Dict[str, str] = {} + # Inverted once: scanning every node's chunk list per candidate is + # quadratic, and at the candidate cap it cost more than the walk it + # feeds (76 ms against 0.5 ms measured). + nodes_by_chunk = _nodes_by_chunk(chunk_links) for chunk_id in candidate_ids: - linked = [n for n, chunks in chunk_links.items() if chunk_id in chunks] - linked = [n for n in linked if n in graph] + linked = [n for n in nodes_by_chunk.get(chunk_id, ()) if n in graph] if not linked: continue passage_node = f"chunk::{chunk_id}" @@ -602,8 +618,9 @@ class GraphRAGRetriever(BaseRetriever): Graph sources keep their own slot in source order; every graphless source collapses into a single ClassicRAG run that occupies the slot of - the first graphless source. Sources whose PPR retrieval raises are - collected and retried as one more classic batch, appended at the end. + the first graphless source. Sources whose PPR retrieval raises, or + answers nothing, are collected and retried as one more classic batch, + appended at the end. """ try: counts = store.count_nodes_many(sources) @@ -629,7 +646,7 @@ class GraphRAGRetriever(BaseRetriever): segments.append([]) graphless.append(source_id) - failed: List[str] = [] + fallback: List[str] = [] query_embedding = None if graphed: # Embedded once for the whole retrieval, not once per graph source. @@ -642,26 +659,39 @@ class GraphRAGRetriever(BaseRetriever): f"GraphRAG query embedding failed, falling back: {e}", exc_info=True, ) - failed, graphed = list(graphed), [] + fallback, graphed = list(graphed), [] for source_id in graphed: try: - segments[graph_slots[source_id]] = self._graph_docs_for_source( - store, source_id, query_embedding - ) + docs = self._graph_docs_for_source(store, source_id, query_embedding) except Exception as e: logging.error( f"GraphRAG retrieval failed for {source_id}, falling back: {e}", exc_info=True, ) - failed.append(source_id) + fallback.append(source_id) + continue + if not docs: + # Empty is not an answer. Every graph read reports its own + # failure and returns nothing, so "no rows" covers a query that + # broke or a half-built graph as much as a walk that found + # nothing — and only a raise reaches the fallback, so the + # source would otherwise contribute nothing at all. Searching + # it classically is what a source with no graph already gets. + logging.info( + "GraphRAG retrieval returned nothing for %s, falling back", + source_id, + ) + fallback.append(source_id) + continue + segments[graph_slots[source_id]] = docs # Every remaining segment is a ClassicRAG fan-out, and each of its legs # checks out of the *same* per-DSN pool this store is holding. Hand the # graph connection back first, or concurrent GraphRAG retrievals occupy # every slot and then block on their own fallbacks until PoolTimeout. # ``close()`` nulls the connection, so ``_get_data``'s finally stays correct. - if graphless or failed: + if graphless or fallback: try: store.close() except Exception as e: @@ -669,8 +699,8 @@ class GraphRAGRetriever(BaseRetriever): if graphless: segments[classic_slot] = self._classic_for_sources(graphless) - if failed: - segments.append(self._classic_for_sources(failed)) + if fallback: + segments.append(self._classic_for_sources(fallback)) return [doc for segment in segments for doc in segment] diff --git a/tests/graphrag/test_retriever_passages.py b/tests/graphrag/test_retriever_passages.py index fcc561ff..d2d9c29c 100644 --- a/tests/graphrag/test_retriever_passages.py +++ b/tests/graphrag/test_retriever_passages.py @@ -130,3 +130,22 @@ class TestChunkSimilaritiesGuard: store = object.__new__(GraphStore) assert store.chunk_similarities("src", chunk_ids, embedding) == {} + + +@pytest.mark.unit +class TestNodesByChunk: + """The passage stage inverts node->chunks once instead of rescanning.""" + + def test_inverts_and_keeps_node_order(self): + from docsgpt.retriever.graph_rag import _nodes_by_chunk + + assert _nodes_by_chunk({"n1": ["c1", "c2"], "n2": ["c2"], "n3": []}) == { + "c1": ["n1"], + "c2": ["n1", "n2"], + } + + def test_no_links_invert_to_nothing(self): + from docsgpt.retriever.graph_rag import _nodes_by_chunk + + assert _nodes_by_chunk({}) == {} + assert _nodes_by_chunk({"n1": None}) == {} diff --git a/tests/retriever/test_graph_rag.py b/tests/retriever/test_graph_rag.py index 71c996d5..5203eee1 100644 --- a/tests/retriever/test_graph_rag.py +++ b/tests/retriever/test_graph_rag.py @@ -180,7 +180,9 @@ class TestGraphRAGPoolDiscipline: mock_store_cls.return_value = store rag = _make_retriever() - with patch.object(rag, "_graph_docs_for_source", return_value=[]): + # A real result: an empty one now falls back like a failure does. + graph_docs = [{"title": "g", "text": "graph text", "source": "src1", "filename": "g"}] + with patch.object(rag, "_graph_docs_for_source", return_value=graph_docs): with patch.object(rag, "_classic_for_sources") as classic: rag._get_data() @@ -355,15 +357,24 @@ class TestGraphRAGHappyPath: @patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10) @patch("docsgpt.retriever.graph_rag.GraphStore") @patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True) - def test_no_seeds_returns_empty( + def test_a_graph_that_answers_nothing_falls_back_to_classic( self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed ): + """Empty is not an answer. Every graph read swallows its own errors and + returns nothing, so "no rows" covers a broken query as much as a walk + that found nothing — and the source would contribute nothing at all, + with no fallback, because only a raise routes one to ClassicRAG.""" store = _store_with_graph([], [], {}, {}, []) store.count_nodes_many.side_effect = lambda ids: {s: 5 for s in ids} mock_store_cls.return_value = store rag = _make_retriever() - assert rag._get_data() == [] + seen = _recording_classic(rag, [_CLASSIC_DOC]) + + docs = rag._get_data() + + assert seen == [["src1"]] + assert [doc["text"] for doc in docs] == ["classic"] # ── IDF down-weighting ──────────────────────────────────────────────────────── From ccd8eb612ffb2476fe59cd435266b20b885eb9a4 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:04 +0100 Subject: [PATCH 105/130] fix(agents): hand the graph tool's connection back, and gate it like the rest Three faults in the graph tool, all on the agent's path: The store was cached on the tool, and the executor caches the tool for the whole agent run -- so one pgvector pooled connection stayed checked out across every LLM round trip of that run, minutes at a time, and enough concurrent runs exhaust the pool. GraphRAGRetriever releases its store before falling back for this reason. The tool now releases it at the end of each action. It gated on GRAPHRAG_ENABLED where everything else asks graphrag_available(), which also requires the pgvector store. Under any other vector store the graph tables are not the ones the sources were ingested into, but the tool was still offered and still queried Postgres. Pages were labelled by hand rather than through labels_from_metadata, which exists so citation labels match across retrievers. A page read by the tool and the same chunk retrieved by internal_search are one document, and citations key on (source, title) -- so the research agent gave that document two citation numbers. The recorded doc also keeps the full chunk text now, so it dedupes against the retriever's copy; only what the model reads is truncated. --- docsgpt/agents/tools/graph_search.py | 48 ++++++++--- tests/graphrag/test_graph_search_tool.py | 104 ++++++++++++++++++++++- 2 files changed, 138 insertions(+), 14 deletions(-) diff --git a/docsgpt/agents/tools/graph_search.py b/docsgpt/agents/tools/graph_search.py index c2c15a7b..b370949e 100644 --- a/docsgpt/agents/tools/graph_search.py +++ b/docsgpt/agents/tools/graph_search.py @@ -28,7 +28,8 @@ import logging from typing import Any, Dict, List, Optional from docsgpt.agents.tools.base import Tool -from docsgpt.core.settings import settings +from docsgpt.graphrag import graphrag_available +from docsgpt.retriever.labels import labels_from_metadata logger = logging.getLogger(__name__) @@ -61,6 +62,24 @@ class GraphSearchTool(Tool): self._store = GraphStore() return self._store + def _release_store(self) -> None: + """Hand the pooled connection back at the end of an action. + + The executor caches this tool for the whole agent run, so a store kept + between actions pins one connection of the shared pgvector pool across + every LLM round trip of that run -- minutes at a time, and enough + concurrent runs exhaust the pool. ``GraphRAGRetriever`` releases its + store before falling back for the same reason. Checking one back out + costs a pool acquire. + """ + store, self._store = self._store, None + if store is None: + return + try: + store.close() + except Exception as exc: # noqa: BLE001 -- releasing must not fail an action + logger.debug(f"Graph tool could not release its store: {exc}") + def _embed(self, text: str) -> Optional[List[float]]: try: from docsgpt.vectorstore.base import get_embeddings @@ -72,7 +91,9 @@ class GraphSearchTool(Tool): # -- actions ------------------------------------------------------------- def execute_action(self, action_name: str, **kwargs): - if not settings.GRAPHRAG_ENABLED: + # The graph lives in the pgvector store, so the flag alone is not + # enough: under another vector store there is no graph to read. + if not graphrag_available(): return "The knowledge graph is not enabled for this deployment." if not self._sources(): return "No graph-backed sources are configured." @@ -86,6 +107,8 @@ class GraphSearchTool(Tool): except Exception as e: # noqa: BLE001 logger.error(f"Graph tool action {action_name} failed: {e}", exc_info=True) return "The graph lookup failed." + finally: + self._release_store() return f"Unknown action: {action_name}" def _search_entities(self, **kwargs) -> str: @@ -137,18 +160,17 @@ class GraphSearchTool(Tool): parts: List[str] = [] for source_id in self._sources(): for page in store.entity_pages(source_id, entity): - metadata = page.get("metadata") or {} - title = ( - metadata.get("file_path") - or metadata.get("title") - or metadata.get("source") - or "document" - ) - text = (page.get("text") or "")[:MAX_PAGE_CHARS] - doc = {"title": title, "text": text, "source": metadata.get("source", "")} + text = page.get("text") or "" + # The retrievers' own labelling: a page read here and the same + # chunk retrieved by internal_search are one document, and + # citations key on (source, title). Labelling it differently + # gives that document two citation numbers. + labels = labels_from_metadata(page.get("metadata"), text, source_id) + doc = {**labels, "text": text} if doc not in self.retrieved_docs: self.retrieved_docs.append(doc) - parts.append(f"--- {title} ---\n{text}") + header = labels["filename"] or labels["title"] + parts.append(f"--- {header} ---\n{text[:MAX_PAGE_CHARS]}") if not parts: return f"No documents mention {entity!r}." return "\n\n".join(parts) @@ -259,7 +281,7 @@ def add_graph_search_tool(tools_dict: Dict, retriever_config: Dict) -> None: tool follows that same per-source exposure choice. A graph source left at ``prefetch`` in a classic agent is used for ranking only. """ - if not settings.GRAPHRAG_ENABLED: + if not graphrag_available(): return source = retriever_config.get("source") or {} if not source.get("active_docs") or not sources_have_graph(source): diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py index 559c911f..465a1c25 100644 --- a/tests/graphrag/test_graph_search_tool.py +++ b/tests/graphrag/test_graph_search_tool.py @@ -9,6 +9,8 @@ must refuse clearly rather than silently when it has nothing to offer. from __future__ import annotations +import pytest + from docsgpt.agents.tools.graph_search import ( GRAPH_TOOL_ID, GraphSearchTool, @@ -38,6 +40,7 @@ class _StubStore: def _tool(monkeypatch, store, enabled=True): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", enabled) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") tool = GraphSearchTool({"source": SOURCE}) tool._store = store monkeypatch.setattr(tool, "_embed", lambda text: [0.0, 0.1]) @@ -52,6 +55,7 @@ class TestGating: def test_reports_when_no_sources_are_configured(self, monkeypatch): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") tool = GraphSearchTool({"source": {"active_docs": []}}) assert "No graph-backed sources" in tool.execute_action( @@ -109,16 +113,49 @@ class TestActions: tool = _tool( monkeypatch, _StubStore( - pages=[{"metadata": {"file_path": "quill-store.md"}, "text": "x" * 5000}] + pages=[ + { + "metadata": {"title": "quill-store.md", "source": "quill-store.md"}, + "text": "x" * 5000, + } + ] ), ) result = tool.execute_action("read_entity_pages", entity="Quill") assert "--- quill-store.md ---" in result + # Only what the model reads is truncated. assert len(result) < 3000 # Accumulated so the answer can cite what the walk actually read. assert tool.retrieved_docs[0]["title"] == "quill-store.md" + assert len(tool.retrieved_docs[0]["text"]) == 5000 + + def test_page_labels_match_what_the_retrievers_record(self, monkeypatch): + """A page read here and the same chunk retrieved by internal_search are + one document. The citation manager keys on (source, title), so labels + derived differently give the same document two citation numbers.""" + from docsgpt.retriever.labels import labels_from_metadata + + metadata = {"title": "Quill Store", "source": "quill-store.md"} + text = "Quill is a write-ahead store." + tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": metadata, "text": text}])) + + tool.execute_action("read_entity_pages", entity="Quill") + + expected = labels_from_metadata(metadata, text, "src-1") + doc = tool.retrieved_docs[0] + assert {k: doc[k] for k in ("title", "source", "filename")} == expected + # The full chunk text, so the doc dedupes against the retriever's copy; + # only what the model reads is truncated. + assert doc["text"] == text + + def test_a_page_with_no_metadata_falls_back_to_its_source_id(self, monkeypatch): + tool = _tool(monkeypatch, _StubStore(pages=[{"metadata": {}, "text": "body"}])) + + tool.execute_action("read_entity_pages", entity="Quill") + + assert tool.retrieved_docs[0]["source"] == "src-1" def test_pages_absent_is_stated_plainly(self, monkeypatch): tool = _tool(monkeypatch, _StubStore(pages=[])) @@ -150,6 +187,7 @@ class TestWiring: def test_not_added_when_the_sources_have_no_graph(self, monkeypatch): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") monkeypatch.setattr( "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: False ) @@ -161,6 +199,7 @@ class TestWiring: def test_added_with_its_sentinel_id_and_source_config(self, monkeypatch): monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") monkeypatch.setattr( "docsgpt.agents.tools.graph_search.sources_have_graph", lambda source: True ) @@ -246,3 +285,66 @@ class TestSourcesHaveGraph: assert sources_have_graph({"active_docs": []}) is False self._patch_counts(monkeypatch, error=RuntimeError("no pgvector")) assert sources_have_graph({"active_docs": ["a"]}) is False + + +class TestGraphsMustBeAvailable: + """The graph lives in the pgvector store, so the flag alone is not enough. + + With another vector store configured the graph tables are not the ones the + sources were ingested into; everything else in the app asks + ``graphrag_available()``, which requires both. + """ + + def test_the_tool_is_not_offered_without_pgvector(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "faiss") + monkeypatch.setattr( + "docsgpt.agents.tools.graph_search.sources_have_graph", + lambda source: pytest.fail("must not reach the database"), + ) + tools = {} + add_graph_search_tool(tools, {"source": SOURCE}) + assert tools == {} + + def test_actions_report_it_rather_than_querying(self, monkeypatch): + monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", True) + monkeypatch.setattr(settings, "VECTOR_STORE", "faiss") + tool = GraphSearchTool({"source": SOURCE}) + tool._store = _StubStore(nodes=[{"name": "Quill", "distance": 0.1}]) + + assert "not enabled" in tool.execute_action("search_entities", query="quill") + + +class TestPooledConnection: + """The tool is cached for the whole agent run; its connection must not be.""" + + class _ClosingStore(_StubStore): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.closed = 0 + + def close(self): + self.closed += 1 + + def test_the_connection_goes_back_after_each_action(self, monkeypatch): + store = self._ClosingStore(relationships=[{"source": "A", "target": "B", "type": "r"}]) + tool = _tool(monkeypatch, store) + + tool.execute_action("get_relationships", entity="A") + + # Held open, one pooled connection would be pinned across every LLM + # round trip of the run. + assert store.closed == 1 + assert tool._store is None + + def test_a_failing_action_still_releases_it(self, monkeypatch): + class _Broken(self._ClosingStore): + def entity_relationships(self, source_id, name, limit=25): + raise RuntimeError("connection lost") + + store = _Broken() + tool = _tool(monkeypatch, store) + + assert tool.execute_action("get_relationships", entity="A") == "The graph lookup failed." + assert store.closed == 1 + assert tool._store is None From 2b6d4d509e13d4be7d667b3d0010f4a77615b583 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:04 +0100 Subject: [PATCH 106/130] fix(graphrag): return a chunk once from entity_pages The subject flag was in the GROUP BY, so a chunk two matching entities link -- one the page is about, one merely mentioned in it -- came back as two identical pages and spent the caller's page budget twice on the same text. It is aggregated with bool_or now, which is what the ordering wanted anyway. Covered by a live test against a pgvector-shaped documents table: the graph tables alone cannot answer this query, so nothing exercised it before. --- docsgpt/graphrag/store.py | 8 +++-- tests/graphrag/test_store.py | 63 ++++++++++++++++++++++++++++++++++++ 2 files changed, 69 insertions(+), 2 deletions(-) diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index 1804d929..0b73d031 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -1265,13 +1265,17 @@ class GraphStore: sql.SQL( """ SELECT d.{metadata}, d.{text}, - (lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject + bool_or(lower(n.name) = %s OR lower(n.name) LIKE %s) AS is_subject FROM graph_node_chunks gc JOIN graph_nodes n ON n.id = gc.node_id JOIN {table} d ON d.id::text = gc.chunk_id WHERE gc.source_id = %s AND d.{source} = %s AND (lower(n.name) = %s OR lower(n.name) LIKE %s OR n.name ILIKE %s) - GROUP BY d.{metadata}, d.{text}, is_subject + -- One page per chunk. Grouping on the subject flag as well + -- split a chunk two entities link -- one naming it, one + -- merely mentioned -- into two identical pages, spending + -- the caller's page budget twice on the same text. + GROUP BY d.{metadata}, d.{text} ORDER BY is_subject DESC, (d.{text} ILIKE %s) DESC LIMIT %s; """ diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 062db1f2..51002d18 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -19,6 +19,7 @@ import uuid from unittest.mock import MagicMock, patch import pytest +from psycopg.types.json import Jsonb import docsgpt.graphrag.store as store_module from docsgpt.vectorstore import pgconn @@ -661,6 +662,68 @@ class TestGraphStoreParameterization: assert params[-1] == embedding +@pytest.mark.integration +class TestEntityPagesLive: + """``entity_pages`` against a real pgvector-shaped table. + + The graph tables alone cannot answer it: the rows it returns live in the + documents table the sources were ingested into, so the test creates a + minimal one with the same column names ``PGVectorStore`` uses. + """ + + @pytest.fixture + def store(self, postgresql): + store = GraphStore(connection_string=_ephemeral_dsn(postgresql.info)) + try: + store._ensure_tables() + except Exception as exc: + pytest.skip(f"pgvector extension unavailable: {exc}") + conn = store._get_connection() + cursor = conn.cursor() + cursor.execute( + """ + CREATE TABLE IF NOT EXISTS documents ( + id SERIAL PRIMARY KEY, + text TEXT, + metadata JSONB, + source_id TEXT + ); + """ + ) + conn.commit() + cursor.close() + yield store + store.close() + + def test_a_page_linked_by_two_entities_is_returned_once(self, store): + """One chunk, two nodes whose names both match: an exact hit and a + mention. They differ only in whether the page is *about* the entity, so + grouping on that flag returned the same page twice and spent a quarter + of the page budget on it.""" + source_id = str(uuid.uuid4()) + conn = store._get_connection() + cursor = conn.cursor() + cursor.execute( + "INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;", + ("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id), + ) + chunk_id = str(cursor.fetchone()[0]) + conn.commit() + cursor.close() + try: + subject = store.upsert_node(source_id, "Quill", "quill") + mention = store.upsert_node(source_id, "Legacy Quill", "legacy quill") + store.link_node_chunk(source_id, subject, chunk_id) + store.link_node_chunk(source_id, mention, chunk_id) + + pages = store.entity_pages(source_id, "Quill", limit=4) + + assert [page["text"] for page in pages] == ["Quill is a write-ahead store."] + assert pages[0]["metadata"] == {"title": "quill.md"} + finally: + store.delete_by_source(source_id) + + @pytest.mark.unit class TestGraphReadQueries: """The reads behind fact seeding and the agent's graph tool, without a DB. From ef5b718f371da4180c1667fe4aaf0ce7ef1ecb81 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:45:05 +0100 Subject: [PATCH 107/130] fix(graphrag): drop entities whose name normalizes to nothing ``canonical_name`` answers "" for a punctuation-only name, which callers are meant to read as "no entity" -- ``_resolve_endpoint`` already does. Entity extraction did not, and nodes merge on that key, so every such entity in a source collapsed onto one shared node that belonged to none of them. --- docsgpt/graphrag/extraction.py | 10 ++++++++-- tests/graphrag/test_extraction.py | 21 +++++++++++++++++++++ 2 files changed, 29 insertions(+), 2 deletions(-) diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index 5ddbe1e1..daa8f910 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -424,7 +424,7 @@ def extract_graph_for_source( def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: - """Normalize the LLM's entity dicts (drop nameless ones).""" + """Normalize the LLM's entity dicts (drop the ones with no usable name).""" entities = [] for e in raw_entities: if not isinstance(e, dict): @@ -432,10 +432,16 @@ def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]: name = str(e.get("name", "")).strip() if not name: continue + normalized_name = normalize_entity_name(name) + if not normalized_name: + # A punctuation-only name normalizes to nothing, and nodes merge on + # that key: keeping it collapses every such entity onto one shared + # node. The relationship side already drops them. + continue entities.append( { "name": name, - "normalized_name": normalize_entity_name(name), + "normalized_name": normalized_name, "type": str(e.get("type") or "") or None, "description": str(e.get("description") or "") or None, } diff --git a/tests/graphrag/test_extraction.py b/tests/graphrag/test_extraction.py index 05bb0a40..1125184f 100644 --- a/tests/graphrag/test_extraction.py +++ b/tests/graphrag/test_extraction.py @@ -1085,6 +1085,27 @@ class TestSummaryCountFailure: assert store.count_nodes.call_args.kwargs.get("strict") is True +@pytest.mark.unit +class TestEntityNormalization: + """An entity whose name normalizes to nothing is not an entity. + + ``canonical_name`` answers "" for a punctuation-only name, and nodes are + merged on that key, so keeping them collapses every such entity onto one + shared node. ``_resolve_endpoint`` already drops them on the relationship + side. + """ + + @pytest.mark.parametrize("name", ["!!!", "--", "?", " *** "]) + def test_a_name_that_normalizes_to_nothing_is_dropped(self, name): + assert extraction_module._build_entities([{"name": name}]) == [] + + def test_real_names_survive(self): + built = extraction_module._build_entities( + [{"name": "Quill Store"}, {"name": "!!!"}, {"name": "Alder"}] + ) + assert [e["normalized_name"] for e in built] == ["quill store", "alder"] + + @pytest.mark.unit class TestParsing: def test_parses_embedded_json(self): From 692edcdd4ea24ee053b21e75ba0dddc4760a87e5 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:54:55 +0100 Subject: [PATCH 108/130] test(agents): give the graph tool tests the vector store they need The tool now asks graphrag_available(), which wants pgvector as well as the flag. One case set only the flag and passed locally off a dev .env, then failed on CI's faiss default. An autouse fixture sets it for the module, so a case that forgets fails for its own reason; the two about a different vector store override it themselves. --- tests/graphrag/test_graph_search_tool.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/tests/graphrag/test_graph_search_tool.py b/tests/graphrag/test_graph_search_tool.py index 465a1c25..ca2f1330 100644 --- a/tests/graphrag/test_graph_search_tool.py +++ b/tests/graphrag/test_graph_search_tool.py @@ -22,6 +22,17 @@ from docsgpt.core.settings import settings SOURCE = {"active_docs": ["src-1"]} +@pytest.fixture(autouse=True) +def _graph_store_is_pgvector(monkeypatch): + """The graph only exists under pgvector, and CI's default is faiss. + + Set here rather than per test so a case that forgets it fails for its own + reason instead of the gate; the cases about a different vector store + override it in the test body. + """ + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector") + + class _StubStore: def __init__(self, nodes=None, relationships=None, pages=None): self._nodes = nodes or [] From 63d93fc4b9ba5baccd268ffadd4169b79a431a53 Mon Sep 17 00:00:00 2001 From: Alex Date: Sun, 20 Sep 2026 10:54:55 +0100 Subject: [PATCH 109/130] test(graphrag): pin that identical pages collapse into one Review asked whether two document rows with the same text and metadata should stay two pages. They should not: the caller gets four pages to hand a model, and a crawl that ingested the same text twice would spend two of them on it. The rows differ only by an id the model never sees. Pinned either way now. --- tests/graphrag/test_store.py | 28 ++++++++++++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/tests/graphrag/test_store.py b/tests/graphrag/test_store.py index 51002d18..d9de877a 100644 --- a/tests/graphrag/test_store.py +++ b/tests/graphrag/test_store.py @@ -724,6 +724,34 @@ class TestEntityPagesLive: store.delete_by_source(source_id) + def test_two_chunks_with_identical_text_collapse_into_one_page(self, store): + """Deliberate: the caller gets at most four pages to hand a model, and a + crawl that ingested the same text twice would spend two of them saying + the same thing. The rows differ only by an id the model never sees.""" + source_id = str(uuid.uuid4()) + conn = store._get_connection() + cursor = conn.cursor() + chunk_ids = [] + for _ in range(2): + cursor.execute( + "INSERT INTO documents (text, metadata, source_id) VALUES (%s, %s, %s) RETURNING id;", + ("Quill is a write-ahead store.", Jsonb({"title": "quill.md"}), source_id), + ) + chunk_ids.append(str(cursor.fetchone()[0])) + conn.commit() + cursor.close() + try: + node = store.upsert_node(source_id, "Quill", "quill") + for chunk_id in chunk_ids: + store.link_node_chunk(source_id, node, chunk_id) + + pages = store.entity_pages(source_id, "Quill", limit=4) + + assert [page["text"] for page in pages] == ["Quill is a write-ahead store."] + finally: + store.delete_by_source(source_id) + + @pytest.mark.unit class TestGraphReadQueries: """The reads behind fact seeding and the agent's graph tool, without a DB. From 7ddb5f64003f1f3b90ae49c9a248fbbdca1d054f Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sun, 20 Sep 2026 12:12:42 +0100 Subject: [PATCH 110/130] fix(pat): align replay scopes, keep 404/405, retire expired names, document the PAT_ENABLED switch The ASGI message events route now accepts the same scopes as its Flask sibling (conversations:read or chat:run) through a shared constant. A token request that fails routing gets Flask's 404/405 instead of a 403. Creating a token retires an expired token that still held the name. allowed_ids uses is_pat instead of a bare literal. PAT_ENABLED is documented as the master switch it is: turning it off stops every existing token from authenticating. The docs explain that sources are matched by name (oldest wins) and point CI flows at sources upload --replace. --- docs/content/Deploying/Settings-Reference.mdx | 2 +- .../Extensions/personal-access-tokens.mdx | 14 +++++++++++--- docsgpt/api/asgi_auth.py | 9 +++++---- docsgpt/api/async_sse.py | 4 +++- docsgpt/api/pat/routes.py | 1 + docsgpt/api/pat/rules.py | 19 +++++++++++++++---- docsgpt/core/settings/auth.py | 6 ++++-- .../db/repositories/personal_access_tokens.py | 17 +++++++++++++++++ tests/api/test_asgi_auth.py | 16 ++++++++++++++++ tests/api/test_pat_routes.py | 9 +++++++++ tests/api/test_pat_rules.py | 5 +++++ .../test_personal_access_tokens.py | 18 ++++++++++++++++++ 12 files changed, 105 insertions(+), 15 deletions(-) diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 9266a90e..2229afc4 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -136,7 +136,7 @@ Bearer token for IdP SCIM clients (required when SCIM is enabled). Type `bool`, default `true`. -Allow users to create personal access tokens. Tokens are only issued under AUTH_TYPE=oidc or unset (None); simple_jwt and session_jwt have no stable user identity to bind a token to. +Master switch for personal access tokens. When false, no token can be created AND every existing token stops authenticating immediately (pipelines using them get 401); tokens are kept and work again when re-enabled. Tokens are only available under AUTH_TYPE=oidc or unset (None); switching to simple_jwt or session_jwt disables them the same way. ### `PAT_DEFAULT_LIFETIME_DAYS` diff --git a/docs/content/Extensions/personal-access-tokens.mdx b/docs/content/Extensions/personal-access-tokens.mdx index c533b528..909b9ef6 100644 --- a/docs/content/Extensions/personal-access-tokens.mdx +++ b/docs/content/Extensions/personal-access-tokens.mdx @@ -69,6 +69,8 @@ curl -X POST "$DOCSGPT_URL/api/import_agent/plan" \ -d "$(jq -Rs '{yaml: .}' support-bot.agent.yaml)" ``` +Sources are matched by **name**, and when several of your sources share a name the oldest one wins. A pipeline that re-uploads documentation on every push should therefore upload with `docsgpt-cli sources upload ... --wait --replace` (which removes the older same-named sources) and run `agents apply` afterwards, or the agent stays bound to the first upload. + `/api/import_agent/plan` is a dry run that reports whether the agent would be created or updated and how each referenced source, tool and prompt resolves. `/api/import_agent` applies it. An agent is matched by `metadata.id`, then `metadata.slug`; when nothing matches, a new draft agent is created. ### GitHub Actions example @@ -83,7 +85,9 @@ jobs: env: DOCSGPT_URL: ${{ vars.DOCSGPT_URL }} DOCSGPT_TOKEN: ${{ secrets.DOCSGPT_TOKEN }} - run: docsgpt-cli agents apply -f agents/ + run: | + docsgpt-cli sources upload docs/*.md --name "Product docs" --wait --replace + docsgpt-cli agents apply -f agents/ ``` ## Scopes @@ -149,19 +153,23 @@ Restrictions and `agents apply`: a token restricted to specific agents can apply - Tokens of a deactivated user (through the admin API or SCIM) stop working immediately and work again if the user is reactivated. - Token creation and revocation are recorded in the authentication audit log (`pat_created`, `pat_revoked`), visible to admins. +An expired token's name can be reused: creating a token with that name retires the expired one. + Each user may hold up to `PAT_MAX_PER_USER` live tokens (25 by default). The token list shows when and from which IP address each token was last used. ## Operator settings | Setting | Default | Purpose | | --- | --- | --- | -| `PAT_ENABLED` | `true` | Allow users to create and use personal access tokens | +| `PAT_ENABLED` | `true` | Master switch. When `false`, tokens cannot be created **and every existing token stops authenticating immediately** | | `PAT_DEFAULT_LIFETIME_DAYS` | `90` | Lifetime of a token created without an explicit expiry | | `PAT_MAX_LIFETIME_DAYS` | `365` | Longest lifetime a user may request | | `PAT_ALLOW_NON_EXPIRING` | `false` | Let users create tokens that never expire | | `PAT_MAX_PER_USER` | `25` | Maximum number of live tokens per user | -Personal access tokens need a stable user identity, so they are available with `AUTH_TYPE=oidc` and with authentication disabled (single-user self-hosting). They are not available with `simple_jwt` or `session_jwt`. See the [Settings Reference](/Deploying/Settings-Reference) for details. +Personal access tokens need a stable user identity, so they are available with `AUTH_TYPE=oidc` and with authentication disabled (single-user self-hosting). They are not available with `simple_jwt` or `session_jwt`. + +Turning `PAT_ENABLED` off, or switching `AUTH_TYPE` to `simple_jwt` or `session_jwt`, is not limited to the settings page: every pipeline that uses a token starts getting `401` right away. Tokens are not deleted and work again once the setting is restored. See the [Settings Reference](/Deploying/Settings-Reference) for details. ## Management API diff --git a/docsgpt/api/asgi_auth.py b/docsgpt/api/asgi_auth.py index a65e83ab..7c6c2cb9 100644 --- a/docsgpt/api/asgi_auth.py +++ b/docsgpt/api/asgi_auth.py @@ -9,7 +9,7 @@ from __future__ import annotations import uuid from contextvars import Token -from typing import Optional, Tuple +from typing import Optional, Sequence, Tuple, Union import anyio from starlette.requests import Request @@ -23,13 +23,13 @@ from docsgpt.core.settings import settings async def authenticate( - request: Request, *, pat_scope: Optional[str] = None + request: Request, *, pat_scope: Union[str, Sequence[str], None] = None ) -> Tuple[Optional[dict], Optional[JSONResponse]]: """Decode the caller's JWT the way Flask's ``authenticate_request`` does. Args: request: The incoming Starlette request. - pat_scope: Scope a personal access token needs for this route. Left + pat_scope: Scope (or any-of scopes) a personal access token needs for this route. Left unset, the route rejects PATs outright (deny by default, matching the Flask rule table in ``docsgpt/api/pat/rules.py``). A token with a resource filter is always rejected. @@ -48,7 +48,8 @@ async def authenticate( if is_pat(decoded): # A PAT lookup already excludes revoked tokens and deactivated users, # so the session denylist below does not apply to it. - if pat_scope is None or pat_scope not in (decoded.get("scopes") or []): + accepted = (pat_scope,) if isinstance(pat_scope, str) else tuple(pat_scope or ()) + if not set(accepted).intersection(decoded.get("scopes") or []): return None, JSONResponse( {"success": False, "message": "Token lacks the required scope", "error": "insufficient_scope"}, status_code=403, diff --git a/docsgpt/api/async_sse.py b/docsgpt/api/async_sse.py index 1bb3af1d..e3ae15d8 100644 --- a/docsgpt/api/async_sse.py +++ b/docsgpt/api/async_sse.py @@ -25,6 +25,7 @@ from starlette.responses import Response from starlette.routing import Route from docsgpt.api.asgi_auth import authenticate, bind_log_context, json_error +from docsgpt.api.pat.rules import MESSAGE_REPLAY_SCOPES from docsgpt.api.asgi_stream import sse_response from docsgpt.core.settings import settings from docsgpt.storage.db.session import db_readonly @@ -94,7 +95,8 @@ async def stream_message_events(request: Request) -> Response: """ # Same JWT decoder and OIDC revocation check as the Flask routes. With # AUTH_TYPE unset the caller resolves to ``{"sub": "local"}``. - decoded, error = await authenticate(request, pat_scope="chat:run") + # Same scopes as its Flask sibling GET /api/messages//tail. + decoded, error = await authenticate(request, pat_scope=MESSAGE_REPLAY_SCOPES) if error is not None: return error user_id = decoded.get("sub") if isinstance(decoded, dict) else None diff --git a/docsgpt/api/pat/routes.py b/docsgpt/api/pat/routes.py index ad22cf6d..bb3ea8c5 100644 --- a/docsgpt/api/pat/routes.py +++ b/docsgpt/api/pat/routes.py @@ -139,6 +139,7 @@ class PersonalAccessTokens(Resource): return _error( f"Token limit reached ({settings.PAT_MAX_PER_USER}); revoke one first", 409 ) + repo.retire_expired_name(user_id, name) if repo.name_in_use(user_id, name): return _error("A token with this name already exists", 409) row = repo.create( diff --git a/docsgpt/api/pat/rules.py b/docsgpt/api/pat/rules.py index 6794390b..88f1df76 100644 --- a/docsgpt/api/pat/rules.py +++ b/docsgpt/api/pat/rules.py @@ -25,6 +25,8 @@ import uuid from dataclasses import dataclass from typing import Any, Callable, Iterable, Optional +from docsgpt.api.pat.tokens import is_pat + VIEW, QUERY, JSON, FORM, BODY = "view", "query", "json", "form", "body" Locator = tuple[str, str] @@ -53,6 +55,11 @@ def _rule(scope: Optional[str] = None, *ids: Locator, any_of: tuple[str, ...] = return Rule(scopes=scopes, family=family, ids=tuple(ids), **kwargs) +#: Scopes that admit a token to message replay. Shared by the Flask tail route +#: below and its ASGI sibling GET /api/messages//events (docsgpt/api/async_sse.py), +#: which sits outside this table. +MESSAGE_REPLAY_SCOPES = ("conversations:read", "chat:run") + _ALL_FAMILIES = ("agents", "sources", "prompts", "tools", "workflows") # Ids of other families that agent create/update accept in their JSON-or-form body. @@ -220,7 +227,7 @@ RULES: dict[tuple[str, str], Rule] = { ("/api/get_single_conversation", "GET"): _rule("conversations:read", blocked_by=("agents",)), # A message cannot be tied to an allowlist from here, so any restricted token is kept out. ("/api/messages//tail", "GET"): _rule( - any_of=("conversations:read", "chat:run"), family=None, blocked_by=_ALL_FAMILIES + any_of=MESSAGE_REPLAY_SCOPES, family=None, blocked_by=_ALL_FAMILIES ), ("/api/delete_conversation", "POST"): _rule("conversations:write", blocked_by=("agents",)), ("/api/delete_all_conversations", "GET"): _rule("conversations:write", blocked_by=("agents",)), @@ -360,7 +367,11 @@ def _all_allowed(ids: Iterable[str], allowed: Iterable[str]) -> bool: def authorize(request, decoded_token: dict) -> Optional[tuple[dict, int]]: """Check a PAT request against the table. ``None`` allows; otherwise ``(body, status)``.""" url_rule = getattr(request, "url_rule", None) - rule = RULES.get((url_rule.rule, request.method)) if url_rule is not None else None + if url_rule is None: + # Routing failed (unknown path or wrong method): no view will run, so + # let Flask answer 404/405 instead of masking it with a 403. + return None + rule = RULES.get((url_rule.rule, request.method)) if rule is None: return ( { @@ -420,8 +431,8 @@ def allowed_ids(request, family: str) -> Optional[set[str]]: Used by listing routes (``listing=True``) and by ``in_route`` handlers. """ - decoded = getattr(request, "decoded_token", None) or {} - if decoded.get("auth_method") != "pat": + decoded = getattr(request, "decoded_token", None) + if not is_pat(decoded): return None ids = (decoded.get("resource_filter") or {}).get(family) if ids is None: diff --git a/docsgpt/core/settings/auth.py b/docsgpt/core/settings/auth.py index 7471dc5a..706a3339 100644 --- a/docsgpt/core/settings/auth.py +++ b/docsgpt/core/settings/auth.py @@ -89,8 +89,10 @@ class AuthSettings(SettingsGroup): PAT_ENABLED: bool = Field( default=True, description=( - "Allow users to create personal access tokens. Tokens are only issued under AUTH_TYPE=oidc or " - "unset (None); simple_jwt and session_jwt have no stable user identity to bind a token to." + "Master switch for personal access tokens. When false, no token can be created AND every existing " + "token stops authenticating immediately (pipelines using them get 401); tokens are kept and work " + "again when re-enabled. Tokens are only available under AUTH_TYPE=oidc or unset (None); switching " + "to simple_jwt or session_jwt disables them the same way." ), ) PAT_DEFAULT_LIFETIME_DAYS: int = Field( diff --git a/docsgpt/storage/db/repositories/personal_access_tokens.py b/docsgpt/storage/db/repositories/personal_access_tokens.py index 56fb5caa..efbe250b 100644 --- a/docsgpt/storage/db/repositories/personal_access_tokens.py +++ b/docsgpt/storage/db/repositories/personal_access_tokens.py @@ -87,6 +87,23 @@ class PersonalAccessTokensRepository: {"user_id": user_id}, ).scalar_one() + def retire_expired_name(self, user_id: str, name: str) -> int: + """Revoke an expired token holding ``name`` so the name can be reused. + + An expired token can no longer authenticate but keeps ``status = 'active'``, + and the unique index on live names would otherwise reserve its name forever. + """ + result = self._conn.execute( + text( + "UPDATE personal_access_tokens " + "SET status = 'revoked', revoked_at = now(), revoke_reason = 'expired' " + "WHERE user_id = :user_id AND name = :name AND status = 'active' " + "AND expires_at IS NOT NULL AND expires_at <= now()" + ), + {"user_id": user_id, "name": name}, + ) + return result.rowcount + def name_in_use(self, user_id: str, name: str) -> bool: return ( self._conn.execute( diff --git a/tests/api/test_asgi_auth.py b/tests/api/test_asgi_auth.py index d40fae0e..fa1768ec 100644 --- a/tests/api/test_asgi_auth.py +++ b/tests/api/test_asgi_auth.py @@ -164,3 +164,19 @@ class TestPersonalAccessTokens: assert decoded is None assert error.status_code == 403 assert json.loads(error.body)["error"] == "resource_not_allowed" + + async def test_any_of_several_scopes_admits_the_token(self): + claims = dict(_PAT_CLAIMS, scopes=["conversations:read"]) + with patch.object(asgi_auth, "handle_auth", return_value=claims): + decoded, error = await asgi_auth.authenticate( + _request(), pat_scope=("conversations:read", "chat:run") + ) + assert error is None and decoded["sub"] == "alice" + + async def test_message_events_accepts_the_same_scopes_as_message_tail(self): + from docsgpt.api import async_sse + from docsgpt.api.pat import rules + + tail = rules.RULES[("/api/messages//tail", "GET")] + assert tail.scopes == rules.MESSAGE_REPLAY_SCOPES + assert async_sse.MESSAGE_REPLAY_SCOPES is rules.MESSAGE_REPLAY_SCOPES diff --git a/tests/api/test_pat_routes.py b/tests/api/test_pat_routes.py index be07b7d5..f7592cf4 100644 --- a/tests/api/test_pat_routes.py +++ b/tests/api/test_pat_routes.py @@ -127,6 +127,15 @@ class TestCreate: assert _create(client).status_code == 201 assert _create(client).status_code == 409 + def test_expired_token_does_not_reserve_its_name(self, client, db): + assert _create(client).status_code == 201 + db.execute(text("UPDATE personal_access_tokens SET expires_at = now() - interval '1 day'")) + assert _create(client).status_code == 201 + rows = db.execute( + text("SELECT status, revoke_reason FROM personal_access_tokens ORDER BY created_at") + ).all() + assert [tuple(r) for r in rows] == [("revoked", "expired"), ("active", None)] + def test_per_user_cap(self, client, db, monkeypatch): monkeypatch.setattr(pat_tokens.settings, "PAT_MAX_PER_USER", 1) assert _create(client, name="one").status_code == 201 diff --git a/tests/api/test_pat_rules.py b/tests/api/test_pat_rules.py index c6c9f0b8..4bdc7070 100644 --- a/tests/api/test_pat_rules.py +++ b/tests/api/test_pat_rules.py @@ -114,6 +114,11 @@ class TestScopeEnforcement: response = _call(client, "GET", "/api/admin/users", _claims(list(SCOPES))) assert _denied(response) == "not_available_to_tokens" + def test_unknown_path_and_wrong_method_keep_their_own_status(self, client): + claims = _claims(list(SCOPES)) + assert _call(client, "GET", "/api/no_such_route", claims).status_code == 404 + assert _call(client, "DELETE", "/api/get_agents", claims).status_code == 405 + def test_missing_scope_is_refused_and_names_the_scope(self, client): response = _call(client, "GET", "/api/get_agents", _claims(["sources:read"])) assert _denied(response) == "insufficient_scope" diff --git a/tests/storage/db/repositories/test_personal_access_tokens.py b/tests/storage/db/repositories/test_personal_access_tokens.py index 0102e625..8a6848a3 100644 --- a/tests/storage/db/repositories/test_personal_access_tokens.py +++ b/tests/storage/db/repositories/test_personal_access_tokens.py @@ -64,6 +64,24 @@ class TestUniqueness: assert not repo.name_in_use("u1", "ci") assert _create(repo, token_hash="h2")["name"] == "ci" + def test_retire_expired_name_only_touches_expired_tokens_with_that_name(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + past = datetime.now(timezone.utc) - timedelta(days=1) + _create(repo, name="ci", token_hash="h1", expires_at=past) + _create(repo, name="other", token_hash="h2", expires_at=past) + _create(repo, user_id="u2", name="ci", token_hash="h3", expires_at=past) + assert repo.retire_expired_name("u1", "ci") == 1 + assert not repo.name_in_use("u1", "ci") + assert repo.name_in_use("u1", "other") and repo.name_in_use("u2", "ci") + + def test_retire_expired_name_leaves_live_tokens(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + _create(repo, token_hash="h1") + _create(repo, name="later", token_hash="h2", expires_at=datetime.now(timezone.utc) + timedelta(days=1)) + assert repo.retire_expired_name("u1", "ci") == 0 + assert repo.retire_expired_name("u1", "later") == 0 + assert repo.name_in_use("u1", "ci") and repo.name_in_use("u1", "later") + def test_same_name_for_other_user_is_fine(self, pg_conn): repo = PersonalAccessTokensRepository(pg_conn) _create(repo, user_id="u1", token_hash="h1") From b1ff8f516e60fc787014feb6c49cf6ee27410865 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sun, 20 Sep 2026 13:32:46 +0100 Subject: [PATCH 111/130] fix(frontend): access tokens dates and names render unescaped, align the scopes block i18next HTML-escaped interpolated values, so expiry dates showed as 26/09/2026 and token names containing & or quotes were mangled in the created and revoke dialogs; React already escapes on render, so those strings opt out like utils/streamingStatusUtils does. The scopes fieldset drops the browser's default padding so it lines up with the other fields, and the native checkboxes follow the dark colour scheme. --- frontend/src/modals/AccessTokenCreatedModal.tsx | 6 +++++- frontend/src/modals/CreateAccessTokenModal.tsx | 6 ++++-- frontend/src/settings/PersonalAccessTokens.tsx | 11 ++++++++--- frontend/src/settings/accessTokenUtils.test.ts | 7 +++++++ frontend/src/settings/accessTokenUtils.ts | 8 ++++++++ 5 files changed, 32 insertions(+), 6 deletions(-) diff --git a/frontend/src/modals/AccessTokenCreatedModal.tsx b/frontend/src/modals/AccessTokenCreatedModal.tsx index 66026edc..2b627f93 100644 --- a/frontend/src/modals/AccessTokenCreatedModal.tsx +++ b/frontend/src/modals/AccessTokenCreatedModal.tsx @@ -5,6 +5,7 @@ import { baseURL } from '../api/client'; import CopyButton from '../components/CopyButton'; import { Button } from '../components/ui/button'; import { Modal } from '../components/ui/modal'; +import { NO_ESCAPE } from '../settings/accessTokenUtils'; interface AccessTokenCreatedModalProps { /** Plaintext secret; `null` keeps the modal closed. Held only in the parent's component state. */ @@ -49,7 +50,10 @@ export default function AccessTokenCreatedModal({ {t('settings.accessTokens.created.title')}

- {t('settings.accessTokens.created.subtitle', { name })} + {t('settings.accessTokens.created.subtitle', { + name, + ...NO_ESCAPE, + })}

diff --git a/frontend/src/modals/CreateAccessTokenModal.tsx b/frontend/src/modals/CreateAccessTokenModal.tsx index 438403a8..87e39a2d 100644 --- a/frontend/src/modals/CreateAccessTokenModal.tsx +++ b/frontend/src/modals/CreateAccessTokenModal.tsx @@ -30,6 +30,7 @@ import { expiryOptions, groupScopesByFamily, isScopeImplied, + NO_ESCAPE, NO_EXPIRY, PICKER_FAMILIES, ResourceOption, @@ -306,12 +307,13 @@ export default function CreateAccessTokenModal({ date: formatDateOnly( new Date(Date.now() + expiry * DAY_MS).toISOString(), ), + ...NO_ESCAPE, })}

-
+
{t('settings.accessTokens.create.scopes')} * @@ -355,7 +357,7 @@ export default function CreateAccessTokenModal({ checked={checked} disabled={implied} onChange={() => toggleScope(scope.name)} - className="accent-primary mt-0.5 size-4 shrink-0 rounded-sm border-gray-300 bg-transparent" + className="accent-primary mt-0.5 size-4 shrink-0 rounded-sm border-gray-300 bg-transparent dark:[color-scheme:dark]" /> diff --git a/frontend/src/settings/PersonalAccessTokens.tsx b/frontend/src/settings/PersonalAccessTokens.tsx index d81416b0..7439e567 100644 --- a/frontend/src/settings/PersonalAccessTokens.tsx +++ b/frontend/src/settings/PersonalAccessTokens.tsx @@ -33,6 +33,7 @@ import { formatDateOnly, formatDateTime } from '../utils/dateTimeUtils'; import { countLiveTokens, expiryStatus, + NO_ESCAPE, relativeTime, restrictionCounts, } from './accessTokenUtils'; @@ -197,8 +198,8 @@ export default function PersonalAccessTokens() { > ); }; @@ -210,7 +211,10 @@ export default function PersonalAccessTokens() { size="sm" className="text-destructive hover:text-destructive border-destructive/40 hover:bg-destructive/10 rounded-full px-4" onClick={() => requestRevoke(item)} - aria-label={t('settings.accessTokens.revokeAria', { name: item.name })} + aria-label={t('settings.accessTokens.revokeAria', { + name: item.name, + ...NO_ESCAPE, + })} > {t('settings.accessTokens.revoke')} @@ -431,6 +435,7 @@ export default function PersonalAccessTokens() { { expect(toResourceOptions('agents', { success: false })).toEqual([]); }); }); + +describe('NO_ESCAPE', () => { + it('turns off i18next HTML escaping so dates and names render as typed', () => { + expect(NO_ESCAPE).toEqual({ interpolation: { escapeValue: false } }); + }); +}); diff --git a/frontend/src/settings/accessTokenUtils.ts b/frontend/src/settings/accessTokenUtils.ts index 134d347d..9e7e56c3 100644 --- a/frontend/src/settings/accessTokenUtils.ts +++ b/frontend/src/settings/accessTokenUtils.ts @@ -222,3 +222,11 @@ export function toResourceOptions( })) .sort((a, b) => a.label.localeCompare(b.label)); } + +/** + * i18next HTML-escapes interpolated values by default, which turns a date like + * 26/09/2026 into `26/09/2026` and mangles token names containing + * `&` or `'`. React already escapes on render, so opt out for those values + * (same approach as utils/streamingStatusUtils.ts). + */ +export const NO_ESCAPE = { interpolation: { escapeValue: false } } as const; From 94b5924e367d9d2e1f5ebba6eef6d62d952a294d Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sun, 20 Sep 2026 19:46:43 +0100 Subject: [PATCH 112/130] fix(pat): close restricted-token paths through workflows, chat, schedules and conversations A resource restriction was checked on ids in the request, not on what the addressed row pulls in or belongs to. Closed: - Workflow writes for tokens restricted on sources, tools or prompts (a graph names those inside its nodes), and attaching a workflow to an agent unless the token is restricted on workflows too. - Chat for tools-restricted tokens (chat executes tools; rejected at token creation as well), and agent-less chat for tokens restricted on prompts or workflows. - conversation_id on chat: it must belong to the agent being run, or to no agent for agent-less chat. Otherwise the server continued, appended to, or resumed pending tool calls of another agent's conversation. - Schedules for tokens restricted on anything but agents; schedule-id routes for every restricted token. - Conversations and analytics for every restricted token, not only agent-restricted ones. Also: create, first publish and adopt return the agent API key masked to a token without agents:keys; token ids must be canonical UUIDs (urn:uuid: gave a 500); an expired token is reported as expired; token creation takes a per-user advisory lock so the cap cannot be raced; admin revoke-sessions writes a pat_revoked event per token. The UI drops a row whose revoke returns 404 and does not offer a tools restriction next to chat:run. --- .../Extensions/personal-access-tokens.mdx | 20 ++- docsgpt/api/admin/routes.py | 13 +- docsgpt/api/pat/routes.py | 23 ++- docsgpt/api/pat/rules.py | 154 +++++++++++++----- docsgpt/api/pat/tokens.py | 3 + docsgpt/api/user/agents/routes.py | 16 +- .../db/repositories/personal_access_tokens.py | 14 +- .../src/settings/PersonalAccessTokens.tsx | 8 + .../src/settings/accessTokenUtils.test.ts | 11 ++ frontend/src/settings/accessTokenUtils.ts | 6 +- tests/api/test_admin_dashboard.py | 13 +- tests/api/test_pat_routes.py | 27 ++- tests/api/test_pat_rules.py | 142 ++++++++++++++++ tests/api/test_pat_tokens.py | 5 + .../test_personal_access_tokens.py | 2 +- 15 files changed, 393 insertions(+), 64 deletions(-) diff --git a/docs/content/Extensions/personal-access-tokens.mdx b/docs/content/Extensions/personal-access-tokens.mdx index 909b9ef6..47d3624d 100644 --- a/docs/content/Extensions/personal-access-tokens.mdx +++ b/docs/content/Extensions/personal-access-tokens.mdx @@ -111,7 +111,7 @@ A `write` scope includes the matching `read` scope. | `teams:read` | View teams, members and resource shares | | `chat:run` | Ask agents and search sources (`/api/answer`, `/stream`, `/api/search`); used for benchmarking | -`agents:keys` is separate from `agents:write` on purpose: a deployment token that updates agents does not need to be able to read or rotate the secrets other systems use to call them. +`agents:keys` is separate from `agents:write` on purpose. Creating, publishing or adopting an agent mints its API key; a token without `agents:keys` gets the key back masked (`1234...90ab`), because an agent key keeps working after the token that saw it is revoked. Put differently: a deployment token that updates agents does not need to be able to read or rotate the secrets other systems use to call them. Some parts of the API are never available to a token, whatever its scopes: token management, the admin API, team management, sign-in flows, device pairing, and the interactive OAuth handshakes used by connectors and MCP servers. A token also never carries the `admin` role, even when its owner is an admin. @@ -135,12 +135,23 @@ For a restricted family the token: - can read, update and delete only the listed resources, and listings show only those; - **cannot create** new resources of that family, since a new resource would be outside the list; - cannot attach a resource outside the list to something else, for example set an agent's source to a source the token may not use; -- is refused (`403`, `"error": "resource_not_allowed"`) wherever DocsGPT cannot prove the request stays inside the list. Schedules, conversations and analytics are closed to agent-restricted tokens for that reason, the message replay endpoints (`/api/messages//tail` and `/api/messages//events`) are closed to every restricted token, and `/api/sources/paginated` is closed to source-restricted tokens (use `/api/sources`). +- is refused (`403`, `"error": "resource_not_allowed"`) wherever DocsGPT cannot prove the request stays inside the list: + - **Conversations, analytics and message replay** (`/api/messages//tail`, `/api/messages//events`) are closed to every restricted token. They span all agents and contain cited source text and tool output. + - **Schedules** run an agent with a free-form instruction and store the output. A token restricted to agents can list and create schedules for its agents; every other schedule route, and schedules altogether for tokens restricted on another family, are closed. + - **Workflow writes** are closed to tokens restricted on sources, tools or prompts, because a workflow graph names those inside its nodes. Such a token also cannot attach a workflow to an agent unless it is restricted on workflows too, in which case only the listed workflows can be attached. + - `/api/sources/paginated` is closed to source-restricted tokens (use `/api/sources`). + +A restriction covers what the token *asks for*, not what an allowed resource already contains: an agent on the list runs with its own sources, prompt and tools even when the token is also restricted on those families. List an agent only if you are happy for the token to use everything that agent uses. Restrictions and chat (`chat:run`): -- A token restricted to specific **agents** must pass `agent_id` in the request body, and it must be one of the listed agents. An agent `api_key` in the body is refused. -- A token restricted to specific **sources** (but not agents) may chat against those sources with `active_docs`. It cannot run agents, because an agent brings its own sources. Restrict the token to agents instead to allow that. +- A token restricted to specific **agents** must pass exactly one `agent_id` in the request body, and it must be one of the listed agents. An agent `api_key` or an inline workflow in the body is refused. +- A token restricted to specific **sources** only may chat against those sources with `active_docs`. It cannot run agents, because an agent brings its own sources. Restrict the token to agents instead to allow that. +- A token restricted on **prompts** or **workflows** must also be restricted to agents to chat. +- A token restricted on **tools** cannot use chat at all, and a tools restriction cannot be combined with `chat:run` when the token is created. Chat executes tools (an agent's own, or your default tools when there is no agent) and those cannot be held to a list. +- A `conversation_id` must belong to the agent being run (or to no agent, for agent-less chat). Otherwise the server would continue, append to, or resume pending tool calls of another agent's conversation. + +`agents:write` and import: applying an agent definition can create the prompt and tools it references and rewrite the agent's workflow, all under `agents:write` alone. It does not need `prompts:write`, `tools:write` or `workflows:write`, so treat `agents:write` as able to create those through an import. Restrictions and `agents apply`: a token restricted to specific agents can apply a definition only when it updates one of those agents. A token restricted on sources, prompts, tools or workflows cannot import agents at all, because an import resolves those references by name and may create them. @@ -151,6 +162,7 @@ Restrictions and `agents apply`: a token restricted to specific agents can apply - Revoking a token in **Settings → Access Tokens** takes effect on the next request. - Admins can list a user's tokens with `GET /api/admin/users//tokens` and revoke any token with `DELETE /api/admin/tokens/`. The admin **revoke sessions** action also revokes all of that user's tokens. - Tokens of a deactivated user (through the admin API or SCIM) stop working immediately and work again if the user is reactivated. +- `GET /api/user/tokens` reports a token past its expiry as `"status": "expired"`. - Token creation and revocation are recorded in the authentication audit log (`pat_created`, `pat_revoked`), visible to admins. An expired token's name can be reused: creating a token with that name retires the expired one. diff --git a/docsgpt/api/admin/routes.py b/docsgpt/api/admin/routes.py index 0bad4772..3d178904 100644 --- a/docsgpt/api/admin/routes.py +++ b/docsgpt/api/admin/routes.py @@ -250,9 +250,18 @@ class AdminUserSessionsResource(Resource): ok = denylist.deny_user(user_id) with db_session() as conn: # A forced logout that left API credentials alive would not be one. - tokens_revoked = PersonalAccessTokensRepository(conn).revoke_all_for_user( + revoked_token_ids = PersonalAccessTokensRepository(conn).revoke_all_for_user( user_id, reason="admin_sessions_revoked" ) + # One pat_revoked event per token, like every other revocation path. + for token_id in revoked_token_ids: + AuthEventsRepository(conn).insert( + user_id, + "pat_revoked", + ip=request.remote_addr, + user_agent=request.headers.get("User-Agent"), + metadata={"token_id": token_id, "by": _actor(), "via": "admin_sessions_revoked"}, + ) AuthEventsRepository(conn).insert( user_id, "admin_sessions_revoked", @@ -262,7 +271,7 @@ class AdminUserSessionsResource(Resource): "by": _actor(), "via": "admin_api", "persisted": ok, - "personal_access_tokens_revoked": tokens_revoked, + "personal_access_tokens_revoked": len(revoked_token_ids), }, ) return make_response(jsonify({"success": True, "revoked": ok}), 200) diff --git a/docsgpt/api/pat/routes.py b/docsgpt/api/pat/routes.py index bb3ea8c5..7c2ecc36 100644 --- a/docsgpt/api/pat/routes.py +++ b/docsgpt/api/pat/routes.py @@ -9,6 +9,7 @@ token can neither mint a replacement nor widen itself. from __future__ import annotations import uuid +from datetime import datetime, timezone from flask import jsonify, make_response, request from flask_restx import Namespace, Resource @@ -50,21 +51,35 @@ def _session_user_id(): def _valid_uuid(value: str) -> bool: + """Canonical form only: ``uuid.UUID`` also accepts ``urn:uuid:…`` and braces, which Postgres does not.""" try: - uuid.UUID(str(value)) + return str(uuid.UUID(str(value))) == str(value).lower() except (ValueError, AttributeError, TypeError): return False - return True + + +def _is_expired(expires_at) -> bool: + if not expires_at: + return False + try: + moment = expires_at if isinstance(expires_at, datetime) else datetime.fromisoformat(str(expires_at)) + except ValueError: + return False + if moment.tzinfo is None: + moment = moment.replace(tzinfo=timezone.utc) + return moment <= datetime.now(timezone.utc) def serialize_token(row: dict) -> dict: + # The row keeps status 'active' until someone revokes it; report what is true for a caller. + status = "expired" if row["status"] == "active" and _is_expired(row.get("expires_at")) else row["status"] return { "id": str(row["id"]), "name": row["name"], "token_prefix": row["token_prefix"], "scopes": list(row.get("scopes") or []), "resource_filter": row.get("resource_filter") or {}, - "status": row["status"], + "status": status, "expires_at": row.get("expires_at"), "last_used_at": row.get("last_used_at"), "last_used_ip": row.get("last_used_ip"), @@ -135,6 +150,8 @@ class PersonalAccessTokens(Resource): try: with db_session() as conn: repo = PersonalAccessTokensRepository(conn) + # Serialise this user's creates so concurrent requests cannot both pass the cap check. + repo.lock_user(user_id) if repo.count_active(user_id) >= settings.PAT_MAX_PER_USER: return _error( f"Token limit reached ({settings.PAT_MAX_PER_USER}); revoke one first", 409 diff --git a/docsgpt/api/pat/rules.py b/docsgpt/api/pat/rules.py index 88f1df76..784f764c 100644 --- a/docsgpt/api/pat/rules.py +++ b/docsgpt/api/pat/rules.py @@ -44,7 +44,7 @@ class Rule: open: bool = False in_route: bool = False blocked_by: tuple[str, ...] = () - check: Optional[Callable[[Any, dict], Optional[str]]] = None + check: Optional[Callable[[Any, dict, Optional[str]], Optional[str]]] = None def _rule(scope: Optional[str] = None, *ids: Locator, any_of: tuple[str, ...] = (), **kwargs) -> Rule: @@ -61,6 +61,9 @@ def _rule(scope: Optional[str] = None, *ids: Locator, any_of: tuple[str, ...] = MESSAGE_REPLAY_SCOPES = ("conversations:read", "chat:run") _ALL_FAMILIES = ("agents", "sources", "prompts", "tools", "workflows") +_WORKFLOW_CONTENT_FAMILIES = ("sources", "tools", "prompts") +# A route reached through an agent id can prove the agent; nothing else about it. +_NON_AGENT_FAMILIES = ("sources", "prompts", "tools", "workflows") # Ids of other families that agent create/update accept in their JSON-or-form body. _AGENT_BODY_REFS = ( @@ -72,15 +75,41 @@ _AGENT_BODY_REFS = ( ) -def _chat_check(request, resource_filter: dict) -> Optional[str]: +def _conversation_agent_id(conversation_id: str, user_id: Optional[str]) -> tuple[bool, str]: + """``(found, agent_id)`` for a conversation the user can reach; ``agent_id`` is "" when it has none.""" + from docsgpt.storage.db.repositories.conversations import ConversationsRepository + from docsgpt.storage.db.session import db_readonly + + if not user_id: + return False, "" + try: + with db_readonly() as conn: + row = ConversationsRepository(conn).get_any(str(conversation_id), user_id) + except Exception: + return False, "" + if not row: + return False, "" + return True, str(row.get("agent_id") or "") + + +def _chat_check(request, resource_filter: dict, user_id: Optional[str]) -> Optional[str]: """Keep a restricted token's chat traffic inside its allowlists. An agent brings its own sources, prompt and tools, which this table cannot - see, so a token restricted on any family must name an allowed agent (or, - when only sources are restricted, chat against allowed sources directly). - An agent ``api_key`` in the body would swap in an arbitrary agent. + see, so a restricted token must name an allowed agent. The one exception is + a token restricted on sources only, which may chat against allowed sources + directly. Everything that could swap in another agent or another set of + resources is refused: an agent ``api_key``, an inline workflow, and a + ``conversation_id`` that belongs to a different agent (the server would + otherwise continue, append to, or resume tool calls of that conversation). + + Chat executes tools: an agent's own, or the user's defaults when there is + no agent. Neither can be held to a tools allowlist from here, so a token + restricted on tools cannot chat at all. """ body = _json_body(request) + if "tools" in resource_filter: + return "A token restricted to specific tools cannot use chat endpoints" if body.get("api_key"): return "A restricted token cannot chat with an agent API key; pass agent_id" if body.get("workflow"): @@ -88,11 +117,33 @@ def _chat_check(request, resource_filter: dict) -> Optional[str]: return "A restricted token cannot run an inline workflow" agent_ids = _as_ids(body.get("agent_id")) if "agents" in resource_filter: - if not agent_ids: + if len(agent_ids) != 1: return "This token is restricted to specific agents; pass agent_id" - return None # the agent id itself is verified through ``refs`` - if agent_ids: - return "This token is restricted to specific resources and cannot run arbitrary agents" + # the agent id itself is verified through ``refs`` + elif set(resource_filter) - {"sources"}: + return "Restrict this token to specific agents to use chat endpoints" + elif agent_ids: + return "This token is restricted to specific sources and cannot run agents" + conversation_id = body.get("conversation_id") + if conversation_id: + found, conversation_agent = _conversation_agent_id(conversation_id, user_id) + expected = agent_ids[0] if agent_ids else "" + if not found or _canonical(conversation_agent) != _canonical(expected): + return "This conversation does not belong to the agent this token may use" + return None + + +def _agent_body_check(request, resource_filter: dict, user_id: Optional[str]) -> Optional[str]: + """A workflow pulls in its own sources, tools and prompts, which a reference check cannot see. + + A token restricted on any of those may attach a workflow to an agent only + when it is also restricted on workflows, so the workflow is one its owner + chose (``refs`` then verifies the id). + """ + if "workflows" in resource_filter or not set(resource_filter) & {"sources", "tools", "prompts"}: + return None + if _read(request, (BODY, "workflow")): + return "A token restricted to specific sources, tools or prompts cannot attach a workflow to an agent" return None @@ -124,9 +175,9 @@ RULES: dict[tuple[str, str], Rule] = { ("/api/guardrails/summary", "GET"): _rule("agents:read", (QUERY, "agent_id")), ("/api/agents/folders/", "GET"): _rule("agents:read", open=True), ("/api/agents/folders/", "GET"): _rule("agents:read"), - ("/api/create_agent", "POST"): _rule("agents:write", refs=_AGENT_BODY_REFS), + ("/api/create_agent", "POST"): _rule("agents:write", refs=_AGENT_BODY_REFS, check=_agent_body_check), ("/api/update_agent/", "PUT"): _rule( - "agents:write", (VIEW, "agent_id"), refs=_AGENT_BODY_REFS + "agents:write", (VIEW, "agent_id"), refs=_AGENT_BODY_REFS, check=_agent_body_check ), ("/api/delete_agent", "DELETE"): _rule("agents:write", (QUERY, "id")), ("/api/adopt_agent", "POST"): _rule("agents:write"), @@ -142,22 +193,25 @@ RULES: dict[tuple[str, str], Rule] = { ("/api/agents/folders/bulk_move", "POST"): _rule("agents:write", (JSON, "agent_ids")), ("/api/regenerate_agent_key/", "POST"): _rule("agents:keys", (VIEW, "agent_id")), ("/api/agent_webhook", "GET"): _rule("agents:keys", (QUERY, "id")), - # Schedules hang off an agent. + # Schedules hang off an agent, and a schedule runs that agent with a free-form + # instruction and stores the output. The agent id proves the agent and nothing + # else, so tokens restricted on any other family are kept out; routes that + # carry only a schedule id prove nothing and are closed to every restricted token. ("/api/agents//schedules", "GET"): _rule( - "schedules:read", refs=(("agents", (VIEW, "agent_id")),) + "schedules:read", refs=(("agents", (VIEW, "agent_id")),), blocked_by=_NON_AGENT_FAMILIES ), ("/api/agents//schedules", "POST"): _rule( - "schedules:write", refs=(("agents", (VIEW, "agent_id")),), blocked_by=("tools",) + "schedules:write", refs=(("agents", (VIEW, "agent_id")),), blocked_by=_NON_AGENT_FAMILIES ), - ("/api/schedules/", "GET"): _rule("schedules:read", blocked_by=("agents",)), - ("/api/schedules//runs", "GET"): _rule("schedules:read", blocked_by=("agents",)), + ("/api/schedules/", "GET"): _rule("schedules:read", blocked_by=_ALL_FAMILIES), + ("/api/schedules//runs", "GET"): _rule("schedules:read", blocked_by=_ALL_FAMILIES), ("/api/schedules//runs/", "GET"): _rule( - "schedules:read", blocked_by=("agents",) + "schedules:read", blocked_by=_ALL_FAMILIES ), - ("/api/schedules/", "PUT"): _rule("schedules:write", blocked_by=("agents", "tools")), - ("/api/schedules/", "PATCH"): _rule("schedules:write", blocked_by=("agents", "tools")), - ("/api/schedules/", "DELETE"): _rule("schedules:write", blocked_by=("agents",)), - ("/api/schedules//run", "POST"): _rule("schedules:write", blocked_by=("agents",)), + ("/api/schedules/", "PUT"): _rule("schedules:write", blocked_by=_ALL_FAMILIES), + ("/api/schedules/", "PATCH"): _rule("schedules:write", blocked_by=_ALL_FAMILIES), + ("/api/schedules/", "DELETE"): _rule("schedules:write", blocked_by=_ALL_FAMILIES), + ("/api/schedules//run", "POST"): _rule("schedules:write", blocked_by=_ALL_FAMILIES), # Sources ("/api/sources", "GET"): _rule("sources:read", listing=True), # Counted and paged in SQL, so it cannot be narrowed here; restricted tokens use /api/sources. @@ -217,28 +271,32 @@ RULES: dict[tuple[str, str], Rule] = { ("/api/user/models/test", "POST"): _rule("models:write"), ("/api/user/models//test", "POST"): _rule("models:write"), # Workflows - ("/api/workflows", "POST"): _rule("workflows:write"), + # A workflow graph names sources, tools and prompts inside its nodes, out of reach of ``refs``. + ("/api/workflows", "POST"): _rule("workflows:write", blocked_by=_WORKFLOW_CONTENT_FAMILIES), ("/api/workflows/", "GET"): _rule("workflows:read", (VIEW, "workflow_id")), - ("/api/workflows/", "PUT"): _rule("workflows:write", (VIEW, "workflow_id")), + ("/api/workflows/", "PUT"): _rule( + "workflows:write", (VIEW, "workflow_id"), blocked_by=_WORKFLOW_CONTENT_FAMILIES + ), ("/api/workflows/", "DELETE"): _rule("workflows:write", (VIEW, "workflow_id")), - # Conversations and analytics span every agent, so an agent-restricted token is kept out. - ("/api/get_conversations", "GET"): _rule("conversations:read", blocked_by=("agents",)), - ("/api/search_conversations", "GET"): _rule("conversations:read", blocked_by=("agents",)), - ("/api/get_single_conversation", "GET"): _rule("conversations:read", blocked_by=("agents",)), + # Conversations and analytics span every agent and carry cited source text and tool + # output, so they are closed to every restricted token. + ("/api/get_conversations", "GET"): _rule("conversations:read", blocked_by=_ALL_FAMILIES), + ("/api/search_conversations", "GET"): _rule("conversations:read", blocked_by=_ALL_FAMILIES), + ("/api/get_single_conversation", "GET"): _rule("conversations:read", blocked_by=_ALL_FAMILIES), # A message cannot be tied to an allowlist from here, so any restricted token is kept out. ("/api/messages//tail", "GET"): _rule( any_of=MESSAGE_REPLAY_SCOPES, family=None, blocked_by=_ALL_FAMILIES ), - ("/api/delete_conversation", "POST"): _rule("conversations:write", blocked_by=("agents",)), - ("/api/delete_all_conversations", "GET"): _rule("conversations:write", blocked_by=("agents",)), - ("/api/update_conversation_name", "POST"): _rule("conversations:write", blocked_by=("agents",)), - ("/api/feedback", "POST"): _rule("conversations:write", blocked_by=("agents",)), - ("/api/get_message_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), - ("/api/get_token_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), - ("/api/get_feedback_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), - ("/api/get_tool_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), - ("/api/get_schedule_analytics", "POST"): _rule("analytics:read", blocked_by=("agents",)), - ("/api/get_user_logs", "POST"): _rule("analytics:read", blocked_by=("agents",)), + ("/api/delete_conversation", "POST"): _rule("conversations:write", blocked_by=_ALL_FAMILIES), + ("/api/delete_all_conversations", "GET"): _rule("conversations:write", blocked_by=_ALL_FAMILIES), + ("/api/update_conversation_name", "POST"): _rule("conversations:write", blocked_by=_ALL_FAMILIES), + ("/api/feedback", "POST"): _rule("conversations:write", blocked_by=_ALL_FAMILIES), + ("/api/get_message_analytics", "POST"): _rule("analytics:read", blocked_by=_ALL_FAMILIES), + ("/api/get_token_analytics", "POST"): _rule("analytics:read", blocked_by=_ALL_FAMILIES), + ("/api/get_feedback_analytics", "POST"): _rule("analytics:read", blocked_by=_ALL_FAMILIES), + ("/api/get_tool_analytics", "POST"): _rule("analytics:read", blocked_by=_ALL_FAMILIES), + ("/api/get_schedule_analytics", "POST"): _rule("analytics:read", blocked_by=_ALL_FAMILIES), + ("/api/get_user_logs", "POST"): _rule("analytics:read", blocked_by=_ALL_FAMILIES), # Teams (read only) ("/api/teams", "GET"): _rule("teams:read"), ("/api/teams/", "GET"): _rule("teams:read"), @@ -395,18 +453,18 @@ def authorize(request, decoded_token: dict) -> Optional[tuple[dict, int]]: resource_filter = decoded_token.get("resource_filter") or {} if not resource_filter: return None - reason = _check_resources(request, rule, resource_filter) + reason = _check_resources(request, rule, resource_filter, decoded_token.get("sub")) if reason is None: return None return ({"success": False, "error": "resource_not_allowed", "message": reason}, 403) -def _check_resources(request, rule: Rule, resource_filter: dict) -> Optional[str]: +def _check_resources(request, rule: Rule, resource_filter: dict, user_id: Optional[str] = None) -> Optional[str]: for family in rule.blocked_by: if family in resource_filter: return f"This endpoint is not available to a token restricted to specific {family}" if rule.check is not None: - reason = rule.check(request, resource_filter) + reason = rule.check(request, resource_filter, user_id) if reason: return reason for family, locator in rule.refs: @@ -459,3 +517,19 @@ def _is_uuid(value: str) -> bool: except (ValueError, AttributeError, TypeError): return False return True + + +def may_see_agent_keys(request) -> bool: + """False for a token without ``agents:keys``: it must not receive a plaintext agent API key. + + Create, first publish and adopt all mint a key and used to return it, which + handed a deploy token a secret that outlives the token's own revocation. + """ + decoded = getattr(request, "decoded_token", None) + if not is_pat(decoded): + return True + return "agents:keys" in (decoded.get("scopes") or []) + + +def mask_agent_key(key: Optional[str]) -> str: + return f"{key[:4]}...{key[-4:]}" if key else "" diff --git a/docsgpt/api/pat/tokens.py b/docsgpt/api/pat/tokens.py index fa4cb9d0..6e4978bd 100644 --- a/docsgpt/api/pat/tokens.py +++ b/docsgpt/api/pat/tokens.py @@ -121,6 +121,9 @@ def normalize_resource_filter(raw: Any, scopes: list[str]) -> dict[str, list[str # chat:run acts on agents and sources, so both may be restricted alongside it. if "chat" in families: families.update({"agents", "sources"}) + if "tools" in raw and "chat:run" in scopes: + # Chat executes tools (an agent's own, or the user's defaults), which cannot be held to an allowlist. + raise ValueError("resource_filter.tools cannot be combined with the chat:run scope") out: dict[str, list[str]] = {} for family, ids in raw.items(): if family not in FILTERABLE_FAMILIES: diff --git a/docsgpt/api/user/agents/routes.py b/docsgpt/api/user/agents/routes.py index 757dd933..6a0ec2b0 100644 --- a/docsgpt/api/user/agents/routes.py +++ b/docsgpt/api/user/agents/routes.py @@ -9,7 +9,7 @@ from flask_restx import fields, Namespace, Resource from pydantic import ValidationError as PydanticValidationError from docsgpt.api import api -from docsgpt.api.pat.rules import filter_listing +from docsgpt.api.pat.rules import filter_listing, mask_agent_key, may_see_agent_keys from docsgpt.guardrails.config import AgentConfig from docsgpt.api.user.base import ( copy_agent_image_for_user, @@ -763,7 +763,9 @@ class CreateAgent(Resource): except Exception as err: current_app.logger.error(f"Error creating agent: {err}", exc_info=True) return make_response(jsonify({"success": False}), 400) - return make_response(jsonify({"id": new_id, "key": key}), 201) + # A token without agents:keys never receives the plaintext agent key. + visible_key = key if may_see_agent_keys(request) else mask_agent_key(key) + return make_response(jsonify({"id": new_id, "key": visible_key}), 201) @agents_ns.route("/update_agent/") @@ -1307,7 +1309,11 @@ class UpdateAgent(Resource): "message": "Agent updated successfully", } if newly_generated_key: - response_data["key"] = newly_generated_key + response_data["key"] = ( + newly_generated_key + if may_see_agent_keys(request) + else mask_agent_key(newly_generated_key) + ) return make_response(jsonify(response_data), 200) @@ -1811,7 +1817,9 @@ class AdoptAgent(Resource): ) response_agent = _format_agent_output(new_agent, include_key_masked=False) - response_agent["key"] = new_key + response_agent["key"] = ( + new_key if may_see_agent_keys(request) else mask_agent_key(new_key) + ) return make_response( jsonify({"success": True, "agent": response_agent}), 200 ) diff --git a/docsgpt/storage/db/repositories/personal_access_tokens.py b/docsgpt/storage/db/repositories/personal_access_tokens.py index efbe250b..12bd80cc 100644 --- a/docsgpt/storage/db/repositories/personal_access_tokens.py +++ b/docsgpt/storage/db/repositories/personal_access_tokens.py @@ -76,6 +76,13 @@ class PersonalAccessTokensRepository: result = self._conn.execute(text(sql), {"user_id": user_id}) return [row_to_dict(r) for r in result.fetchall()] + def lock_user(self, user_id: str) -> None: + """Hold a per-user advisory lock until the transaction ends.""" + self._conn.execute( + text("SELECT pg_advisory_xact_lock(hashtextextended(:key, 0))"), + {"key": f"personal_access_tokens:{user_id}"}, + ) + def count_active(self, user_id: str) -> int: """Live tokens only: revoked and expired rows don't count against the per-user cap.""" return self._conn.execute( @@ -161,13 +168,14 @@ class PersonalAccessTokensRepository: params["user_id"] = user_id return self._conn.execute(text(sql), params).rowcount > 0 - def revoke_all_for_user(self, user_id: str, *, reason: str = "admin_revoked") -> int: + def revoke_all_for_user(self, user_id: str, *, reason: str = "admin_revoked") -> list[str]: + """Revoke every live token of a user. Returns the revoked token ids (for the audit trail).""" result = self._conn.execute( text( "UPDATE personal_access_tokens " "SET status = 'revoked', revoked_at = now(), revoke_reason = :reason " - "WHERE user_id = :user_id AND status = 'active'" + "WHERE user_id = :user_id AND status = 'active' RETURNING id" ), {"user_id": user_id, "reason": reason}, ) - return result.rowcount + return [str(row[0]) for row in result.fetchall()] diff --git a/frontend/src/settings/PersonalAccessTokens.tsx b/frontend/src/settings/PersonalAccessTokens.tsx index 7439e567..f417200c 100644 --- a/frontend/src/settings/PersonalAccessTokens.tsx +++ b/frontend/src/settings/PersonalAccessTokens.tsx @@ -4,6 +4,7 @@ import { useTranslation } from 'react-i18next'; import { useSelector } from 'react-redux'; import patService, { + AccessTokenApiError, AccessTokenPolicy, AccessTokenScope, CreateAccessTokenResponse, @@ -134,6 +135,13 @@ export default function PersonalAccessTokens() { setTokens((prev) => prev.filter((item) => item.id !== target.id)); setError(null); } catch (err) { + if (err instanceof AccessTokenApiError && err.status === 404) { + // Already revoked elsewhere (another tab, an admin): it is gone, so + // drop the stale row instead of reporting a failure. + setTokens((prev) => prev.filter((item) => item.id !== target.id)); + setError(null); + return; + } console.error('Failed to revoke access token:', err); setError( (err instanceof Error && err.message) || diff --git a/frontend/src/settings/accessTokenUtils.test.ts b/frontend/src/settings/accessTokenUtils.test.ts index c820fc8d..a9df9ed5 100644 --- a/frontend/src/settings/accessTokenUtils.test.ts +++ b/frontend/src/settings/accessTokenUtils.test.ts @@ -294,3 +294,14 @@ describe('NO_ESCAPE', () => { expect(NO_ESCAPE).toEqual({ interpolation: { escapeValue: false } }); }); }); + +describe('eligibleFilterFamilies with chat:run', () => { + const all = ['agents', 'sources', 'prompts', 'tools', 'workflows']; + it('never offers a tools restriction next to chat:run', () => { + expect(eligibleFilterFamilies(['tools:write', 'chat:run'], all)).toEqual([ + 'agents', + 'sources', + ]); + expect(eligibleFilterFamilies(['tools:write'], all)).toEqual(['tools']); + }); +}); diff --git a/frontend/src/settings/accessTokenUtils.ts b/frontend/src/settings/accessTokenUtils.ts index 9e7e56c3..ddfc4094 100644 --- a/frontend/src/settings/accessTokenUtils.ts +++ b/frontend/src/settings/accessTokenUtils.ts @@ -75,7 +75,11 @@ export function eligibleFilterFamilies( filterableFamilies: string[], ): string[] { const families = new Set(selected.map(scopeFamily)); - if (families.has('chat')) CHAT_FAMILIES.forEach((f) => families.add(f)); + const chat = families.has('chat'); + if (chat) CHAT_FAMILIES.forEach((f) => families.add(f)); + // Chat executes tools, which cannot be held to an allowlist, so the server + // rejects a tools restriction on a token that also has chat:run. + if (chat) families.delete('tools'); return filterableFamilies.filter((family) => families.has(family)); } diff --git a/tests/api/test_admin_dashboard.py b/tests/api/test_admin_dashboard.py index 50f7dfe0..59c67e4c 100644 --- a/tests/api/test_admin_dashboard.py +++ b/tests/api/test_admin_dashboard.py @@ -178,13 +178,20 @@ class TestUserLifecycle: def test_force_logout(self, client): events = Mock() - with _admin(AuthEventsRepository=Mock(return_value=events)), patch( - "docsgpt.api.admin.routes.denylist" - ) as dl: + tokens = Mock() + tokens.revoke_all_for_user.return_value = ["t1", "t2"] + with _admin( + AuthEventsRepository=Mock(return_value=events), + PersonalAccessTokensRepository=Mock(return_value=tokens), + ), patch("docsgpt.api.admin.routes.denylist") as dl: dl.deny_user.return_value = True resp = client.post("/api/admin/users/bob/revoke-sessions") assert resp.status_code == 200 dl.deny_user.assert_called_once_with("bob") + # A forced logout also revokes the user's API tokens, one audit event each. + tokens.revoke_all_for_user.assert_called_once_with("bob", reason="admin_sessions_revoked") + recorded = [call.args[1] for call in events.insert.call_args_list] + assert recorded == ["pat_revoked", "pat_revoked", "admin_sessions_revoked"] def test_user_detail(self, client): users = Mock() diff --git a/tests/api/test_pat_routes.py b/tests/api/test_pat_routes.py index f7592cf4..cf21c8c5 100644 --- a/tests/api/test_pat_routes.py +++ b/tests/api/test_pat_routes.py @@ -166,6 +166,16 @@ class TestCreate: assert "dgpt_pat_" not in json.dumps(metadata) +class TestExpiredStatus: + def test_expired_token_is_reported_as_expired_not_active(self, client, db): + _create(client) + with _session(): + assert json.loads(client.get("/api/user/tokens").data)["tokens"][0]["status"] == "active" + db.execute(text("UPDATE personal_access_tokens SET expires_at = now() - interval '1 day'")) + with _session(): + assert json.loads(client.get("/api/user/tokens").data)["tokens"][0]["status"] == "expired" + + class TestList: def test_includes_scope_catalog_and_policy(self, client, db): with _session(): @@ -219,10 +229,16 @@ class TestRevoke: with _session(sub="bob"): assert client.delete(f"/api/user/tokens/{token_id}").status_code == 404 - def test_unknown_and_malformed_ids(self, client, db): + @pytest.mark.parametrize( + "token_id", + [AGENT_A, "not-a-uuid", f"urn:uuid:{AGENT_A}", "{" + AGENT_A + "}", AGENT_A.replace("-", "")], + ) + def test_unknown_and_malformed_ids(self, client, db, token_id): + # uuid.UUID() accepts urn:/braced/bare-hex spellings that Postgres rejects; none may reach the cast. with _session(): - assert client.delete(f"/api/user/tokens/{AGENT_A}").status_code == 404 - assert client.delete("/api/user/tokens/not-a-uuid").status_code == 404 + assert client.delete(f"/api/user/tokens/{token_id}").status_code == 404 + with _session(sub="root", roles=("admin", "user")): + assert client.delete(f"/api/admin/tokens/{token_id}").status_code == 404 class TestAdmin: @@ -256,3 +272,8 @@ class TestAdmin: assert client.post("/api/admin/users/alice/revoke-sessions").status_code == 200 headers = {"Authorization": f"Bearer {created['token']}"} assert client.get("/api/user/me", headers=headers).status_code == 401 + events = db.execute( + text("SELECT metadata FROM auth_events WHERE user_id = 'alice' AND event = 'pat_revoked'") + ).all() + assert [e[0]["token_id"] for e in events] == [created["personal_access_token"]["id"]] + assert events[0][0]["via"] == "admin_sessions_revoked" diff --git a/tests/api/test_pat_rules.py b/tests/api/test_pat_rules.py index 4bdc7070..69e532a3 100644 --- a/tests/api/test_pat_rules.py +++ b/tests/api/test_pat_rules.py @@ -246,6 +246,102 @@ class TestResourceRestrictions: assert _denied(_call(client, "GET", "/api/sources/paginated", claims)) == "resource_not_allowed" +WORKFLOW_A = "eeeeeeee-eeee-eeee-eeee-eeeeeeeeeeee" +TOOL_A = "ffffffff-ffff-ffff-ffff-ffffffffffff" + + +@pytest.mark.unit +class TestRelationshipsBeyondIds: + """Rows whose content or parent the table cannot see are closed to restricted tokens.""" + + @pytest.mark.parametrize("family", ["sources", "tools", "prompts"]) + def test_workflow_writes_are_closed_to_tokens_restricted_on_what_a_graph_can_name(self, client, family): + claims = _claims(["workflows:write"], {family: [SOURCE_A]}) + assert _denied(_call(client, "POST", "/api/workflows", claims, json={})) == "resource_not_allowed" + assert ( + _denied(_call(client, "PUT", f"/api/workflows/{WORKFLOW_A}", claims, json={})) + == "resource_not_allowed" + ) + assert _denied(_call(client, "GET", f"/api/workflows/{WORKFLOW_A}", claims)) is None + + def test_workflow_restricted_token_can_still_edit_its_workflows(self, client): + claims = _claims(["workflows:write"], {"workflows": [WORKFLOW_A]}) + assert _denied(_call(client, "PUT", f"/api/workflows/{WORKFLOW_A}", claims, json={})) is None + + @pytest.mark.parametrize("family", ["sources", "tools", "prompts"]) + def test_agent_cannot_be_pointed_at_a_workflow_by_a_token_restricted_on_its_contents(self, client, family): + claims = _claims(["agents:write"], {family: [SOURCE_A]}) + for body in ({"workflow": WORKFLOW_A}, {"workflow": {"id": WORKFLOW_A}}): + response = _call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, json=body) + assert _denied(response) == "resource_not_allowed" + form = _call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, data={"workflow": WORKFLOW_A}) + assert _denied(form) == "resource_not_allowed" + assert _denied(_call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, json={"name": "x"})) is None + + def test_workflow_allowlist_lets_the_agent_use_those_workflows_only(self, client): + claims = _claims(["agents:write"], {"sources": [SOURCE_A], "workflows": [WORKFLOW_A]}) + ok = _call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, json={"workflow": WORKFLOW_A}) + bad = _call(client, "PUT", f"/api/update_agent/{AGENT_A}", claims, json={"workflow": SOURCE_B}) + assert _denied(ok) is None + assert _denied(bad) == "resource_not_allowed" + + @pytest.mark.parametrize("family", ["sources", "prompts", "tools", "workflows"]) + def test_schedules_are_closed_to_tokens_restricted_on_anything_but_agents(self, client, family): + claims = _claims(["schedules:write"], {family: [SOURCE_A]}) + for method, path in ( + ("GET", f"/api/agents/{AGENT_A}/schedules"), + ("POST", f"/api/agents/{AGENT_A}/schedules"), + ("GET", "/api/schedules/s1"), + ("POST", "/api/schedules/s1/run"), + ("GET", "/api/schedules/s1/runs"), + ): + assert _denied(_call(client, method, path, claims, json={})) == "resource_not_allowed", path + + @pytest.mark.parametrize("family", ["agents", "sources", "prompts", "tools", "workflows"]) + def test_conversations_and_analytics_are_closed_to_every_restricted_token(self, client, family): + claims = _claims(["conversations:write", "analytics:read"], {family: [SOURCE_A]}) + for method, path in ( + ("GET", "/api/get_conversations"), + ("GET", "/api/get_single_conversation?id=c1"), + ("GET", "/api/search_conversations?q=x"), + ("POST", "/api/delete_conversation"), + ("POST", "/api/feedback"), + ("POST", "/api/get_message_analytics"), + ("POST", "/api/get_user_logs"), + ): + assert _denied(_call(client, method, path, claims, json={})) == "resource_not_allowed", path + + +@pytest.mark.unit +class TestAgentKeyVisibility: + def _request(self, flask_app, claims): + from flask import request + + ctx = flask_app.test_request_context("/") + ctx.push() + request.decoded_token = claims + return ctx, request + + def test_sessions_and_tokens_with_the_keys_scope_see_the_key(self, flask_app): + for claims in ({"sub": "alice"}, _claims(["agents:write", "agents:keys"])): + ctx, request = self._request(flask_app, claims) + try: + assert rules.may_see_agent_keys(request) is True + finally: + ctx.pop() + + def test_token_without_the_keys_scope_does_not(self, flask_app): + ctx, request = self._request(flask_app, _claims(["agents:write"])) + try: + assert rules.may_see_agent_keys(request) is False + finally: + ctx.pop() + + def test_mask(self): + assert rules.mask_agent_key("12345678-aaaa-bbbb-cccc-1234567890ab") == "1234...90ab" + assert rules.mask_agent_key("") == "" and rules.mask_agent_key(None) == "" + + @pytest.mark.unit class TestChatRestrictions: def _chat(self, client, claims, body): @@ -274,6 +370,52 @@ class TestChatRestrictions: assert self._chat(client, claims, {"question": "hi", "active_docs": [SOURCE_A, SOURCE_B]}) == "resource_not_allowed" assert self._chat(client, claims, {"question": "hi", "agent_id": AGENT_A}) == "resource_not_allowed" + @pytest.mark.parametrize("extra", [{}, {"agents": [AGENT_A]}]) + def test_tools_restricted_token_cannot_chat_at_all(self, client, extra): + claims = _claims(["chat:run"], {"tools": [TOOL_A], **extra}) + assert self._chat(client, claims, {"question": "hi"}) == "resource_not_allowed" + assert self._chat(client, claims, {"question": "hi", "agent_id": AGENT_A}) == "resource_not_allowed" + + @pytest.mark.parametrize("family", ["prompts", "workflows"]) + def test_agentless_chat_needs_an_agent_restriction_unless_only_sources_are_restricted(self, client, family): + claims = _claims(["chat:run"], {family: [SOURCE_A]}) + assert self._chat(client, claims, {"question": "hi"}) == "resource_not_allowed" + assert self._chat(client, claims, {"question": "hi", "agent_id": AGENT_A}) == "resource_not_allowed" + + def test_conversation_must_belong_to_the_agent_being_run(self, client): + claims = _claims(["chat:run"], {"agents": [AGENT_A]}) + body = {"question": "hi", "agent_id": AGENT_A, "conversation_id": "c1"} + with patch.object(rules, "_conversation_agent_id", return_value=(True, AGENT_A.upper())) as lookup: + assert self._chat(client, claims, body) is None + lookup.assert_called_once_with("c1", "alice") + for result in ((True, AGENT_B), (True, ""), (False, "")): + with patch.object(rules, "_conversation_agent_id", return_value=result): + assert self._chat(client, claims, body) == "resource_not_allowed", result + + def test_conversation_resume_with_tool_actions_is_held_to_the_same_rule(self, client): + claims = _claims(["chat:run"], {"agents": [AGENT_A]}) + body = {"agent_id": AGENT_A, "conversation_id": "c1", "tool_actions": [{"call_id": "x"}]} + with patch.object(rules, "_conversation_agent_id", return_value=(True, AGENT_B)): + assert self._chat(client, claims, body) == "resource_not_allowed" + + def test_agentless_token_cannot_continue_an_agent_conversation(self, client): + claims = _claims(["chat:run"], {"sources": [SOURCE_A]}) + body = {"question": "hi", "active_docs": SOURCE_A, "conversation_id": "c1"} + with patch.object(rules, "_conversation_agent_id", return_value=(True, AGENT_B)): + assert self._chat(client, claims, body) == "resource_not_allowed" + with patch.object(rules, "_conversation_agent_id", return_value=(True, "")): + assert self._chat(client, claims, body) is None + + def test_unrestricted_token_never_pays_for_the_conversation_lookup(self, client): + with patch.object(rules, "_conversation_agent_id") as lookup: + self._chat(client, _claims(["chat:run"]), {"question": "hi", "conversation_id": "c1"}) + lookup.assert_not_called() + + def test_conversation_lookup_fails_closed(self): + with patch("docsgpt.storage.db.session.db_readonly", side_effect=RuntimeError("db down")): + assert rules._conversation_agent_id("c1", "alice") == (False, "") + assert rules._conversation_agent_id("c1", None) == (False, "") + def test_retrieval_test_honours_the_source_allowlist(self, client): claims = _claims(["chat:run"], {"sources": [SOURCE_A]}) ok = _call(client, "POST", f"/api/sources/{SOURCE_A}/search", claims, json={"query": "q"}) diff --git a/tests/api/test_pat_tokens.py b/tests/api/test_pat_tokens.py index 1e38d0ce..24f6a5bb 100644 --- a/tests/api/test_pat_tokens.py +++ b/tests/api/test_pat_tokens.py @@ -113,6 +113,11 @@ class TestResourceFilter: ) assert set(out) == {"agents", "sources"} + def test_tools_restriction_cannot_be_combined_with_chat(self): + with pytest.raises(ValueError, match="chat:run"): + tokens.normalize_resource_filter({"tools": [self.UUID_A]}, ["tools:read", "chat:run"]) + assert tokens.normalize_resource_filter({"tools": [self.UUID_A]}, ["tools:read"]) + @pytest.mark.parametrize( "raw,scopes", [ diff --git a/tests/storage/db/repositories/test_personal_access_tokens.py b/tests/storage/db/repositories/test_personal_access_tokens.py index 8a6848a3..0f4732f0 100644 --- a/tests/storage/db/repositories/test_personal_access_tokens.py +++ b/tests/storage/db/repositories/test_personal_access_tokens.py @@ -150,7 +150,7 @@ class TestRevoke: _create(repo, name="a", token_hash="h1") _create(repo, name="b", token_hash="h2") _create(repo, user_id="u2", token_hash="h3") - assert repo.revoke_all_for_user("u1") == 2 + assert len(repo.revoke_all_for_user("u1")) == 2 assert repo.list_for_user("u1") == [] assert len(repo.list_for_user("u2")) == 1 From 3e38dc19eca562b8b3e27005e965bd0eebe533b2 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Sun, 20 Sep 2026 19:47:26 +0100 Subject: [PATCH 113/130] docs: key CI source uploads by commit so a revert is ingested again --- docs/content/Extensions/personal-access-tokens.mdx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/content/Extensions/personal-access-tokens.mdx b/docs/content/Extensions/personal-access-tokens.mdx index 47d3624d..7e5be596 100644 --- a/docs/content/Extensions/personal-access-tokens.mdx +++ b/docs/content/Extensions/personal-access-tokens.mdx @@ -86,7 +86,7 @@ jobs: DOCSGPT_URL: ${{ vars.DOCSGPT_URL }} DOCSGPT_TOKEN: ${{ secrets.DOCSGPT_TOKEN }} run: | - docsgpt-cli sources upload docs/*.md --name "Product docs" --wait --replace + docsgpt-cli sources upload docs/*.md --name "Product docs" --wait --replace --idempotency-key "docs-${{ github.sha }}" docsgpt-cli agents apply -f agents/ ``` From 91f2ec2f092a583655e4440b1b6685fcb94f99dd Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Mon, 21 Sep 2026 10:32:17 +0100 Subject: [PATCH 114/130] feat(pat): regenerate a token's secret and reset its expiry POST /api/user/tokens//regenerate swaps the secret of an existing token in place: name, scopes and restrictions stay, the old secret stops matching at once, and the expiry is reset. The lifetime defaults to the one the token was last issued with (clamped to today's policy) or to expires_in_days when given. An expired token can be renewed this way; a revoked one cannot. It is session only like the rest of token management, and writes a pat_regenerated audit event. regenerated_at records the rotation. --- .../Extensions/personal-access-tokens.mdx | 6 +- .../versions/0032_personal_access_tokens.py | 4 +- docsgpt/api/pat/routes.py | 62 ++++++++++ docsgpt/api/pat/rules.py | 1 + docsgpt/api/pat/tokens.py | 27 +++++ docsgpt/storage/db/models.py | 1 + .../db/repositories/personal_access_tokens.py | 37 +++++- tests/api/test_pat_routes.py | 114 ++++++++++++++++++ tests/api/test_pat_rules.py | 1 + .../test_personal_access_tokens.py | 26 ++++ 10 files changed, 275 insertions(+), 4 deletions(-) diff --git a/docs/content/Extensions/personal-access-tokens.mdx b/docs/content/Extensions/personal-access-tokens.mdx index 7e5be596..e3ae3647 100644 --- a/docs/content/Extensions/personal-access-tokens.mdx +++ b/docs/content/Extensions/personal-access-tokens.mdx @@ -20,7 +20,7 @@ Every token is limited in three ways: 3. Optionally restrict it to specific resources and pick an expiry. 4. Copy the token. It starts with `dgpt_pat_` and is shown **once**. DocsGPT stores only a hash of it, so a lost token cannot be recovered. Revoke it and create a new one. -Tokens can only be created and revoked from a signed-in session. A token cannot create, list or revoke tokens, so a leaked token cannot mint a replacement for itself. +Tokens can only be created, regenerated and revoked from a signed-in session. A token cannot create, list, regenerate or revoke tokens, so a leaked token cannot mint a replacement for itself. ## Using a token @@ -159,11 +159,12 @@ Restrictions and `agents apply`: a token restricted to specific agents can apply - A token created without an explicit lifetime expires after `PAT_DEFAULT_LIFETIME_DAYS` (90 by default). Users can choose any lifetime up to `PAT_MAX_LIFETIME_DAYS` (365 by default). - Non-expiring tokens are available only when the operator sets `PAT_ALLOW_NON_EXPIRING=true`. +- **Regenerate** in **Settings → Access Tokens** issues a new secret for the same token and resets its expiry. The name, scopes and restrictions stay; the old secret stops working immediately, so update whatever uses it. The new lifetime defaults to the one the token was last issued with, and an expired token can be renewed this way (a revoked one cannot). This is the way to rotate a secret or extend a token without rebuilding its scopes. - Revoking a token in **Settings → Access Tokens** takes effect on the next request. - Admins can list a user's tokens with `GET /api/admin/users//tokens` and revoke any token with `DELETE /api/admin/tokens/`. The admin **revoke sessions** action also revokes all of that user's tokens. - Tokens of a deactivated user (through the admin API or SCIM) stop working immediately and work again if the user is reactivated. - `GET /api/user/tokens` reports a token past its expiry as `"status": "expired"`. -- Token creation and revocation are recorded in the authentication audit log (`pat_created`, `pat_revoked`), visible to admins. +- Token creation and revocation are recorded in the authentication audit log (`pat_created`, `pat_regenerated`, `pat_revoked`), visible to admins. An expired token's name can be reused: creating a token with that name retires the expired one. @@ -191,6 +192,7 @@ These endpoints need a signed-in session and cannot be called with a token. | --- | --- | | `GET /api/user/tokens` | List your tokens, the scope catalog and the server's token policy | | `POST /api/user/tokens` | Create a token. Body: `name`, `scopes`, optional `resource_filter`, optional `expires_in_days` (`0` = never, when allowed). The response carries the plaintext `token` once | +| `POST /api/user/tokens//regenerate` | New secret and new expiry for the same token. Optional body `expires_in_days`; omitted = the lifetime it was last issued with. The response carries the plaintext `token` once | | `DELETE /api/user/tokens/` | Revoke a token | ## Good practice diff --git a/docsgpt/alembic/versions/0032_personal_access_tokens.py b/docsgpt/alembic/versions/0032_personal_access_tokens.py index 7d5090bf..6135bc3d 100644 --- a/docsgpt/alembic/versions/0032_personal_access_tokens.py +++ b/docsgpt/alembic/versions/0032_personal_access_tokens.py @@ -10,7 +10,8 @@ user can tell their tokens apart in the UI. itself). ``resource_filter`` optionally narrows a resource family to specific ids, e.g. ``{"agents": [""]}``; an absent family is unrestricted within the token's scopes. ``expires_at`` is NULL only when the operator allows -non-expiring tokens. +non-expiring tokens. Regenerating a token swaps its secret in place and stamps +``regenerated_at``; the row, its name, scopes and restrictions stay. ``user_id`` is the auth ``sub``; no FK or trigger, mirroring ``devices`` and ``user_roles`` so a token row never blocks user deletion. @@ -47,6 +48,7 @@ def upgrade() -> None: last_used_at TIMESTAMPTZ, last_used_ip TEXT, created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + regenerated_at TIMESTAMPTZ, revoked_at TIMESTAMPTZ, revoke_reason TEXT ); diff --git a/docsgpt/api/pat/routes.py b/docsgpt/api/pat/routes.py index 7c2ecc36..b8a9e00d 100644 --- a/docsgpt/api/pat/routes.py +++ b/docsgpt/api/pat/routes.py @@ -23,6 +23,7 @@ from docsgpt.api.pat.tokens import ( is_pat, normalize_resource_filter, normalize_scopes, + renewal_lifetime_days, resolve_expiry, ) from docsgpt.api.user.authz import admin_required @@ -84,6 +85,7 @@ def serialize_token(row: dict) -> dict: "last_used_at": row.get("last_used_at"), "last_used_ip": row.get("last_used_ip"), "created_at": row.get("created_at"), + "regenerated_at": row.get("regenerated_at"), "revoked_at": row.get("revoked_at"), } @@ -214,6 +216,66 @@ class PersonalAccessToken(Resource): return make_response(jsonify({"success": True}), 200) +@pat_ns.route("/user/tokens//regenerate") +class PersonalAccessTokenRegenerate(Resource): + def post(self, token_id): + """Issue a new secret for a token and reset its expiry. + + Name, scopes and restrictions stay; the old secret stops working at + once. ``expires_in_days`` is optional and defaults to the lifetime the + token was last issued with. An expired token can be renewed this way; a + revoked one cannot. The plaintext ``token`` is returned here and never again. + """ + user_id = _session_user_id() + if not user_id: + return _error("Authentication required", 401) + if not auth_type_supports_pats(): + return _error("Personal access tokens are not available on this server", 403) + if not _valid_uuid(token_id): + return _error("Token not found", 404) + body = request.get_json(silent=True) + if body is None: + body = {} + if not isinstance(body, dict): + return _error("Request body must be a JSON object", 400) + + token, token_hash, token_prefix = generate_token() + with db_session() as conn: + repo = PersonalAccessTokensRepository(conn) + current = repo.get(token_id, user_id) + if not current or current["status"] != "active": + return _error("Token not found", 404) + requested = body.get("expires_in_days") + if requested is None: + requested = renewal_lifetime_days(current) + try: + expires_at = resolve_expiry(requested) + except ValueError as exc: + return _error(str(exc), 400) + row = repo.regenerate( + token_id, user_id, token_hash=token_hash, token_prefix=token_prefix, expires_at=expires_at + ) + if not row: + # Revoked between the read and the write. + return _error("Token not found", 404) + AuthEventsRepository(conn).insert( + user_id, + "pat_regenerated", + ip=request.remote_addr, + user_agent=request.headers.get("User-Agent"), + metadata={ + "token_id": token_id, + "name": row["name"], + "expires_at": row.get("expires_at"), + "previous_expires_at": current.get("expires_at"), + }, + ) + return make_response( + jsonify({"success": True, "token": token, "personal_access_token": serialize_token(row)}), + 200, + ) + + @pat_ns.route("/admin/users//tokens") class AdminUserTokens(Resource): @admin_required diff --git a/docsgpt/api/pat/rules.py b/docsgpt/api/pat/rules.py index 784f764c..9b69d1a7 100644 --- a/docsgpt/api/pat/rules.py +++ b/docsgpt/api/pat/rules.py @@ -321,6 +321,7 @@ DENIED: dict[str, tuple[str, ...]] = { "/": ("*",), "/api/user/tokens": ("*",), "/api/user/tokens/": ("*",), + "/api/user/tokens//regenerate": ("*",), "/api/generate_token": ("*",), "/api/combine": ("*",), "/api/download": ("*",), diff --git a/docsgpt/api/pat/tokens.py b/docsgpt/api/pat/tokens.py index 6e4978bd..6a2050a9 100644 --- a/docsgpt/api/pat/tokens.py +++ b/docsgpt/api/pat/tokens.py @@ -169,6 +169,33 @@ def resolve_expiry(expires_in_days: Any) -> Optional[datetime]: return datetime.now(timezone.utc) + timedelta(days=days) +def _parse_moment(value: Any) -> Optional[datetime]: + if not value: + return None + try: + moment = value if isinstance(value, datetime) else datetime.fromisoformat(str(value)) + except ValueError: + return None + return moment if moment.tzinfo else moment.replace(tzinfo=timezone.utc) + + +def renewal_lifetime_days(row: dict) -> Optional[int]: + """The lifetime a token was last issued with, for renewing it on the same terms. + + ``0`` for a non-expiring token, ``None`` when it cannot be derived (the + caller then falls back to the default). The result is clamped to today's + maximum, since the policy may have tightened since the token was issued. + """ + issued = _parse_moment(row.get("regenerated_at")) or _parse_moment(row.get("created_at")) + expires = _parse_moment(row.get("expires_at")) + if expires is None: + return 0 if settings.PAT_ALLOW_NON_EXPIRING else None + if issued is None: + return None + days = round((expires - issued).total_seconds() / 86400) + return max(1, min(days, settings.PAT_MAX_LIFETIME_DAYS)) + + def _client_ip(request) -> Optional[str]: # Flask exposes remote_addr; Starlette exposes client.host. ip = getattr(request, "remote_addr", None) diff --git a/docsgpt/storage/db/models.py b/docsgpt/storage/db/models.py index 65527c76..a0df1a06 100644 --- a/docsgpt/storage/db/models.py +++ b/docsgpt/storage/db/models.py @@ -1098,6 +1098,7 @@ personal_access_tokens_table = Table( Column("last_used_at", DateTime(timezone=True)), Column("last_used_ip", Text), Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()), + Column("regenerated_at", DateTime(timezone=True)), Column("revoked_at", DateTime(timezone=True)), Column("revoke_reason", Text), CheckConstraint("status IN ('active', 'revoked')", name="personal_access_tokens_status_check"), diff --git a/docsgpt/storage/db/repositories/personal_access_tokens.py b/docsgpt/storage/db/repositories/personal_access_tokens.py index 12bd80cc..8d161bf4 100644 --- a/docsgpt/storage/db/repositories/personal_access_tokens.py +++ b/docsgpt/storage/db/repositories/personal_access_tokens.py @@ -14,7 +14,7 @@ from docsgpt.storage.db.base_repository import row_to_dict # token_hash never leaves the repository except through find_active_by_hash. _PUBLIC_COLUMNS = ( "id, user_id, name, token_prefix, scopes, resource_filter, status, " - "expires_at, last_used_at, last_used_ip, created_at, revoked_at, revoke_reason" + "expires_at, last_used_at, last_used_ip, created_at, regenerated_at, revoked_at, revoke_reason" ) @@ -155,6 +155,41 @@ class PersonalAccessTokensRepository: {"id": token_id, "ip": ip, "min_interval": min_interval_seconds}, ) + def regenerate( + self, + token_id: str, + user_id: str, + *, + token_hash: str, + token_prefix: str, + expires_at: Optional[datetime], + ) -> Optional[dict]: + """Swap a live token's secret and expiry in place. The old secret stops matching at once. + + Expired tokens qualify (renewal is the point); revoked ones do not. + Usage fields are cleared because they described the old secret. + """ + row = self._conn.execute( + text( + f""" + UPDATE personal_access_tokens + SET token_hash = :token_hash, token_prefix = :token_prefix, + expires_at = :expires_at, regenerated_at = now(), + last_used_at = NULL, last_used_ip = NULL + WHERE id = CAST(:id AS uuid) AND user_id = :user_id AND status = 'active' + RETURNING {_PUBLIC_COLUMNS} + """ + ), + { + "id": token_id, + "user_id": user_id, + "token_hash": token_hash, + "token_prefix": token_prefix, + "expires_at": expires_at, + }, + ).fetchone() + return row_to_dict(row) if row is not None else None + def revoke(self, token_id: str, user_id: Optional[str] = None, *, reason: str = "user_revoked") -> bool: """Revoke one token. ``user_id=None`` is the admin path (any owner).""" sql = ( diff --git a/tests/api/test_pat_routes.py b/tests/api/test_pat_routes.py index cf21c8c5..f4530553 100644 --- a/tests/api/test_pat_routes.py +++ b/tests/api/test_pat_routes.py @@ -223,6 +223,120 @@ class TestEndToEnd: assert response.status_code == 401 +class TestRegenerate: + def _regen(self, client, token_id, **body): + with _session(): + return client.post(f"/api/user/tokens/{token_id}/regenerate", json=body) + + def test_new_secret_works_old_one_stops_and_the_rest_is_kept(self, client, db): + created = json.loads( + _create(client, scopes=["prompts:read"], resource_filter={"prompts": [AGENT_A]}).data + ) + old, token_id = created["token"], created["personal_access_token"]["id"] + response = self._regen(client, token_id) + assert response.status_code == 200 + body = json.loads(response.data) + new, public = body["token"], body["personal_access_token"] + assert new.startswith("dgpt_pat_") and new != old + assert public["id"] == token_id and public["name"] == "ci" + assert public["scopes"] == ["prompts:read"] + assert public["resource_filter"] == {"prompts": [AGENT_A]} + assert public["token_prefix"] == new[:15] + assert public["regenerated_at"] is not None + assert public["last_used_at"] is None + + assert client.get("/api/user/me", headers={"Authorization": f"Bearer {old}"}).status_code == 401 + me = client.get("/api/user/me", headers={"Authorization": f"Bearer {new}"}) + assert me.status_code == 200 + assert json.loads(me.data)["token"]["id"] == token_id + stored = db.execute(text("SELECT token_hash FROM personal_access_tokens")).scalar_one() + assert stored == pat_tokens.hash_token(new) + + def test_expiry_is_reset_to_the_original_lifetime(self, client, db): + token_id = json.loads(_create(client, expires_in_days=30).data)["personal_access_token"]["id"] + # 20 days in: 10 days left. + db.execute( + text( + "UPDATE personal_access_tokens SET created_at = now() - interval '20 days', " + "expires_at = now() + interval '10 days'" + ) + ) + self._regen(client, token_id) + days_left = db.execute( + text("SELECT extract(epoch FROM expires_at - now()) / 86400 FROM personal_access_tokens") + ).scalar_one() + assert 29.9 < float(days_left) < 30.1 + + def test_explicit_lifetime_is_honoured_and_capped(self, client, db): + token_id = json.loads(_create(client).data)["personal_access_token"]["id"] + assert self._regen(client, token_id, expires_in_days=7).status_code == 200 + days_left = db.execute( + text("SELECT extract(epoch FROM expires_at - now()) / 86400 FROM personal_access_tokens") + ).scalar_one() + assert 6.9 < float(days_left) < 7.1 + assert self._regen(client, token_id, expires_in_days=366).status_code == 400 + assert self._regen(client, token_id, expires_in_days=0).status_code == 400 + + def test_an_expired_token_can_be_renewed(self, client, db): + created = json.loads(_create(client, expires_in_days=30).data) + db.execute( + text( + "UPDATE personal_access_tokens SET created_at = now() - interval '31 days', " + "expires_at = now() - interval '1 day'" + ) + ) + old_headers = {"Authorization": f"Bearer {created['token']}"} + assert client.get("/api/user/me", headers=old_headers).status_code == 401 + body = json.loads(self._regen(client, created["personal_access_token"]["id"]).data) + assert body["personal_access_token"]["status"] == "active" + headers = {"Authorization": f"Bearer {body['token']}"} + assert client.get("/api/user/me", headers=headers).status_code == 200 + + def test_non_expiring_token_stays_non_expiring_only_while_the_operator_allows_it( + self, client, db, monkeypatch + ): + monkeypatch.setattr(pat_tokens.settings, "PAT_ALLOW_NON_EXPIRING", True) + token_id = json.loads(_create(client, expires_in_days=0).data)["personal_access_token"]["id"] + assert json.loads(self._regen(client, token_id).data)["personal_access_token"]["expires_at"] is None + monkeypatch.setattr(pat_tokens.settings, "PAT_ALLOW_NON_EXPIRING", False) + renewed = json.loads(self._regen(client, token_id).data)["personal_access_token"] + assert renewed["expires_at"] is not None # falls back to the default lifetime + + def test_revoked_foreign_unknown_and_malformed_tokens_are_not_found(self, client, db): + token_id = json.loads(_create(client).data)["personal_access_token"]["id"] + with _session(sub="bob"): + assert client.post(f"/api/user/tokens/{token_id}/regenerate").status_code == 404 + with _session(): + client.delete(f"/api/user/tokens/{token_id}") + assert self._regen(client, token_id).status_code == 404 + assert self._regen(client, AGENT_A).status_code == 404 + assert self._regen(client, f"urn:uuid:{AGENT_A}").status_code == 404 + + def test_a_token_cannot_regenerate_itself_or_others(self, client, db): + created = json.loads(_create(client, scopes=list(pat_tokens.SCOPES)).data) + headers = {"Authorization": f"Bearer {created['token']}"} + token_id = created["personal_access_token"]["id"] + response = client.post(f"/api/user/tokens/{token_id}/regenerate", headers=headers) + assert response.status_code == 403 + assert json.loads(response.data)["error"] == "not_available_to_tokens" + + def test_requires_a_session_and_an_object_body(self, client, db): + token_id = json.loads(_create(client).data)["personal_access_token"]["id"] + with patch("docsgpt.app.handle_auth", return_value=None): + assert client.post(f"/api/user/tokens/{token_id}/regenerate").status_code == 401 + with _session(): + assert client.post(f"/api/user/tokens/{token_id}/regenerate", json=[1]).status_code == 400 + + def test_audited_without_the_secret(self, client, db): + token_id = json.loads(_create(client).data)["personal_access_token"]["id"] + self._regen(client, token_id) + metadata = db.execute( + text("SELECT metadata FROM auth_events WHERE event = 'pat_regenerated'") + ).scalar_one() + assert metadata["token_id"] == token_id + assert "dgpt_pat_" not in json.dumps(metadata) + + class TestRevoke: def test_cannot_revoke_someone_elses_token(self, client, db): token_id = json.loads(_create(client).data)["personal_access_token"]["id"] diff --git a/tests/api/test_pat_rules.py b/tests/api/test_pat_rules.py index 69e532a3..a1bb0977 100644 --- a/tests/api/test_pat_rules.py +++ b/tests/api/test_pat_rules.py @@ -91,6 +91,7 @@ class TestClassification: ("/api/user/tokens", "POST"), ("/api/user/tokens", "GET"), ("/api/user/tokens/", "DELETE"), + ("/api/user/tokens//regenerate", "POST"), ("/api/admin/users", "GET"), ("/api/admin/tokens/", "DELETE"), ("/api/generate_token", "GET"), diff --git a/tests/storage/db/repositories/test_personal_access_tokens.py b/tests/storage/db/repositories/test_personal_access_tokens.py index 0f4732f0..8d2ca531 100644 --- a/tests/storage/db/repositories/test_personal_access_tokens.py +++ b/tests/storage/db/repositories/test_personal_access_tokens.py @@ -179,3 +179,29 @@ class TestCountAndUsage: assert repo.get(str(row["id"]))["last_used_ip"] == "10.0.0.1" repo.touch_last_used(str(row["id"]), "10.0.0.3", min_interval_seconds=0) assert repo.get(str(row["id"]))["last_used_ip"] == "10.0.0.3" + + +class TestRegenerate: + def test_swaps_the_secret_and_resets_expiry_and_usage(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + row = _create(repo, expires_at=datetime.now(timezone.utc) - timedelta(days=1)) + repo.touch_last_used(str(row["id"]), "10.0.0.1") + future = datetime.now(timezone.utc) + timedelta(days=30) + renewed = repo.regenerate( + str(row["id"]), "u1", token_hash="h-new", token_prefix="dgpt_pat_new123", expires_at=future + ) + assert renewed["id"] == row["id"] and renewed["name"] == "ci" + assert renewed["token_prefix"] == "dgpt_pat_new123" + assert renewed["regenerated_at"] is not None + assert renewed["last_used_at"] is None and renewed["last_used_ip"] is None + assert repo.find_active_by_hash("h1") is None + assert repo.find_active_by_hash("h-new")["id"] == row["id"] + + def test_owner_scoped_and_never_revives_a_revoked_token(self, pg_conn): + repo = PersonalAccessTokensRepository(pg_conn) + row = _create(repo) + kwargs = dict(token_hash="h-new", token_prefix="dgpt_pat_new123", expires_at=None) + assert repo.regenerate(str(row["id"]), "someone-else", **kwargs) is None + repo.revoke(str(row["id"]), "u1") + assert repo.regenerate(str(row["id"]), "u1", **kwargs) is None + assert repo.find_active_by_hash("h-new") is None From 52069e3f8a8a2aec4a22234712b84c33e2510d68 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Mon, 21 Sep 2026 10:32:17 +0100 Subject: [PATCH 115/130] feat(frontend): regenerate access tokens; fix resource pickers inside the create modal Each token gets a Regenerate action: a confirmation with the new expiration (preselecting the lifetime the token was issued with), then the one-time secret view. The "restrict to specific resources" pickers could not be scrolled and did not close on an outside click. Inside a Modal a non-modal popover is portalled outside the dialog, so the dialog's scroll lock swallowed the wheel, and Radix defers its outside-click dismissal to the document click, which Modal stops from propagating. MultiSelect takes a `modal` prop for that case. --- frontend/src/api/endpoints.ts | 2 + frontend/src/api/services/patService.test.ts | 44 ++++ frontend/src/api/services/patService.ts | 20 ++ frontend/src/components/ui/multi-select.tsx | 12 +- frontend/src/locale/de.json | 14 ++ frontend/src/locale/en.json | 14 ++ frontend/src/locale/es.json | 14 ++ frontend/src/locale/jp.json | 14 ++ frontend/src/locale/ru.json | 14 ++ frontend/src/locale/zh-TW.json | 14 ++ frontend/src/locale/zh.json | 14 ++ .../src/modals/AccessTokenCreatedModal.tsx | 12 +- .../src/modals/CreateAccessTokenModal.tsx | 1 + .../src/modals/RegenerateAccessTokenModal.tsx | 191 ++++++++++++++++++ .../src/settings/PersonalAccessTokens.tsx | 53 ++++- .../src/settings/accessTokenUtils.test.ts | 57 ++++++ frontend/src/settings/accessTokenUtils.ts | 25 +++ 17 files changed, 509 insertions(+), 6 deletions(-) create mode 100644 frontend/src/modals/RegenerateAccessTokenModal.tsx diff --git a/frontend/src/api/endpoints.ts b/frontend/src/api/endpoints.ts index cd19899c..6f5f36e6 100644 --- a/frontend/src/api/endpoints.ts +++ b/frontend/src/api/endpoints.ts @@ -154,6 +154,8 @@ const endpoints = { `/api/devices/pairings/${deviceCode}`, ACCESS_TOKENS: '/api/user/tokens', ACCESS_TOKEN: (id: string) => `/api/user/tokens/${id}`, + ACCESS_TOKEN_REGENERATE: (id: string) => + `/api/user/tokens/${id}/regenerate`, }, V1: { CHAT_COMPLETIONS: '/v1/chat/completions', diff --git a/frontend/src/api/services/patService.test.ts b/frontend/src/api/services/patService.test.ts index 63b7a88e..e6f479c9 100644 --- a/frontend/src/api/services/patService.test.ts +++ b/frontend/src/api/services/patService.test.ts @@ -154,3 +154,47 @@ describe('patService.revoke', () => { ); }); }); + +describe('patService.regenerate', () => { + it('POSTs the chosen lifetime to the regenerate endpoint', async () => { + const body = { + success: true, + token: 'dgpt_pat_newsecret', + personal_access_token: TOKEN_ROW, + }; + const spy = vi.spyOn(apiClient, 'post').mockResolvedValue(response(body)); + + const result = await patService.regenerate(TOKEN_ROW.id, 30, 'session-jwt'); + + expect(spy).toHaveBeenCalledWith( + `/api/user/tokens/${TOKEN_ROW.id}/regenerate`, + { expires_in_days: 30 }, + 'session-jwt', + ); + expect(result.token).toBe('dgpt_pat_newsecret'); + }); + + it('sends an empty body to keep the original lifetime', async () => { + const spy = vi + .spyOn(apiClient, 'post') + .mockResolvedValue(response({ success: true, token: 't' })); + + await patService.regenerate(TOKEN_ROW.id, undefined, null); + + expect(spy).toHaveBeenCalledWith( + `/api/user/tokens/${TOKEN_ROW.id}/regenerate`, + {}, + null, + ); + }); + + it('surfaces the server message', async () => { + vi.spyOn(apiClient, 'post').mockResolvedValue( + response({ success: false, message: 'Token not found' }, 404), + ); + + await expect(patService.regenerate('missing', 30, null)).rejects.toThrow( + 'Token not found', + ); + }); +}); diff --git a/frontend/src/api/services/patService.ts b/frontend/src/api/services/patService.ts index 09eb197a..fda7017f 100644 --- a/frontend/src/api/services/patService.ts +++ b/frontend/src/api/services/patService.ts @@ -14,6 +14,8 @@ export interface PersonalAccessToken { last_used_at: string | null; last_used_ip: string | null; created_at: string | null; + /** Set once the secret has been regenerated; the lifetime then counts from here. */ + regenerated_at?: string | null; revoked_at: string | null; } @@ -95,6 +97,24 @@ const patService = { await apiClient.post(endpoints.USER.ACCESS_TOKENS, payload, token), ), + /** + * New secret for the same token (name, scopes and restrictions stay); the old + * secret stops working at once. `expiresInDays` omitted = the lifetime the + * token was last issued with. + */ + regenerate: async ( + id: string, + expiresInDays: number | undefined, + token: string | null, + ): Promise => + parse( + await apiClient.post( + endpoints.USER.ACCESS_TOKEN_REGENERATE(encodeURIComponent(id)), + expiresInDays === undefined ? {} : { expires_in_days: expiresInDays }, + token, + ), + ), + revoke: async ( id: string, token: string | null, diff --git a/frontend/src/components/ui/multi-select.tsx b/frontend/src/components/ui/multi-select.tsx index 826c658d..e436dd8f 100644 --- a/frontend/src/components/ui/multi-select.tsx +++ b/frontend/src/components/ui/multi-select.tsx @@ -32,6 +32,15 @@ interface MultiSelectProps { emptyText?: string; searchPlaceholder?: string; className?: string; + /** + * Set when the MultiSelect sits inside a Modal. A non-modal popover there + * cannot scroll (the dialog's scroll lock swallows the wheel, since the + * dropdown is portalled outside it) and never closes on an outside click + * (Radix defers that to the document `click`, which Modal stops from + * propagating). A modal popover owns its own scroll lock and dismisses on + * pointerdown instead. + */ + modal?: boolean; } export function MultiSelect({ @@ -42,6 +51,7 @@ export function MultiSelect({ emptyText = 'No results found.', searchPlaceholder = 'Search...', className, + modal = false, }: MultiSelectProps) { const [open, setOpen] = React.useState(false); @@ -63,7 +73,7 @@ export function MultiSelect({ .map((option) => option.label); return ( - + + + + } + > +
+
+

+ {t('settings.accessTokens.regenerate.title')} +

+

+ {t('settings.accessTokens.regenerate.warning', { + name: item?.name ?? '', + ...NO_ESCAPE, + })} +

+
+ +
+ + +

+ {expiry === NO_EXPIRY + ? t('settings.accessTokens.create.noExpirationHint') + : t('settings.accessTokens.create.expiresOn', { + date: formatDateOnly( + new Date(Date.now() + expiry * DAY_MS).toISOString(), + ), + ...NO_ESCAPE, + })} +

+
+ + {error && ( +

+ {error} +

+ )} +
+ + ); +} diff --git a/frontend/src/settings/PersonalAccessTokens.tsx b/frontend/src/settings/PersonalAccessTokens.tsx index f417200c..27cebbb0 100644 --- a/frontend/src/settings/PersonalAccessTokens.tsx +++ b/frontend/src/settings/PersonalAccessTokens.tsx @@ -28,6 +28,7 @@ import { useDarkTheme } from '../hooks'; import AccessTokenCreatedModal from '../modals/AccessTokenCreatedModal'; import ConfirmationModal from '../modals/ConfirmationModal'; import CreateAccessTokenModal from '../modals/CreateAccessTokenModal'; +import RegenerateAccessTokenModal from '../modals/RegenerateAccessTokenModal'; import { ActiveState } from '../models/misc'; import { selectToken } from '../preferences/preferenceSlice'; import { formatDateOnly, formatDateTime } from '../utils/dateTimeUtils'; @@ -89,6 +90,10 @@ export default function PersonalAccessTokens() { // The plaintext secret lives only here, and only until the modal closes. const [created, setCreated] = React.useState(null); + // True while `created` holds a regenerated (not brand-new) secret. + const [createdByRegenerate, setCreatedByRegenerate] = React.useState(false); + const [tokenToRegenerate, setTokenToRegenerate] = + React.useState(null); const [revokeState, setRevokeState] = React.useState('INACTIVE'); const [tokenToRevoke, setTokenToRevoke] = React.useState(null); @@ -116,8 +121,20 @@ export default function PersonalAccessTokens() { loadTokens(true); }, [loadTokens]); + const handleRegenerated = (response: CreateAccessTokenResponse) => { + const updated = response.personal_access_token; + setTokenToRegenerate(null); + setCreatedByRegenerate(true); + setCreated(response); + setTokens((prev) => + prev.map((item) => (item.id === updated.id ? updated : item)), + ); + setError(null); + }; + const handleCreated = (response: CreateAccessTokenResponse) => { setCreateOpen(false); + setCreatedByRegenerate(false); setCreated(response); setTokens((prev) => [response.personal_access_token, ...prev]); }; @@ -212,6 +229,23 @@ export default function PersonalAccessTokens() { ); }; + const renderRegenerateButton = (item: PersonalAccessToken) => + policy?.enabled ? ( + + ) : null; + const renderRevokeButton = (item: PersonalAccessToken) => ( + ) : null} + + + + ); +} diff --git a/frontend/src/admin/Quotas.tsx b/frontend/src/admin/Quotas.tsx new file mode 100644 index 00000000..7d5d93ed --- /dev/null +++ b/frontend/src/admin/Quotas.tsx @@ -0,0 +1,343 @@ +import { useCallback, useEffect, useMemo, useState } from 'react'; +import { useSelector } from 'react-redux'; + +import adminService, { type QuotaScope } from '../api/services/adminService'; +import teamsService from '../api/services/teamsService'; +import { Button } from '../components/ui/button'; +import { Modal } from '../components/ui/modal'; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from '../components/ui/select'; +import { + Table, + TableBody, + TableCell, + TableContainer, + TableHead, + TableHeader, + TableRow, +} from '../components/ui/table'; +import { selectToken } from '../preferences/preferenceSlice'; +import { + LoadError, + Loading, + Pill, + fmtDate, + fmtNumber, + fmtRelative, +} from './AdminUI'; +import QuotaEditor from './QuotaEditor'; +import { describeBudget, type QuotaPolicy } from './quotaUtils'; + +type TeamPolicy = QuotaPolicy & { + team_name?: string | null; + team_slug?: string | null; + member_count?: number | null; +}; + +type Editing = { + scope: QuotaScope; + subjectId: string | null; + title: string; + policy: QuotaPolicy | null; +}; + +const HINTS: Record = { + instance: + 'Applies to every user that no team allowance or user override covers.', + team: 'Each member gets this allowance; it is not a shared pool. A member of several teams gets the most generous one.', + user: 'Overrides team allowances and the instance default for this user.', +}; + +function PolicyCells({ policy }: { policy: QuotaPolicy }) { + return ( + <> + + {describeBudget(policy.token_limit, policy.token_unlimited, 'tokens')} + + + {describeBudget(policy.cost_limit_usd, policy.cost_unlimited, 'cost')} + + + {policy.note || '—'} + + + {fmtRelative(policy.updated_at)} + + + ); +} + +export default function Quotas() { + const token = useSelector(selectToken); + const [data, setData] = useState(null); + const [teams, setTeams] = useState([]); + const [loading, setLoading] = useState(true); + const [editing, setEditing] = useState(null); + const [teamPick, setTeamPick] = useState(''); + + const load = useCallback(async () => { + setLoading(true); + try { + const [quotasRes, teamsJson] = await Promise.all([ + adminService.getQuotas(token), + teamsService.listAll(token).catch(() => ({})), + ]); + setData(await quotasRes.json().catch(() => ({ success: false }))); + setTeams(teamsJson?.teams ?? []); + } catch { + setData({ success: false }); + } finally { + setLoading(false); + } + }, [token]); + + useEffect(() => { + load(); + }, [load]); + + // The editor covers the ``all`` bucket; other buckets are listed read-only. + const isAll = (p: QuotaPolicy) => p.bucket === 'all'; + const instancePolicy: QuotaPolicy | null = + (data?.instance ?? []).find(isAll) ?? null; + const teamPolicies: TeamPolicy[] = data?.teams ?? []; + const userPolicies: QuotaPolicy[] = data?.users ?? []; + const teamsWithoutPolicy = useMemo(() => { + const covered = new Set( + teamPolicies.filter(isAll).map((p) => String(p.subject_id)), + ); + return teams.filter((team) => !covered.has(String(team.id))); + }, [teams, teamPolicies]); + + if (data === null && loading) return ; + if (!data?.success) return ; + + const bucketPill = (policy: QuotaPolicy) => + isAll(policy) ? null : {policy.bucket} traffic; + + return ( +
+

+ Usage is counted per user over each calendar {data.period} (UTC). The + current window resets {fmtDate(data.resets_at)}. A request is refused + once a budget is used up; the request that crosses it still completes. +

+ + {(data.unpriced_models ?? []).length > 0 ? ( +
+

+ Models without a price are invisible to cost limits +

+

+ These were used this {data.period} and recorded at $0:{' '} + {(data.unpriced_models as any[]) + .map((m) => `${m.model_id} (${fmtNumber(m.tokens)} tokens)`) + .join(', ')} + . Use a token limit for them, or declare their rates in the model + catalog. +

+
+ ) : null} + +
+
+

Instance default

+ +
+

+ {instancePolicy + ? `Tokens: ${describeBudget(instancePolicy.token_limit, instancePolicy.token_unlimited, 'tokens')} · Cost: ${describeBudget(instancePolicy.cost_limit_usd, instancePolicy.cost_unlimited, 'cost')}` + : 'No default: users without a team allowance or override are unlimited.'} +

+
+ +
+
+

Team allowances

+ {teamsWithoutPolicy.length > 0 ? ( +
+ + +
+ ) : null} +
+ {teamPolicies.length === 0 ? ( +

+ No team has an allowance. +

+ ) : ( + + + + + Team + Members + Tokens + Cost + Note + Updated + Actions + + + + {teamPolicies.map((policy) => ( + + + + {policy.team_name ?? policy.subject_id} + + {bucketPill(policy)} + + + {fmtNumber(policy.member_count)} + + + + {isAll(policy) ? ( + + ) : null} + + + ))} + +
+
+ )} +
+ +
+

User overrides

+ {userPolicies.length === 0 ? ( +

+ No user has an override. Add one from a user's menu on the + Users tab. +

+ ) : ( + + + + + User + Tokens + Cost + Note + Updated + Actions + + + + {userPolicies.map((policy) => ( + + + + {policy.subject_id} + + {bucketPill(policy)} + + + + {isAll(policy) ? ( + + ) : null} + + + ))} + +
+
+ )} +
+ + { + if (!open) setEditing(null); + }} + title={editing ? `Quota · ${editing.title}` : 'Quota'} + > + {editing ? ( + { + setEditing(null); + setTeamPick(''); + load(); + }} + /> + ) : null} + +
+ ); +} diff --git a/frontend/src/admin/UserQuotaModal.tsx b/frontend/src/admin/UserQuotaModal.tsx new file mode 100644 index 00000000..f6ce58b2 --- /dev/null +++ b/frontend/src/admin/UserQuotaModal.tsx @@ -0,0 +1,115 @@ +import { useCallback, useEffect, useState } from 'react'; +import { useSelector } from 'react-redux'; + +import adminService from '../api/services/adminService'; +import teamsService from '../api/services/teamsService'; +import { Modal } from '../components/ui/modal'; +import { selectToken } from '../preferences/preferenceSlice'; +import { LoadError, Loading, fmtDate } from './AdminUI'; +import QuotaEditor, { UsageBar } from './QuotaEditor'; +import { + sourceLabel, + type BucketStatus, + type Budget, + type QuotaPolicy, +} from './quotaUtils'; + +/** A user's effective limits and usage, with the editor for their override. */ +export default function UserQuotaModal({ + userId, + onClose, +}: { + userId: string | null; + onClose: () => void; +}) { + const token = useSelector(selectToken); + const [data, setData] = useState(null); + const [teamNames, setTeamNames] = useState>({}); + + const load = useCallback(async () => { + if (!userId) return; + setData(null); + try { + const [res, teamsJson] = await Promise.all([ + adminService.getUserQuota(userId, token), + teamsService.listAll(token).catch(() => ({})), + ]); + setData(await res.json().catch(() => ({ success: false }))); + setTeamNames( + Object.fromEntries( + (teamsJson?.teams ?? []).map((team: any) => [ + String(team.id), + team.name, + ]), + ), + ); + } catch { + setData({ success: false }); + } + }, [userId, token]); + + useEffect(() => { + load(); + }, [load]); + + const overall: BucketStatus | undefined = (data?.effective ?? []).find( + (status: BucketStatus) => status.bucket === 'all', + ); + const override: QuotaPolicy | null = + (data?.policies ?? []).find((p: QuotaPolicy) => p.bucket === 'all') ?? null; + const caption = (budget: Budget) => + sourceLabel( + budget, + budget.source_id ? teamNames[budget.source_id] : undefined, + ); + + return ( + { + if (!open) onClose(); + }} + title={userId ? `Quota · ${userId}` : 'Quota'} + > + {data === null ? ( + + ) : !data.success ? ( + + ) : ( +
+ {overall ? ( +
+ + +

+ Resets {fmtDate(overall.resets_at)} +

+
+ ) : null} +
+

+ User override +

+ +
+
+ )} +
+ ); +} diff --git a/frontend/src/admin/Users.tsx b/frontend/src/admin/Users.tsx index 13e40d13..d7ae7f8e 100644 --- a/frontend/src/admin/Users.tsx +++ b/frontend/src/admin/Users.tsx @@ -1,5 +1,6 @@ import { Eye, + Gauge, LogOut, ShieldCheck, ShieldOff, @@ -40,6 +41,7 @@ import { fmtNumber, fmtRelative, } from './AdminUI'; +import UserQuotaModal from './UserQuotaModal'; type AdminUser = { user_id: string; @@ -70,6 +72,7 @@ export default function Users() { const [busy, setBusy] = useState(null); const [menuUserId, setMenuUserId] = useState(null); const [detail, setDetail] = useState(null); + const [quotaUserId, setQuotaUserId] = useState(null); const [feedback, setFeedback] = useState<{ ok: boolean; message: string; @@ -157,7 +160,14 @@ export default function Users() { isAdmin: boolean, active: boolean, ): Action[] => { - const acts: Action[] = []; + const acts: Action[] = [ + { + key: 'quota', + label: 'Quota', + icon: Gauge, + perform: () => setQuotaUserId(userId), + }, + ]; if (isAdmin) { acts.push({ key: 'revoke', @@ -433,6 +443,11 @@ export default function Users() { /> ) : null} + setQuotaUserId(null)} + /> + { diff --git a/frontend/src/admin/index.tsx b/frontend/src/admin/index.tsx index 6a147f06..5a21c9f4 100644 --- a/frontend/src/admin/index.tsx +++ b/frontend/src/admin/index.tsx @@ -11,6 +11,7 @@ import { Tabs, TabsList, TabsTrigger } from '../components/ui/tabs'; import Admins from './Admins'; import Audit from './Audit'; import Overview from './Overview'; +import Quotas from './Quotas'; import Usage from './Usage'; import Users from './Users'; @@ -19,6 +20,7 @@ const TABS = [ { key: 'users', label: 'Users', path: '/admin/users' }, { key: 'admins', label: 'Admins', path: '/admin/roles' }, { key: 'usage', label: 'Usage', path: '/admin/usage' }, + { key: 'quotas', label: 'Quotas', path: '/admin/quotas' }, { key: 'audit', label: 'Audit', path: '/admin/audit' }, ]; @@ -63,6 +65,7 @@ export default function Admin() { } /> } /> } /> + } /> } /> } /> diff --git a/frontend/src/admin/quotaUtils.test.ts b/frontend/src/admin/quotaUtils.test.ts new file mode 100644 index 00000000..60e4f468 --- /dev/null +++ b/frontend/src/admin/quotaUtils.test.ts @@ -0,0 +1,120 @@ +import { describe, expect, it } from 'vitest'; + +import { + describeBudget, + formToPolicy, + isEmptyForm, + policyToForm, + sourceLabel, + usagePercent, + type QuotaPolicy, +} from './quotaUtils'; + +const policy = (fields: Partial): QuotaPolicy => ({ + scope: 'user', + subject_id: 'u1', + bucket: 'all', + token_limit: null, + token_unlimited: false, + cost_limit_usd: null, + cost_unlimited: false, + enabled: true, + ...fields, +}); + +describe('policyToForm', () => { + it('starts a missing policy as inherit', () => { + const form = policyToForm(null); + expect(form.tokenMode).toBe('inherit'); + expect(form.costMode).toBe('inherit'); + expect(isEmptyForm(form)).toBe(true); + }); + + it('keeps zero as a limit, not as inherit', () => { + const form = policyToForm(policy({ token_limit: 0, cost_unlimited: true })); + expect(form.tokenMode).toBe('limit'); + expect(form.tokenLimit).toBe('0'); + expect(form.costMode).toBe('unlimited'); + }); +}); + +describe('formToPolicy', () => { + const base = policyToForm(null); + + it('round-trips limits and trims the note', () => { + const result = formToPolicy({ + ...base, + tokenMode: 'limit', + tokenLimit: ' 5000 ', + costMode: 'limit', + costLimit: '2.5', + note: ' trial ', + }); + expect(result).toEqual({ + ok: true, + policy: { + bucket: 'all', + token_limit: 5000, + token_unlimited: false, + cost_limit_usd: 2.5, + cost_unlimited: false, + note: 'trial', + }, + }); + }); + + it('sends unlimited without a limit', () => { + const result = formToPolicy({ + ...base, + tokenMode: 'unlimited', + tokenLimit: '99', + }); + expect(result.ok && result.policy.token_limit).toBeNull(); + expect(result.ok && result.policy.token_unlimited).toBe(true); + }); + + it.each(['', '1.5', '-1', 'abc', '1e3'])( + 'rejects token limit %j', + (tokenLimit) => { + expect(formToPolicy({ ...base, tokenMode: 'limit', tokenLimit }).ok).toBe( + false, + ); + }, + ); + + it.each(['', '-0.01', 'abc', 'Infinity'])( + 'rejects cost limit %j', + (costLimit) => { + expect(formToPolicy({ ...base, costMode: 'limit', costLimit }).ok).toBe( + false, + ); + }, + ); +}); + +describe('usagePercent', () => { + it('handles unlimited, zero and overshoot', () => { + expect(usagePercent(50, null)).toBe(0); + expect(usagePercent(0, 0)).toBe(100); + expect(usagePercent(25, 100)).toBe(25); + expect(usagePercent(500, 100)).toBe(100); + }); +}); + +describe('labels', () => { + it('describes budgets', () => { + expect(describeBudget(null, true, 'tokens')).toBe('Unlimited'); + expect(describeBudget(null, false, 'cost')).toBe('—'); + expect(describeBudget(1000, false, 'tokens')).toContain('tokens'); + }); + + it('names the layer a limit came from', () => { + expect(sourceLabel({ limit: null, used: 0 })).toBe('No limit set'); + expect(sourceLabel({ limit: 1, used: 0, source: 'team' }, 'Eng')).toBe( + 'Team: Eng', + ); + expect(sourceLabel({ limit: 1, used: 0, source: 'default' })).toBe( + 'Plan default', + ); + }); +}); diff --git a/frontend/src/admin/quotaUtils.ts b/frontend/src/admin/quotaUtils.ts new file mode 100644 index 00000000..10038574 --- /dev/null +++ b/frontend/src/admin/quotaUtils.ts @@ -0,0 +1,130 @@ +// Pure helpers behind the quota editor and usage bars. + +export type BudgetMode = 'inherit' | 'limit' | 'unlimited'; + +export type QuotaPolicy = { + scope: 'instance' | 'team' | 'user'; + subject_id: string | null; + bucket: string; + token_limit: number | null; + token_unlimited: boolean; + cost_limit_usd: number | null; + cost_unlimited: boolean; + enabled: boolean; + note?: string | null; + updated_by?: string | null; + updated_at?: string | null; +}; + +export type Budget = { + limit: number | null; + used: number; + source?: string | null; + source_id?: string | null; +}; + +export type BucketStatus = { + bucket: string; + tokens: Budget; + cost: Budget; + resets_at: string; +}; + +export type QuotaForm = { + tokenMode: BudgetMode; + tokenLimit: string; + costMode: BudgetMode; + costLimit: string; + note: string; +}; + +const mode = (limit: number | null, unlimited: boolean): BudgetMode => { + if (unlimited) return 'unlimited'; + return limit === null || limit === undefined ? 'inherit' : 'limit'; +}; + +export function policyToForm(policy?: QuotaPolicy | null): QuotaForm { + return { + tokenMode: policy + ? mode(policy.token_limit, policy.token_unlimited) + : 'inherit', + tokenLimit: policy?.token_limit != null ? String(policy.token_limit) : '', + costMode: policy + ? mode(policy.cost_limit_usd, policy.cost_unlimited) + : 'inherit', + costLimit: + policy?.cost_limit_usd != null ? String(policy.cost_limit_usd) : '', + note: policy?.note ?? '', + }; +} + +export type FormResult = + { ok: true; policy: Record } | { ok: false; error: string }; + +// An empty form (both budgets inherited) is not a policy: the caller deletes instead. +export function isEmptyForm(form: QuotaForm): boolean { + return form.tokenMode === 'inherit' && form.costMode === 'inherit'; +} + +export function formToPolicy(form: QuotaForm): FormResult { + const policy: Record = { + bucket: 'all', + token_limit: null, + token_unlimited: form.tokenMode === 'unlimited', + cost_limit_usd: null, + cost_unlimited: form.costMode === 'unlimited', + note: form.note.trim() || null, + }; + if (form.tokenMode === 'limit') { + const raw = form.tokenLimit.trim(); + if (!/^\d+$/.test(raw)) + return { ok: false, error: 'Token limit must be a whole number.' }; + const tokens = Number(raw); + if (!Number.isSafeInteger(tokens)) + return { ok: false, error: 'Token limit is too large.' }; + policy.token_limit = tokens; + } + if (form.costMode === 'limit') { + const raw = form.costLimit.trim(); + const cost = Number(raw); + if (raw === '' || !Number.isFinite(cost) || cost < 0) + return { ok: false, error: 'Cost limit must be a number, 0 or more.' }; + policy.cost_limit_usd = cost; + } + return { ok: true, policy }; +} + +export function usagePercent(used: number, limit: number | null): number { + if (limit === null || limit === undefined) return 0; + if (limit <= 0) return 100; + return Math.min(100, Math.max(0, (used / limit) * 100)); +} + +export function fmtUsd(value?: number | null): string { + return new Intl.NumberFormat(undefined, { + style: 'currency', + currency: 'USD', + maximumFractionDigits: value != null && value < 1 ? 4 : 2, + }).format(value ?? 0); +} + +export function describeBudget( + limit: number | null, + unlimited: boolean, + kind: 'tokens' | 'cost', +): string { + if (unlimited) return 'Unlimited'; + if (limit === null || limit === undefined) return '—'; + return kind === 'cost' + ? fmtUsd(limit) + : `${new Intl.NumberFormat().format(limit)} tokens`; +} + +export function sourceLabel(budget: Budget, teamName?: string): string { + if (!budget.source) return 'No limit set'; + if (budget.source === 'user') return 'User override'; + if (budget.source === 'team') + return teamName ? `Team: ${teamName}` : 'Team allowance'; + if (budget.source === 'instance') return 'Instance default'; + return 'Plan default'; +} diff --git a/frontend/src/api/endpoints.ts b/frontend/src/api/endpoints.ts index 6f5f36e6..89187528 100644 --- a/frontend/src/api/endpoints.ts +++ b/frontend/src/api/endpoints.ts @@ -2,6 +2,7 @@ const endpoints = { USER: { CONFIG: '/api/config', ME: '/api/user/me', + QUOTA: '/api/user/quota', NEW_TOKEN: '/api/generate_token', OIDC_LOGIN: '/api/auth/oidc/login', OIDC_TOKEN: '/api/auth/oidc/token', @@ -173,6 +174,12 @@ const endpoints = { USAGE: '/api/admin/usage', AUDIT: '/api/admin/audit', DEVICE_AUDIT: '/api/admin/devices/audit', + QUOTAS: '/api/admin/quotas', + QUOTA_INSTANCE: '/api/admin/quotas/instance', + QUOTA_TEAM: (id: string) => + `/api/admin/quotas/teams/${encodeURIComponent(id)}`, + QUOTA_USER: (id: string) => + `/api/admin/quotas/users/${encodeURIComponent(id)}`, }, CONVERSATION: { ANSWER: '/api/answer', diff --git a/frontend/src/api/services/adminService.ts b/frontend/src/api/services/adminService.ts index 7359aed4..adf67e71 100644 --- a/frontend/src/api/services/adminService.ts +++ b/frontend/src/api/services/adminService.ts @@ -10,6 +10,14 @@ const qs = (params: Record): string => { return str ? `?${str}` : ''; }; +export type QuotaScope = 'instance' | 'team' | 'user'; + +const quotaUrl = (scope: QuotaScope, subjectId?: string | null): string => { + if (scope === 'team') return endpoints.ADMIN.QUOTA_TEAM(subjectId ?? ''); + if (scope === 'user') return endpoints.ADMIN.QUOTA_USER(subjectId ?? ''); + return endpoints.ADMIN.QUOTA_INSTANCE; +}; + const adminService = { getOverview: (token: string | null): Promise => apiClient.get(endpoints.ADMIN.OVERVIEW, token), @@ -54,6 +62,23 @@ const adminService = { token: string | null, ): Promise => apiClient.get(`${endpoints.ADMIN.DEVICE_AUDIT}${qs(params)}`, token), + getQuotas: (token: string | null): Promise => + apiClient.get(endpoints.ADMIN.QUOTAS, token), + getUserQuota: (userId: string, token: string | null): Promise => + apiClient.get(endpoints.ADMIN.QUOTA_USER(userId), token), + setQuota: ( + scope: QuotaScope, + subjectId: string | null, + policy: Record, + token: string | null, + ): Promise => apiClient.put(quotaUrl(scope, subjectId), policy, token), + deleteQuota: ( + scope: QuotaScope, + subjectId: string | null, + bucket: string, + token: string | null, + ): Promise => + apiClient.delete(`${quotaUrl(scope, subjectId)}${qs({ bucket })}`, token), }; export default adminService; diff --git a/frontend/src/api/services/userService.ts b/frontend/src/api/services/userService.ts index b9de3bc0..fd9d81f6 100644 --- a/frontend/src/api/services/userService.ts +++ b/frontend/src/api/services/userService.ts @@ -7,6 +7,8 @@ const userService = { throttledApiClient.get(endpoints.USER.CONFIG, null), getMe: (token: string | null): Promise => apiClient.get(endpoints.USER.ME, token), + getQuota: (token: string | null): Promise => + apiClient.get(endpoints.USER.QUOTA, token), getNewToken: (): Promise => throttledApiClient.get(endpoints.USER.NEW_TOKEN, null), // Token deliberately null: a stale Authorization header must not be able diff --git a/frontend/src/conversation/conversationHandlers.ts b/frontend/src/conversation/conversationHandlers.ts index 4578c8d2..dd519628 100644 --- a/frontend/src/conversation/conversationHandlers.ts +++ b/frontend/src/conversation/conversationHandlers.ts @@ -1,7 +1,10 @@ +import i18n from 'i18next'; + import { baseURL } from '../api/client'; import conversationService from '../api/services/conversationService'; import { Doc } from '../models/misc'; import { Answer, FEEDBACK, RetrievalPayload } from './conversationModels'; +import { isQuotaError, quotaErrorMessage } from './quotaError'; import { ToolCallsType } from './types'; /** @@ -48,7 +51,9 @@ async function _handlePreStreamHttpError( if (text) { try { const parsed = JSON.parse(text); - if (parsed && typeof parsed === 'object') { + if (isQuotaError(parsed)) { + message = quotaErrorMessage(parsed, i18n.t.bind(i18n), i18n.language); + } else if (parsed && typeof parsed === 'object') { message = (typeof parsed.message === 'string' && parsed.message) || (typeof parsed.error === 'string' && parsed.error) || diff --git a/frontend/src/conversation/quotaError.test.ts b/frontend/src/conversation/quotaError.test.ts new file mode 100644 index 00000000..78220511 --- /dev/null +++ b/frontend/src/conversation/quotaError.test.ts @@ -0,0 +1,52 @@ +import { describe, expect, it } from 'vitest'; + +import { isQuotaError, quotaErrorMessage } from './quotaError'; + +const t = ((key: string, values: Record) => + `${key}|${values.used}|${values.limit}|${values.resetsAt}`) as any; + +describe('isQuotaError', () => { + it('matches only the quota error code', () => { + expect(isQuotaError({ error_code: 'quota-exceeded' })).toBe(true); + expect(isQuotaError({ message: 'Exceeding usage limit' })).toBe(false); + expect(isQuotaError(null)).toBe(false); + expect(isQuotaError('quota-exceeded')).toBe(false); + }); +}); + +describe('quotaErrorMessage', () => { + it('formats token budgets as numbers', () => { + const message = quotaErrorMessage( + { + dimension: 'tokens', + usage: 1200000, + limit: 1000000, + resets_at: '2026-10-01T00:00:00+00:00', + }, + t, + 'en-US', + ); + const [key, used, limit, resetsAt] = message.split('|'); + expect(key).toBe('conversation.quotaExceeded.tokens'); + expect([used, limit]).toEqual(['1,200,000', '1,000,000']); + expect(resetsAt).not.toBe(''); + }); + + it('formats cost budgets as dollars', () => { + const message = quotaErrorMessage( + { dimension: 'cost', usage: 5.25, limit: 5 }, + t, + 'en-US', + ); + expect(message).toBe('conversation.quotaExceeded.cost|$5.25|$5.00|'); + }); + + it('tolerates a malformed reset time', () => { + const message = quotaErrorMessage( + { dimension: 'tokens', usage: 1, limit: 1, resets_at: 'soon' }, + t, + 'en-US', + ); + expect(message.endsWith('|')).toBe(true); + }); +}); diff --git a/frontend/src/conversation/quotaError.ts b/frontend/src/conversation/quotaError.ts new file mode 100644 index 00000000..86bf46dd --- /dev/null +++ b/frontend/src/conversation/quotaError.ts @@ -0,0 +1,47 @@ +import type { TFunction } from 'i18next'; + +export type QuotaErrorBody = { + error_code?: string; + dimension?: string; + usage?: number; + limit?: number; + resets_at?: string; +}; + +export function isQuotaError(body: unknown): body is QuotaErrorBody { + return ( + !!body && + typeof body === 'object' && + (body as QuotaErrorBody).error_code === 'quota-exceeded' + ); +} + +/** The chat message for a 429 ``quota-exceeded`` body, in the user's language. */ +export function quotaErrorMessage( + body: QuotaErrorBody, + t: TFunction, + locale?: string, +): string { + const isCost = body.dimension === 'cost'; + const amount = (value?: number) => + isCost + ? new Intl.NumberFormat(locale, { + style: 'currency', + currency: 'USD', + }).format(value ?? 0) + : new Intl.NumberFormat(locale).format(value ?? 0); + const reset = body.resets_at ? new Date(body.resets_at) : null; + const resetsAt = + reset && !Number.isNaN(reset.getTime()) + ? new Intl.DateTimeFormat(locale, { + dateStyle: 'medium', + timeStyle: 'short', + }).format(reset) + : ''; + return t( + isCost + ? 'conversation.quotaExceeded.cost' + : 'conversation.quotaExceeded.tokens', + { used: amount(body.usage), limit: amount(body.limit), resetsAt }, + ); +} diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json index ae7d2630..9c6879f7 100644 --- a/frontend/src/locale/de.json +++ b/frontend/src/locale/de.json @@ -375,6 +375,13 @@ "toolCalls": "Werkzeugaufrufe", "runSuccess": "Erfolgsquote", "feedback": "Feedback" + }, + "quota": { + "title": "Ihr Nutzungskontingent", + "resets": "Wird zurückgesetzt: {{resetsAt}}", + "tokens": "Tokens", + "cost": "Kosten", + "usedOf": "{{used}} von {{limit}}" } }, "logs": { @@ -1262,6 +1269,10 @@ "running": "Läuft…", "denied": "Vom Benutzer abgelehnt", "failed": "fehlgeschlagen" + }, + "quotaExceeded": { + "tokens": "Sie haben {{used}} von Ihrem Kontingent von {{limit}} Tokens verbraucht. Es wird am {{resetsAt}} zurückgesetzt.", + "cost": "Sie haben {{used}} von Ihrem Nutzungsbudget von {{limit}} verbraucht. Es wird am {{resetsAt}} zurückgesetzt." } }, "agents": { diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index 93674148..ba269327 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -380,6 +380,13 @@ "toolCalls": "Tool Calls", "runSuccess": "Run Success", "feedback": "Feedback" + }, + "quota": { + "title": "Your usage quota", + "resets": "Resets {{resetsAt}}", + "tokens": "Tokens", + "cost": "Cost", + "usedOf": "{{used}} of {{limit}}" } }, "logs": { @@ -1273,6 +1280,10 @@ "running": "Running…", "denied": "Denied by user", "failed": "failed" + }, + "quotaExceeded": { + "tokens": "You've used {{used}} of your {{limit}} token quota. It resets {{resetsAt}}.", + "cost": "You've used {{used}} of your {{limit}} usage budget. It resets {{resetsAt}}." } }, "agents": { diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json index 3f19b13b..b844829b 100644 --- a/frontend/src/locale/es.json +++ b/frontend/src/locale/es.json @@ -375,6 +375,13 @@ "toolCalls": "Llamadas a Herramientas", "runSuccess": "Éxito de Ejecución", "feedback": "Retroalimentación" + }, + "quota": { + "title": "Tu cuota de uso", + "resets": "Se restablece el {{resetsAt}}", + "tokens": "Tokens", + "cost": "Coste", + "usedOf": "{{used}} de {{limit}}" } }, "logs": { @@ -1262,6 +1269,10 @@ "running": "Ejecutando…", "denied": "Denegado por el usuario", "failed": "falló" + }, + "quotaExceeded": { + "tokens": "Has usado {{used}} de tu cuota de {{limit}} tokens. Se restablece el {{resetsAt}}.", + "cost": "Has usado {{used}} de tu presupuesto de uso de {{limit}}. Se restablece el {{resetsAt}}." } }, "agents": { diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json index f60227d1..7c5aec0b 100644 --- a/frontend/src/locale/jp.json +++ b/frontend/src/locale/jp.json @@ -375,6 +375,13 @@ "toolCalls": "ツール呼び出し", "runSuccess": "実行成功率", "feedback": "フィードバック" + }, + "quota": { + "title": "利用クォータ", + "resets": "{{resetsAt}} にリセット", + "tokens": "トークン", + "cost": "コスト", + "usedOf": "{{used}} / {{limit}}" } }, "logs": { @@ -1262,6 +1269,10 @@ "running": "実行中…", "denied": "ユーザーによって拒否されました", "failed": "失敗" + }, + "quotaExceeded": { + "tokens": "トークンクォータ {{limit}} のうち {{used}} を使用しました。{{resetsAt}} にリセットされます。", + "cost": "利用予算 {{limit}} のうち {{used}} を使用しました。{{resetsAt}} にリセットされます。" } }, "agents": { diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json index 878451ba..69c0c25f 100644 --- a/frontend/src/locale/ru.json +++ b/frontend/src/locale/ru.json @@ -375,6 +375,13 @@ "toolCalls": "Вызовы инструментов", "runSuccess": "Успешность запусков", "feedback": "Обратная связь" + }, + "quota": { + "title": "Ваша квота использования", + "resets": "Сброс: {{resetsAt}}", + "tokens": "Токены", + "cost": "Стоимость", + "usedOf": "{{used}} из {{limit}}" } }, "logs": { @@ -1282,6 +1289,10 @@ "running": "Выполняется…", "denied": "Отклонено пользователем", "failed": "не удалось" + }, + "quotaExceeded": { + "tokens": "Вы использовали {{used}} из квоты в {{limit}} токенов. Квота сбросится {{resetsAt}}.", + "cost": "Вы использовали {{used}} из бюджета в {{limit}}. Бюджет сбросится {{resetsAt}}." } }, "agents": { diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json index 5323803d..4de54e9d 100644 --- a/frontend/src/locale/zh-TW.json +++ b/frontend/src/locale/zh-TW.json @@ -375,6 +375,13 @@ "toolCalls": "工具呼叫", "runSuccess": "執行成功率", "feedback": "回饋" + }, + "quota": { + "title": "您的用量配額", + "resets": "{{resetsAt}} 重設", + "tokens": "權杖", + "cost": "費用", + "usedOf": "{{used}} / {{limit}}" } }, "logs": { @@ -1262,6 +1269,10 @@ "running": "執行中…", "denied": "已被使用者拒絕", "failed": "失敗" + }, + "quotaExceeded": { + "tokens": "您已使用 {{limit}} 權杖配額中的 {{used}}。配額將於 {{resetsAt}} 重設。", + "cost": "您已使用 {{limit}} 用量預算中的 {{used}}。預算將於 {{resetsAt}} 重設。" } }, "agents": { diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json index 93ba0088..6a542506 100644 --- a/frontend/src/locale/zh.json +++ b/frontend/src/locale/zh.json @@ -375,6 +375,13 @@ "toolCalls": "工具调用", "runSuccess": "运行成功率", "feedback": "反馈" + }, + "quota": { + "title": "您的用量配额", + "resets": "{{resetsAt}} 重置", + "tokens": "令牌", + "cost": "费用", + "usedOf": "{{used}} / {{limit}}" } }, "logs": { @@ -1262,6 +1269,10 @@ "running": "正在运行…", "denied": "已被用户拒绝", "failed": "失败" + }, + "quotaExceeded": { + "tokens": "您已使用 {{limit}} 令牌配额中的 {{used}}。配额将于 {{resetsAt}} 重置。", + "cost": "您已使用 {{limit}} 用量预算中的 {{used}}。预算将于 {{resetsAt}} 重置。" } }, "agents": { diff --git a/frontend/src/settings/Analytics.tsx b/frontend/src/settings/Analytics.tsx index 1183f2af..4eaccd5c 100644 --- a/frontend/src/settings/Analytics.tsx +++ b/frontend/src/settings/Analytics.tsx @@ -26,6 +26,7 @@ import { useDarkTheme, useLoaderState } from '../hooks'; import { selectToken } from '../preferences/preferenceSlice'; import { htmlLegendPlugin } from '../utils/chartUtils'; import { formatDate } from '../utils/dateTimeUtils'; +import UsageQuota from './components/UsageQuota'; /** * Resolve a CSS custom property on `:root` to a concrete color string. @@ -377,6 +378,7 @@ export default function Analytics({ agentId }: AnalyticsProps) { return (
+ {agentId ? null : }

{t('settings.analytics.subtitle')} @@ -412,7 +414,7 @@ export default function Analytics({ agentId }: AnalyticsProps) {

{card.label}

diff --git a/frontend/src/settings/components/UsageQuota.tsx b/frontend/src/settings/components/UsageQuota.tsx new file mode 100644 index 00000000..7cb9ceff --- /dev/null +++ b/frontend/src/settings/components/UsageQuota.tsx @@ -0,0 +1,130 @@ +import { useEffect, useState } from 'react'; +import { useTranslation } from 'react-i18next'; +import { useSelector } from 'react-redux'; + +import userService from '../../api/services/userService'; +import { selectToken } from '../../preferences/preferenceSlice'; + +type Budget = { limit: number | null; used: number }; +type Bucket = { + bucket: string; + tokens: Budget; + cost: Budget; + resets_at: string; +}; + +function Meter({ + label, + budget, + format, +}: { + label: string; + budget: Budget; + format: (value: number) => string; +}) { + const { t } = useTranslation(); + if (budget.limit === null) return null; + const percent = + budget.limit <= 0 + ? 100 + : Math.min(100, Math.max(0, (budget.used / budget.limit) * 100)); + const tone = + percent >= 100 + ? 'bg-red-500' + : percent >= 80 + ? 'bg-amber-500' + : 'bg-[#7D54D1]'; + return ( +

+
+ {label} + + {t('settings.analytics.quota.usedOf', { + used: format(budget.used), + limit: format(budget.limit), + })} + +
+
+
+
+
+ ); +} + +/** The caller's usage against the quota an admin set; renders nothing when unlimited. */ +export default function UsageQuota() { + const { t, i18n } = useTranslation(); + const token = useSelector(selectToken); + const [bucket, setBucket] = useState(null); + + useEffect(() => { + let cancelled = false; + userService + .getQuota(token) + .then((res: Response) => (res.ok ? res.json() : null)) + .then((json: { buckets?: Bucket[] } | null) => { + if (cancelled) return; + const buckets = json?.buckets ?? []; + setBucket( + buckets.find((b) => b.bucket === 'all') ?? buckets[0] ?? null, + ); + }) + .catch(() => undefined); + return () => { + cancelled = true; + }; + }, [token]); + + if (!bucket) return null; + + const number = new Intl.NumberFormat(i18n.language); + const usd = new Intl.NumberFormat(i18n.language, { + style: 'currency', + currency: 'USD', + }); + const reset = new Date(bucket.resets_at); + const resetsAt = Number.isNaN(reset.getTime()) + ? '' + : new Intl.DateTimeFormat(i18n.language, { + dateStyle: 'medium', + timeStyle: 'short', + }).format(reset); + + return ( +
+
+

+ {t('settings.analytics.quota.title')} +

+ {resetsAt ? ( +

+ {t('settings.analytics.quota.resets', { resetsAt })} +

+ ) : null} +
+
+ number.format(value)} + /> + usd.format(value)} + /> +
+
+ ); +} From 1eacfdd3d00fa514ff39aa30fe4c2b5582229e8a Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 12:03:02 +0100 Subject: [PATCH 123/130] docs: usage quotas How the instance default, team allowances and user overrides resolve (including users in several teams), the quota window, who is charged for agent traffic, how cost budgets price models and what happens to unpriced ones, and the admin and user API. --- .env-template | 6 ++ docs/content/Deploying/Access-Control.mdx | 4 +- docs/content/Deploying/Usage-Quotas.mdx | 109 ++++++++++++++++++++++ docs/content/Deploying/_meta.js | 4 + 4 files changed, 122 insertions(+), 1 deletion(-) create mode 100644 docs/content/Deploying/Usage-Quotas.mdx diff --git a/.env-template b/.env-template index 5e868f5e..7f3c05a6 100644 --- a/.env-template +++ b/.env-template @@ -109,3 +109,9 @@ MICROSOFT_AUTHORITY=https://{tenantId}.ciamlogin.com/{tenantId} # PAT_MAX_LIFETIME_DAYS=365 # PAT_ALLOW_NON_EXPIRING=false # PAT_MAX_PER_USER=25 + +# Usage quotas (set limits in Admin → Quotas). Usage is counted per calendar +# day, week or month in UTC. Models without a declared price are recorded at $0 +# unless a fallback [input, output] USD rate per 1M tokens is given. +# QUOTA_PERIOD=month +# QUOTA_UNPRICED_RATE_PER_MILLION=[0.5, 1.5] diff --git a/docs/content/Deploying/Access-Control.mdx b/docs/content/Deploying/Access-Control.mdx index 1c452004..dc238349 100644 --- a/docs/content/Deploying/Access-Control.mdx +++ b/docs/content/Deploying/Access-Control.mdx @@ -88,6 +88,7 @@ Admins get a dashboard backed by a REST surface under `/api/admin` (every endpoi | `GET` | `/api/admin/audit` | Authentication/admin audit feed. | | `GET` | `/api/admin/devices/audit` | Remote-device audit feed. | | `GET` | `/api/admin/teams` | Instance-wide oversight of all teams. | +| `GET` `PUT` `DELETE` | `/api/admin/quotas/...` | [Usage quotas](/Deploying/Usage-Quotas) for the instance, teams and users. | Deactivating a user via the dashboard works for any auth type, while OIDC deployments can also offboard through [SCIM](/Deploying/OIDC-SSO#scim-user-provisioning). Both revoke live sessions immediately. @@ -141,9 +142,10 @@ Sharing rules: ## Audit log -Access-control actions are appended to the `auth_events` table alongside the [authentication events](/Deploying/OIDC-SSO#login-auditing). This includes admin actions — `admin_user_activated` / `admin_user_deactivated`, `admin_sessions_revoked`, `role_granted` / `role_revoked` (with `metadata.source` = `manual` or `oidc_group`) — and team events (`team.create`, `team.member_add`, `team.member_role`, `team.member_remove`, `team.share`, `team.unshare`, `team.transfer_owner`, `team.delete`). The acting admin is recorded in the event metadata. +Access-control actions are appended to the `auth_events` table alongside the [authentication events](/Deploying/OIDC-SSO#login-auditing). This includes admin actions — `admin_user_activated` / `admin_user_deactivated`, `admin_sessions_revoked`, `role_granted` / `role_revoked` (with `metadata.source` = `manual` or `oidc_group`), `quota_policy_set` / `quota_policy_deleted` — and team events (`team.create`, `team.member_add`, `team.member_role`, `team.member_remove`, `team.share`, `team.unshare`, `team.transfer_owner`, `team.delete`). The acting admin is recorded in the event metadata. ## Related - [SSO with OIDC](/Deploying/OIDC-SSO) — sign-in, group allowlists, and the `auth_events` table. +- [Usage Quotas](/Deploying/Usage-Quotas) — token and cost limits per user and per team. - [App Configuration](/Deploying/DocsGPT-Settings) — the full settings reference. diff --git a/docs/content/Deploying/Usage-Quotas.mdx b/docs/content/Deploying/Usage-Quotas.mdx new file mode 100644 index 00000000..701b9c83 --- /dev/null +++ b/docs/content/Deploying/Usage-Quotas.mdx @@ -0,0 +1,109 @@ +--- +title: Usage Quotas +description: Cap how many tokens or dollars each user may spend per day, week or month, with an instance default, per-team allowances and per-user overrides. +--- + +import { Callout } from 'nextra/components' + +# Usage Quotas + +An instance admin can limit how much each user spends on language models. A quota has two independent budgets: + +- **Tokens** — prompt plus generated tokens. Works for every model, including local ones. +- **Cost (USD)** — tokens priced at the model's catalog rate. Only sees models that declare a price. + +Set either, both or neither. Quotas are managed from **Admin → Quotas**, or through the [API](#api). With no quota set, nothing is limited. + +## Layers + +Limits are set at three layers. For each budget, the first layer that says something wins: + +1. **User override** — one user's own limit. +2. **Team allowance** — what each member of a team gets. +3. **Instance default** — everyone else. + +At each layer a budget is either *not set* (defer to the next layer), a *limit*, or *unlimited*. A limit of `0` blocks the user. The two budgets resolve separately, so a user's token limit can come from their team while their cost limit comes from the instance default. + +### Teams + +A team allowance is **per member**, not a pool the team shares: if the allowance is 2M tokens, each member may use 2M. + +A user in several teams gets the **most generous** allowance among them, and allowances are never added together. Usage is always counted per user, whichever teams they belong to. To hold one person below their team's allowance, give them a user override. + + +Team membership can change without an instance admin — team admins, OIDC group sync and SCIM all add members — so joining a team can only raise a user's allowance to what you granted that team, never lower it. Only instance admins set allowances; team admins cannot. + + +## Windows and enforcement + +Usage is counted over a calendar window in UTC, chosen for the whole instance with [`QUOTA_PERIOD`](/Deploying/Settings-Reference#quotas): `day` (from 00:00), `week` (from Monday) or `month` (from the 1st, the default). Windows are worked out when a request arrives, so there is no reset job to run. + +The quota is checked **before** a request starts. The request that crosses a limit completes; the next one is refused with HTTP `429`: + +```json +{ + "success": false, + "error_code": "quota-exceeded", + "message": "Usage quota reached (1,000,000 of 1,000,000 tokens). It resets at 2026-10-01T00:00:00+00:00.", + "dimension": "tokens", + "unit": "tokens", + "usage": 1000000, + "limit": 1000000, + "bucket": "all", + "source": "instance", + "resets_at": "2026-10-01T00:00:00+00:00" +} +``` + +The response carries a `Retry-After` header. The check covers chat, the agent and OpenAI-compatible APIs, scheduled runs (recorded as `budget_exceeded`) and webhook runs. If the quota check itself fails, the request is allowed. + +Who is charged: + +| Traffic | Charged to | +| --- | --- | +| Chat without an agent | The user | +| A user's own agent, its API key, webhooks and schedules | The agent's owner | +| An agent shared with the user | The user | + +Per-agent token and request limits still apply on top of the owner's quota. + +Users with a quota see their usage and the reset time under **Settings → Analytics**. + +## Pricing + +Cost budgets use the rates in the [model catalog](/Models/cloud-providers), in USD per million tokens: + +```yaml +models: + - id: my-model + input_cost_per_million: 3.0 + output_cost_per_million: 15.0 + cached_input_cost_per_million: 0.3 # optional, prompt-cache reads + cache_write_cost_per_million: 3.75 # optional, prompt-cache writes +``` + +The built-in catalogs ship list prices for hosted models. Override or add rates by dropping a YAML with the same model `id` into `MODELS_CONFIG_DIR`. The cost of each call is stored with its usage row when the call is made, so later price changes do not rewrite history. + + +A model with no declared price is recorded at $0, so a cost budget cannot see it. The Quotas tab lists such models once they have been used. Either limit them with a token budget, declare their rates, or set [`QUOTA_UNPRICED_RATE_PER_MILLION`](/Deploying/Settings-Reference#quotas) to charge a fallback rate. Models a user adds with their own API key are always $0, but their tokens still count. + + +## API + +Every admin endpoint requires the admin role, and every change is written to the [audit log](/Deploying/Access-Control#audit-log) as `quota_policy_set` or `quota_policy_deleted`. + +| Method | Path | Description | +| --- | --- | --- | +| `GET` | `/api/admin/quotas` | All policies by layer, the current window, and used models without a price. | +| `PUT` `DELETE` | `/api/admin/quotas/instance` | The instance default. | +| `GET` `PUT` `DELETE` | `/api/admin/quotas/teams/` | A team's per-member allowance. | +| `GET` `PUT` `DELETE` | `/api/admin/quotas/users/` | A user's override. `GET` also returns the limits the user ends up with, the layer each came from, and their usage. | +| `GET` | `/api/user/quota` | The caller's own limits, usage and reset time. | + +A `PUT` body sets, per budget, a limit or the unlimited flag; leave both out to defer to the next layer: + +```json +{ "token_limit": 2000000, "cost_unlimited": true, "note": "Research team" } +``` + +`bucket` (default `all`) narrows a policy to `direct` traffic (chat without an agent) or `agent` traffic (anything that runs through an agent). A request must fit both its own bucket and `all`. The dashboard edits `all`. diff --git a/docs/content/Deploying/_meta.js b/docs/content/Deploying/_meta.js index 105b9efd..9afb460d 100644 --- a/docs/content/Deploying/_meta.js +++ b/docs/content/Deploying/_meta.js @@ -15,6 +15,10 @@ export default { "title": "👥 Access Control & Teams", "href": "/Deploying/Access-Control" }, + "Usage-Quotas": { + "title": "📊 Usage Quotas", + "href": "/Deploying/Usage-Quotas" + }, "Docker-Deploying": { "title": "🛳️ Docker Setup", "href": "/Deploying/Docker-Deploying" From 6db9014dfb146b37f65dcaafcedba0ce74ed8ce5 Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 12:14:14 +0100 Subject: [PATCH 124/130] fix(quotas): list unpriced models by recorded cost; integer token limits in status The unpriced-model notice asked the live registry whether a model has a price, so a priced model whose provider was later disabled showed up as unpriced. It now lists models whose calls this period were all recorded at $0. Token limits are serialized as integers. --- docsgpt/api/admin/quotas.py | 8 +++++--- docsgpt/quotas/service.py | 13 +++++++------ tests/api/test_quota_endpoints.py | 13 ++++++++----- tests/quotas/test_service.py | 2 +- 4 files changed, 21 insertions(+), 15 deletions(-) diff --git a/docsgpt/api/admin/quotas.py b/docsgpt/api/admin/quotas.py index 8210ba8d..58b49387 100644 --- a/docsgpt/api/admin/quotas.py +++ b/docsgpt/api/admin/quotas.py @@ -151,13 +151,15 @@ def _delete_policy(scope: str, subject_id: Optional[str]): def _unpriced_models(conn) -> list[dict]: - """Catalog models used this period that no cost limit can see.""" + """Models used this period whose calls were all recorded at $0 for want of a price.""" start, _ = window_bounds(settings.QUOTA_PERIOD) return [ row for row in TokenUsageRepository(conn).tokens_by_model(start=start) - # BYOM ids are UUIDs; those calls are $0 by design, not by omission. - if not looks_like_uuid(row["model_id"]) and not is_priced(row["model_id"]) + # Judged by what was recorded, so a priced model whose provider has since + # been disabled is not listed. BYOM ids are UUIDs and $0 by design; a + # model explicitly priced at $0 is free, not unpriced. + if row["cost"] == 0 and not looks_like_uuid(row["model_id"]) and not is_priced(row["model_id"]) ] diff --git a/docsgpt/quotas/service.py b/docsgpt/quotas/service.py index 7659fe1d..fff7178a 100644 --- a/docsgpt/quotas/service.py +++ b/docsgpt/quotas/service.py @@ -43,18 +43,19 @@ class BucketStatus: def to_dict(self) -> dict: """Return the JSON shape shared by the admin and user quota endpoints.""" - def budget(limit: ResolvedLimit, used: float) -> dict: + def budget(limit: Optional[float], resolved: ResolvedLimit, used: float) -> dict: return { - "limit": limit.limit, + "limit": limit, "used": used, - "source": limit.source, - "source_id": limit.source_id, + "source": resolved.source, + "source_id": resolved.source_id, } + tokens, cost = self.limits.tokens, self.limits.cost return { "bucket": self.bucket, - "tokens": budget(self.limits.tokens, self.tokens_used), - "cost": budget(self.limits.cost, round(self.cost_used, 6)), + "tokens": budget(None if tokens.unlimited else int(tokens.limit), tokens, self.tokens_used), + "cost": budget(cost.limit, cost, round(self.cost_used, 6)), "resets_at": self.resets_at.isoformat(), } diff --git a/tests/api/test_quota_endpoints.py b/tests/api/test_quota_endpoints.py index 44c424db..c7337d1b 100644 --- a/tests/api/test_quota_endpoints.py +++ b/tests/api/test_quota_endpoints.py @@ -222,7 +222,8 @@ class TestUserPolicy: body = _body(client.get("/api/admin/quotas/users/u1")) overall = body["effective"][0] assert overall["bucket"] == "all" - assert overall["tokens"] == {"limit": 900.0, "used": 40, "source": "team", "source_id": big} + assert overall["tokens"] == {"limit": 900, "used": 40, "source": "team", "source_id": big} + assert isinstance(overall["tokens"]["limit"], int) assert overall["cost"] == {"limit": 1.0, "used": 0.25, "source": "instance", "source_id": None} assert body["policies"] == [] @@ -243,12 +244,14 @@ class TestUserPolicy: class TestUnpricedModels: - def test_lists_used_catalog_models_without_a_price(self, client, db): + def test_lists_models_recorded_at_zero_for_want_of_a_price(self, client, db): usage = TokenUsageRepository(db) usage.insert(user_id="u1", prompt_tokens=10, model_id="local-llama") - usage.insert(user_id="u1", prompt_tokens=5, model_id="claude-haiku-4-5", cost=0.1) + # Priced when called; its provider may be disabled by now. + usage.insert(user_id="u1", prompt_tokens=5, model_id="retired-priced-model", cost=0.1) + usage.insert(user_id="u1", prompt_tokens=3, model_id="free-model") usage.insert(user_id="u1", prompt_tokens=7, model_id="7d0c1a52-2f5e-4c53-9a0e-111111111111") - with _admin(), patch("docsgpt.api.admin.quotas.is_priced", lambda m: m == "claude-haiku-4-5"): + with _admin(), patch("docsgpt.api.admin.quotas.is_priced", lambda m: m == "free-model"): unpriced = _body(client.get("/api/admin/quotas"))["unpriced_models"] assert unpriced == [{"model_id": "local-llama", "tokens": 10, "cost": 0.0}] @@ -266,7 +269,7 @@ class TestMyQuota: body = _body(client.get("/api/user/quota")) (bucket,) = body["buckets"] assert bucket["bucket"] == "all" - assert bucket["tokens"] == {"limit": 100.0, "used": 30} + assert bucket["tokens"] == {"limit": 100, "used": 30} assert bucket["cost"] == {"limit": None, "used": 0.0} assert "source" not in json.dumps(body) and "secret" not in json.dumps(body) diff --git a/tests/quotas/test_service.py b/tests/quotas/test_service.py index 0a365799..13975bb3 100644 --- a/tests/quotas/test_service.py +++ b/tests/quotas/test_service.py @@ -210,7 +210,7 @@ class TestStatusAndPayload: (status,) = QuotaService.status("u1", now=NOW) assert status.to_dict() == { "bucket": "all", - "tokens": {"limit": 100.0, "used": 40, "source": "instance", "source_id": None}, + "tokens": {"limit": 100, "used": 40, "source": "instance", "source_id": None}, "cost": {"limit": 2.0, "used": 0.5, "source": "instance", "source_id": None}, "resets_at": "2026-10-01T00:00:00+00:00", } From 69f55b74cb0ae069ac5699002412f557ad726d38 Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 12:15:42 +0100 Subject: [PATCH 125/130] refactor(quotas): validate policy bodies without exception text in responses Validation problems are returned as values rather than raised and echoed with str(exc), and a huge integer limit is rejected as out of range instead of overflowing. Tests no longer call mutating endpoints inside asserts. --- docsgpt/api/admin/quotas.py | 111 +++++++++--------- tests/api/test_quota_endpoints.py | 12 +- .../db/repositories/test_quota_policies.py | 7 +- 3 files changed, 65 insertions(+), 65 deletions(-) diff --git a/docsgpt/api/admin/quotas.py b/docsgpt/api/admin/quotas.py index 58b49387..0235f36b 100644 --- a/docsgpt/api/admin/quotas.py +++ b/docsgpt/api/admin/quotas.py @@ -30,10 +30,8 @@ from docsgpt.storage.db.session import db_readonly, db_session _MAX_TOKEN_LIMIT = 2**62 _MAX_COST_LIMIT = 99_999_999.0 _MAX_NOTE_LENGTH = 500 - - -class _BadPolicy(ValueError): - """The request body does not describe a valid policy.""" +_FLAG_DEFAULTS = {"token_unlimited": False, "cost_unlimited": False, "enabled": True} +_BUCKET_MESSAGE = f"bucket must be one of: {', '.join(BUCKETS)}" def _policy_json(row: dict) -> dict: @@ -57,60 +55,59 @@ def _error(message: str, status: int): return make_response(jsonify({"success": False, "message": message}), status) -def _bucket(value: Any) -> str: +def _limit_error(value: Any, name: str, whole: bool, maximum: float) -> Optional[str]: + """Return why ``value`` is not a valid limit, or ``None``.""" if value is None: - return "all" - if value not in BUCKETS: - raise _BadPolicy(f"bucket must be one of: {', '.join(BUCKETS)}") - return value + return None + number = (int,) if whole else (int, float) + if isinstance(value, bool) or not isinstance(value, number): + return f"{name} must be a {'whole number' if whole else 'number'} or null" + # Range first: ``isfinite`` overflows on an int too large for a float. + if not 0 <= value <= maximum or not math.isfinite(value): + return f"{name} is out of range" + return None -def _flag(data: dict, key: str, default: bool) -> bool: - value = data.get(key, default) - if not isinstance(value, bool): - raise _BadPolicy(f"{key} must be a boolean") - return value +def _parse_policy(data: Any) -> tuple[Optional[dict], Optional[str]]: + """Validate a policy body. - -def _parse_policy(data: Any) -> dict: - """Validate a policy body into ``QuotaPoliciesRepository.upsert`` kwargs.""" + Returns: + ``(fields, None)`` with ``QuotaPoliciesRepository.upsert`` kwargs, or + ``(None, message)`` describing the first problem. + """ if not isinstance(data, dict): - raise _BadPolicy("Body must be a JSON object") - token_limit = data.get("token_limit") - if token_limit is not None: - if isinstance(token_limit, bool) or not isinstance(token_limit, int): - raise _BadPolicy("token_limit must be a whole number or null") - if not 0 <= token_limit <= _MAX_TOKEN_LIMIT: - raise _BadPolicy("token_limit is out of range") - cost_limit = data.get("cost_limit_usd") - if cost_limit is not None: - if isinstance(cost_limit, bool) or not isinstance(cost_limit, (int, float)): - raise _BadPolicy("cost_limit_usd must be a number or null") - if not math.isfinite(cost_limit) or not 0 <= cost_limit <= _MAX_COST_LIMIT: - raise _BadPolicy("cost_limit_usd is out of range") - cost_limit = round(float(cost_limit), 4) - token_unlimited = _flag(data, "token_unlimited", False) - cost_unlimited = _flag(data, "cost_unlimited", False) - if token_unlimited and token_limit is not None: - raise _BadPolicy("Set token_limit or token_unlimited, not both") - if cost_unlimited and cost_limit is not None: - raise _BadPolicy("Set cost_limit_usd or cost_unlimited, not both") - if token_limit is None and cost_limit is None and not token_unlimited and not cost_unlimited: - raise _BadPolicy("Set a limit or mark a budget unlimited; delete the policy to remove it") + return None, "Body must be a JSON object" + token_limit, cost_limit = data.get("token_limit"), data.get("cost_limit_usd") + problem = _limit_error(token_limit, "token_limit", True, _MAX_TOKEN_LIMIT) or _limit_error( + cost_limit, "cost_limit_usd", False, _MAX_COST_LIMIT + ) + if problem: + return None, problem + flags = {key: data.get(key, default) for key, default in _FLAG_DEFAULTS.items()} + for key, value in flags.items(): + if not isinstance(value, bool): + return None, f"{key} must be a boolean" + bucket = data.get("bucket", "all") + if bucket not in BUCKETS: + return None, _BUCKET_MESSAGE note = data.get("note") - if note is not None: - if not isinstance(note, str): - raise _BadPolicy("note must be a string") - note = note.strip()[:_MAX_NOTE_LENGTH] or None + if note is not None and not isinstance(note, str): + return None, "note must be a string" + if flags["token_unlimited"] and token_limit is not None: + return None, "Set token_limit or token_unlimited, not both" + if flags["cost_unlimited"] and cost_limit is not None: + return None, "Set cost_limit_usd or cost_unlimited, not both" + if token_limit is None and cost_limit is None and not flags["token_unlimited"] and not flags["cost_unlimited"]: + return None, "Set a limit or mark a budget unlimited; delete the policy to remove it" return { - "bucket": _bucket(data.get("bucket")), + "bucket": bucket, "token_limit": token_limit, - "token_unlimited": token_unlimited, - "cost_limit_usd": cost_limit, - "cost_unlimited": cost_unlimited, - "enabled": _flag(data, "enabled", True), - "note": note, - } + "token_unlimited": flags["token_unlimited"], + "cost_limit_usd": round(float(cost_limit), 4) if cost_limit is not None else None, + "cost_unlimited": flags["cost_unlimited"], + "enabled": flags["enabled"], + "note": (note.strip()[:_MAX_NOTE_LENGTH] or None) if note else None, + }, None def _audit(conn, event: str, scope: str, subject_id: Optional[str], detail: dict) -> None: @@ -126,10 +123,9 @@ def _audit(conn, event: str, scope: str, subject_id: Optional[str], detail: dict def _put_policy(scope: str, subject_id: Optional[str]): - try: - fields = _parse_policy(request.get_json(silent=True)) - except _BadPolicy as exc: - return _error(str(exc), 400) + fields, problem = _parse_policy(request.get_json(silent=True)) + if fields is None: + return _error(problem or "Invalid policy", 400) with db_session() as conn: row = QuotaPoliciesRepository(conn).upsert( scope=scope, subject_id=subject_id, actor=_actor(), **fields @@ -139,10 +135,9 @@ def _put_policy(scope: str, subject_id: Optional[str]): def _delete_policy(scope: str, subject_id: Optional[str]): - try: - bucket = _bucket(request.args.get("bucket")) if "bucket" in request.args else None - except _BadPolicy as exc: - return _error(str(exc), 400) + bucket = request.args.get("bucket") + if bucket is not None and bucket not in BUCKETS: + return _error(_BUCKET_MESSAGE, 400) with db_session() as conn: deleted = QuotaPoliciesRepository(conn).delete(scope, subject_id, bucket) if deleted: diff --git a/tests/api/test_quota_endpoints.py b/tests/api/test_quota_endpoints.py index c7337d1b..69224f9c 100644 --- a/tests/api/test_quota_endpoints.py +++ b/tests/api/test_quota_endpoints.py @@ -125,9 +125,10 @@ class TestInstancePolicy: ] assert overview["period"] == "month" - assert _body(client.delete("/api/admin/quotas/instance?bucket=agent"))["deleted"] == 1 - assert _body(client.delete("/api/admin/quotas/instance"))["deleted"] == 1 - assert _body(client.get("/api/admin/quotas"))["instance"] == [] + one_bucket = _body(client.delete("/api/admin/quotas/instance?bucket=agent")) + the_rest = _body(client.delete("/api/admin/quotas/instance")) + remaining = _body(client.get("/api/admin/quotas"))["instance"] + assert (one_bucket["deleted"], the_rest["deleted"], remaining) == (1, 1, []) def test_writes_are_audited(self, client, db): with _admin(): @@ -158,6 +159,8 @@ class TestInstancePolicy: {"token_limit": True}, {"token_limit": "10"}, {"token_limit": 2**63}, + {"token_limit": 10**400}, + {"cost_limit_usd": 10**400}, {"cost_limit_usd": -0.01}, {"cost_limit_usd": "5"}, {"cost_limit_usd": float("inf")}, @@ -178,7 +181,8 @@ class TestInstancePolicy: def test_unknown_bucket_on_delete(self, client, db): with _admin(): - assert client.delete("/api/admin/quotas/instance?bucket=nope").status_code == 400 + resp = client.delete("/api/admin/quotas/instance?bucket=nope") + assert resp.status_code == 400 def test_zero_is_accepted_as_a_block(self, client, db): with _admin(): diff --git a/tests/storage/db/repositories/test_quota_policies.py b/tests/storage/db/repositories/test_quota_policies.py index 25396210..83a9302c 100644 --- a/tests/storage/db/repositories/test_quota_policies.py +++ b/tests/storage/db/repositories/test_quota_policies.py @@ -137,7 +137,8 @@ class TestDelete: repo.upsert(scope="user", subject_id="u1", token_limit=1) repo.upsert(scope="user", subject_id="u1", bucket="agent", token_limit=2) repo.upsert(scope="user", subject_id="u2", token_limit=3) - assert repo.delete("user", "u1", "agent") == 1 - assert repo.delete("user", "u1", "agent") == 0 - assert repo.delete("user", "u1") == 1 + first = repo.delete("user", "u1", "agent") + again = repo.delete("user", "u1", "agent") + rest = repo.delete("user", "u1") + assert (first, again, rest) == (1, 0, 1) assert repo.get("user", "u2") is not None From 4bd259fc091932a6ee0c017bb3595111089b5ea6 Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 12:44:26 +0100 Subject: [PATCH 126/130] fix(quotas): address review: resume claims, agent bucket rule, UI races - A tool continuation refused for usage now releases the resume claim it took; before, retries got a 409 until the stale claim was reverted. - Agent traffic is any row with an agent key or an agent id, so keyless agents and workflow nodes count toward the agent bucket, not direct. - The user quota modal discards responses for a previously opened user. - The usage meter shows every limited bucket, not only 'all'. - Restore the class separator on the analytics stat card that a formatter run removed, and align the OpenRouter DeepSeek description with its rates. --- docs/content/Deploying/Usage-Quotas.mdx | 2 +- docsgpt/agents/headless_runner.py | 4 +- docsgpt/api/answer/routes/answer.py | 4 +- docsgpt/api/answer/routes/base.py | 26 +++++++ docsgpt/api/answer/routes/stream.py | 6 +- docsgpt/api/v1/routes.py | 8 ++- docsgpt/core/models/openrouter.yaml | 2 +- .../storage/db/repositories/token_usage.py | 9 +-- frontend/src/admin/UserQuotaModal.tsx | 15 +++-- frontend/src/locale/de.json | 6 +- frontend/src/locale/en.json | 6 +- frontend/src/locale/es.json | 6 +- frontend/src/locale/jp.json | 6 +- frontend/src/locale/ru.json | 6 +- frontend/src/locale/zh-TW.json | 6 +- frontend/src/locale/zh.json | 6 +- frontend/src/settings/Analytics.tsx | 2 +- .../src/settings/components/UsageQuota.tsx | 50 ++++++++------ tests/quotas/test_enforcement.py | 67 +++++++++++++++++++ .../db/repositories/test_token_usage.py | 10 ++- 20 files changed, 202 insertions(+), 45 deletions(-) diff --git a/docs/content/Deploying/Usage-Quotas.mdx b/docs/content/Deploying/Usage-Quotas.mdx index 701b9c83..c7391625 100644 --- a/docs/content/Deploying/Usage-Quotas.mdx +++ b/docs/content/Deploying/Usage-Quotas.mdx @@ -106,4 +106,4 @@ A `PUT` body sets, per budget, a limit or the unlimited flag; leave both out to { "token_limit": 2000000, "cost_unlimited": true, "note": "Research team" } ``` -`bucket` (default `all`) narrows a policy to `direct` traffic (chat without an agent) or `agent` traffic (anything that runs through an agent). A request must fit both its own bucket and `all`. The dashboard edits `all`. +`bucket` (default `all`) narrows a policy to `direct` traffic (chat without an agent) or `agent` traffic (anything that runs through an agent, whether or not the agent has an API key). A request must fit both its own bucket and `all`. The dashboard edits `all`. diff --git a/docsgpt/agents/headless_runner.py b/docsgpt/agents/headless_runner.py index 344e5fe3..c2a4d713 100644 --- a/docsgpt/agents/headless_runner.py +++ b/docsgpt/agents/headless_runner.py @@ -87,7 +87,9 @@ def run_agent_headless( if not owner: raise ValueError("Agent config is missing user_id; cannot run headless.") decoded_token = {"sub": owner} - exceeded = QuotaService.check(owner, "agent" if agent_config.get("key") else "direct") + # An agent run is agent traffic whether or not the agent has a key yet. + is_agent_run = bool(agent_config.get("key") or _resolve_agent_id(agent_config)) + exceeded = QuotaService.check(owner, "agent" if is_agent_run else "direct") if exceeded is not None: raise QuotaExceededError(exceeded) diff --git a/docsgpt/api/answer/routes/answer.py b/docsgpt/api/answer/routes/answer.py index cdc63cca..b5819e3e 100644 --- a/docsgpt/api/answer/routes/answer.py +++ b/docsgpt/api/answer/routes/answer.py @@ -103,8 +103,8 @@ class AnswerResource(Resource, BaseAnswerResource): ) if not processor.decoded_token: return make_response({"error": "Unauthorized"}, 401) - if error := self.check_usage( - processor.agent_config, processor.decoded_token + if error := self.check_usage_on_resume( + processor, data["conversation_id"] ): return error stream = self.complete_stream( diff --git a/docsgpt/api/answer/routes/base.py b/docsgpt/api/answer/routes/base.py index c61cfd48..c2b3f929 100644 --- a/docsgpt/api/answer/routes/base.py +++ b/docsgpt/api/answer/routes/base.py @@ -211,6 +211,32 @@ class BaseAnswerResource: ) return None + def check_usage_on_resume(self, processor: Any, conversation_id: Any) -> Optional[Response]: + """Run ``check_usage`` for a tool continuation, releasing its claim on refusal. + + ``resume_from_tool_actions`` has already claimed the paused turn by the + time the limits can be checked (the agent config comes from the claimed + state). A refusal returns before ``complete_stream`` and its cleanup, so + the claim is released here; otherwise retries get a 409 until the stale + claim is reverted. + + Args: + processor: The ``StreamProcessor`` that resumed the turn. + conversation_id: The conversation whose pending state was claimed. + + Returns: + None, or the refusal Response. + """ + error = self.check_usage(processor.agent_config, processor.decoded_token) + if error is None or not conversation_id: + return error + user = processor.initial_user_id or (processor.decoded_token or {}).get("sub") + try: + ContinuationService().release_claim(str(conversation_id), user) + except Exception: + logger.exception("Failed to release resume claim after a usage refusal") + return error + def complete_stream( self, question: str, diff --git a/docsgpt/api/answer/routes/stream.py b/docsgpt/api/answer/routes/stream.py index fdc5a00b..54b62ec8 100644 --- a/docsgpt/api/answer/routes/stream.py +++ b/docsgpt/api/answer/routes/stream.py @@ -115,9 +115,9 @@ class StreamResource(Resource, BaseAnswerResource): status=401, mimetype="text/event-stream", ) - if error := self.check_usage( - processor.agent_config, processor.decoded_token - ): + if error := self.check_usage_on_resume( + processor, data["conversation_id"] + ): return error return Response( with_sse_keepalive( diff --git a/docsgpt/api/v1/routes.py b/docsgpt/api/v1/routes.py index cd208e84..a9902a4a 100644 --- a/docsgpt/api/v1/routes.py +++ b/docsgpt/api/v1/routes.py @@ -258,6 +258,8 @@ def chat_completions(): try: processor = StreamProcessor(internal_data, decoded_token) + # Set when this request took the resume claim, so a refusal can release it. + claimed_conversation_id = None if internal_data.get("tool_actions"): conversation_id = internal_data.get("conversation_id") @@ -282,6 +284,7 @@ def chat_completions(): claimed_state=pending_state, ) processor.conversation_id = conversation_id + claimed_conversation_id = conversation_id else: # Compatibility fallback for old/completed conversations and # clients that resend the full transcript without resumable @@ -338,7 +341,10 @@ def chat_completions(): ) helper = _V1AnswerHelper() - usage_error = helper.check_usage(processor.agent_config, processor.decoded_token) + if claimed_conversation_id: + usage_error = helper.check_usage_on_resume(processor, claimed_conversation_id) + else: + usage_error = helper.check_usage(processor.agent_config, processor.decoded_token) if usage_error: return usage_error diff --git a/docsgpt/core/models/openrouter.yaml b/docsgpt/core/models/openrouter.yaml index 6ac373ee..f0b2cffd 100644 --- a/docsgpt/core/models/openrouter.yaml +++ b/docsgpt/core/models/openrouter.yaml @@ -15,7 +15,7 @@ models: - id: deepseek/deepseek-v3.2 display_name: DeepSeek V3.2 - description: Open-weights reasoning model, very low cost (~$0.25 in / $0.38 out per 1M) + description: Open-weights reasoning model, very low cost (~$0.23 in / $0.34 out per 1M) context_window: 131072 attachments: [] supports_structured_output: true diff --git a/docsgpt/storage/db/repositories/token_usage.py b/docsgpt/storage/db/repositories/token_usage.py index 2c8038fa..fb256b55 100644 --- a/docsgpt/storage/db/repositories/token_usage.py +++ b/docsgpt/storage/db/repositories/token_usage.py @@ -144,16 +144,17 @@ class TokenUsageRepository: Args: user_id: The billable user (auth ``sub``). start: Inclusive window start. - bucket: ``all``, ``direct`` (rows without an agent key) or - ``agent`` (rows with one). + bucket: ``all``, ``agent`` (rows carrying an agent key or an agent + id) or ``direct`` (rows with neither). Rollup rows are excluded; side-channel calls count, they are real spend. """ clauses = ["user_id = :user_id", "timestamp >= :start", "source <> ALL(:rollup_sources)"] + # Keyless agents and workflow nodes carry an agent id without a key. if bucket == "direct": - clauses.append("api_key IS NULL") + clauses.append("api_key IS NULL AND agent_id IS NULL") elif bucket == "agent": - clauses.append("api_key IS NOT NULL") + clauses.append("(api_key IS NOT NULL OR agent_id IS NOT NULL)") elif bucket != "all": raise ValueError(f"unknown usage bucket: {bucket!r}") row = self._conn.execute( diff --git a/frontend/src/admin/UserQuotaModal.tsx b/frontend/src/admin/UserQuotaModal.tsx index f6ce58b2..241da310 100644 --- a/frontend/src/admin/UserQuotaModal.tsx +++ b/frontend/src/admin/UserQuotaModal.tsx @@ -1,4 +1,4 @@ -import { useCallback, useEffect, useState } from 'react'; +import { useCallback, useEffect, useRef, useState } from 'react'; import { useSelector } from 'react-redux'; import adminService from '../api/services/adminService'; @@ -26,15 +26,22 @@ export default function UserQuotaModal({ const [data, setData] = useState(null); const [teamNames, setTeamNames] = useState>({}); + // Bumped per request so a slow response for a previous user is discarded + // instead of showing (and letting the editor save) that user's policy. + const requestRef = useRef(0); + const load = useCallback(async () => { - if (!userId) return; + const request = ++requestRef.current; setData(null); + if (!userId) return; try { const [res, teamsJson] = await Promise.all([ adminService.getUserQuota(userId, token), teamsService.listAll(token).catch(() => ({})), ]); - setData(await res.json().catch(() => ({ success: false }))); + const json = await res.json().catch(() => ({ success: false })); + if (request !== requestRef.current) return; + setData(json); setTeamNames( Object.fromEntries( (teamsJson?.teams ?? []).map((team: any) => [ @@ -44,7 +51,7 @@ export default function UserQuotaModal({ ), ); } catch { - setData({ success: false }); + if (request === requestRef.current) setData({ success: false }); } }, [userId, token]); diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json index 9c6879f7..df6dabfa 100644 --- a/frontend/src/locale/de.json +++ b/frontend/src/locale/de.json @@ -381,7 +381,11 @@ "resets": "Wird zurückgesetzt: {{resetsAt}}", "tokens": "Tokens", "cost": "Kosten", - "usedOf": "{{used}} von {{limit}}" + "usedOf": "{{used}} von {{limit}}", + "scope": { + "direct": "Chat ohne Agent", + "agent": "Über Agenten" + } } }, "logs": { diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index ba269327..cdf03a26 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -386,7 +386,11 @@ "resets": "Resets {{resetsAt}}", "tokens": "Tokens", "cost": "Cost", - "usedOf": "{{used}} of {{limit}}" + "usedOf": "{{used}} of {{limit}}", + "scope": { + "direct": "Chat without an agent", + "agent": "Through agents" + } } }, "logs": { diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json index b844829b..76058628 100644 --- a/frontend/src/locale/es.json +++ b/frontend/src/locale/es.json @@ -381,7 +381,11 @@ "resets": "Se restablece el {{resetsAt}}", "tokens": "Tokens", "cost": "Coste", - "usedOf": "{{used}} de {{limit}}" + "usedOf": "{{used}} de {{limit}}", + "scope": { + "direct": "Chat sin agente", + "agent": "A través de agentes" + } } }, "logs": { diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json index 7c5aec0b..18626d19 100644 --- a/frontend/src/locale/jp.json +++ b/frontend/src/locale/jp.json @@ -381,7 +381,11 @@ "resets": "{{resetsAt}} にリセット", "tokens": "トークン", "cost": "コスト", - "usedOf": "{{used}} / {{limit}}" + "usedOf": "{{used}} / {{limit}}", + "scope": { + "direct": "エージェントなしのチャット", + "agent": "エージェント経由" + } } }, "logs": { diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json index 69c0c25f..081860c6 100644 --- a/frontend/src/locale/ru.json +++ b/frontend/src/locale/ru.json @@ -381,7 +381,11 @@ "resets": "Сброс: {{resetsAt}}", "tokens": "Токены", "cost": "Стоимость", - "usedOf": "{{used}} из {{limit}}" + "usedOf": "{{used}} из {{limit}}", + "scope": { + "direct": "Чат без агента", + "agent": "Через агентов" + } } }, "logs": { diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json index 4de54e9d..7902cf5a 100644 --- a/frontend/src/locale/zh-TW.json +++ b/frontend/src/locale/zh-TW.json @@ -381,7 +381,11 @@ "resets": "{{resetsAt}} 重設", "tokens": "權杖", "cost": "費用", - "usedOf": "{{used}} / {{limit}}" + "usedOf": "{{used}} / {{limit}}", + "scope": { + "direct": "不使用代理的聊天", + "agent": "透過代理" + } } }, "logs": { diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json index 6a542506..7e647c5e 100644 --- a/frontend/src/locale/zh.json +++ b/frontend/src/locale/zh.json @@ -381,7 +381,11 @@ "resets": "{{resetsAt}} 重置", "tokens": "令牌", "cost": "费用", - "usedOf": "{{used}} / {{limit}}" + "usedOf": "{{used}} / {{limit}}", + "scope": { + "direct": "不使用代理的聊天", + "agent": "通过代理" + } } }, "logs": { diff --git a/frontend/src/settings/Analytics.tsx b/frontend/src/settings/Analytics.tsx index 4eaccd5c..f67855bd 100644 --- a/frontend/src/settings/Analytics.tsx +++ b/frontend/src/settings/Analytics.tsx @@ -414,7 +414,7 @@ export default function Analytics({ agentId }: AnalyticsProps) {

{card.label}

diff --git a/frontend/src/settings/components/UsageQuota.tsx b/frontend/src/settings/components/UsageQuota.tsx index 7cb9ceff..30511081 100644 --- a/frontend/src/settings/components/UsageQuota.tsx +++ b/frontend/src/settings/components/UsageQuota.tsx @@ -66,7 +66,7 @@ function Meter({ export default function UsageQuota() { const { t, i18n } = useTranslation(); const token = useSelector(selectToken); - const [bucket, setBucket] = useState(null); + const [buckets, setBuckets] = useState([]); useEffect(() => { let cancelled = false; @@ -75,10 +75,7 @@ export default function UsageQuota() { .then((res: Response) => (res.ok ? res.json() : null)) .then((json: { buckets?: Bucket[] } | null) => { if (cancelled) return; - const buckets = json?.buckets ?? []; - setBucket( - buckets.find((b) => b.bucket === 'all') ?? buckets[0] ?? null, - ); + setBuckets(json?.buckets ?? []); }) .catch(() => undefined); return () => { @@ -86,14 +83,14 @@ export default function UsageQuota() { }; }, [token]); - if (!bucket) return null; + if (buckets.length === 0) return null; const number = new Intl.NumberFormat(i18n.language); const usd = new Intl.NumberFormat(i18n.language, { style: 'currency', currency: 'USD', }); - const reset = new Date(bucket.resets_at); + const reset = new Date(buckets[0].resets_at); const resetsAt = Number.isNaN(reset.getTime()) ? '' : new Intl.DateTimeFormat(i18n.language, { @@ -101,6 +98,12 @@ export default function UsageQuota() { timeStyle: 'short', }).format(reset); + // A request must fit its own bucket and ``all``, so each limited one is shown. + const scopeLabel = (name: string) => + name === 'direct' || name === 'agent' + ? t(`settings.analytics.quota.scope.${name}`) + : null; + return (

@@ -113,18 +116,27 @@ export default function UsageQuota() {

) : null}
-
- number.format(value)} - /> - usd.format(value)} - /> -
+ {buckets.map((bucket) => ( +
+ {scopeLabel(bucket.bucket) ? ( +

+ {scopeLabel(bucket.bucket)} +

+ ) : null} +
+ number.format(value)} + /> + usd.format(value)} + /> +
+
+ ))}
); } diff --git a/tests/quotas/test_enforcement.py b/tests/quotas/test_enforcement.py index b02dbae1..c6300b22 100644 --- a/tests/quotas/test_enforcement.py +++ b/tests/quotas/test_enforcement.py @@ -106,3 +106,70 @@ class TestHeadless: assert raised.value.exceeded.source == "instance" assert "10 of 10 tokens" in str(raised.value) + + + def test_a_keyless_agent_run_is_agent_traffic(self, db): + from docsgpt.agents.headless_runner import run_agent_headless + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="owner", bucket="agent", token_limit=0) + config = {"user_id": "owner", "id": "22222222-2222-2222-2222-222222222222"} + + with patch("docsgpt.agents.headless_runner.RetrieverCreator"): + with pytest.raises(QuotaExceededError) as raised: + run_agent_headless(config, "hello") + + assert raised.value.exceeded.bucket == "agent" + + +class TestResumeRefusal: + def _processor(self, user_id="u1"): + from types import SimpleNamespace + + return SimpleNamespace(agent_config={}, decoded_token={"sub": user_id}, initial_user_id=user_id) + + def test_a_refused_resume_releases_its_claim(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + response = BaseAnswerResource().check_usage_on_resume(self._processor(), "conv-1") + + assert response.status_code == 429 + service.return_value.release_claim.assert_called_once_with("conv-1", "u1") + + def test_an_admitted_resume_keeps_its_claim(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + response = BaseAnswerResource().check_usage_on_resume(self._processor(), "conv-1") + + assert response is None + service.return_value.release_claim.assert_not_called() + + def test_no_claim_means_nothing_to_release(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + response = BaseAnswerResource().check_usage_on_resume(self._processor(), None) + + assert response.status_code == 429 + service.return_value.release_claim.assert_not_called() + + def test_a_failed_release_still_returns_the_refusal(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + service.return_value.release_claim.side_effect = RuntimeError("db down") + response = BaseAnswerResource().check_usage_on_resume(self._processor(), "conv-1") + + assert response.status_code == 429 diff --git a/tests/storage/db/repositories/test_token_usage.py b/tests/storage/db/repositories/test_token_usage.py index 9698c610..6ce4cb96 100644 --- a/tests/storage/db/repositories/test_token_usage.py +++ b/tests/storage/db/repositories/test_token_usage.py @@ -61,6 +61,14 @@ class TestUsageTotals: repo.insert(user_id="u-tot", prompt_tokens=100, generated_tokens=10, cost=0.5) repo.insert(user_id="u-tot", api_key="k", prompt_tokens=20, generated_tokens=2, cost=0.25) repo.insert(user_id="u-tot", prompt_tokens=7, generated_tokens=0, cost=0.125, source="title") + # A keyless agent (or workflow node): an agent id without a key. + from docsgpt.storage.db.repositories.agents import AgentsRepository + + agent = AgentsRepository(repo._conn).create("u-tot", "keyless", "draft") + repo.insert( + user_id="u-tot", agent_id=str(agent["id"]), + prompt_tokens=4, generated_tokens=0, cost=0.0625, + ) repo.insert(user_id="u-tot", prompt_tokens=999, generated_tokens=0, source="schedule") repo.insert(user_id="u-other", prompt_tokens=999, generated_tokens=0, cost=9) repo.insert( @@ -70,7 +78,7 @@ class TestUsageTotals: @pytest.mark.parametrize( "bucket, expected", - [("all", (139, 0.875)), ("direct", (117, 0.625)), ("agent", (22, 0.25))], + [("all", (143, 0.9375)), ("direct", (117, 0.625)), ("agent", (26, 0.3125))], ) def test_totals_per_bucket(self, pg_conn, bucket, expected): repo = _repo(pg_conn) From 3a30f5cce9d57464d3c0bfb61ee1de2cab9a3d2d Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 13:36:22 +0100 Subject: [PATCH 127/130] feat(pricing): rates for the default DocsGPT model $0.15 input, $0.50 output and $0.03 cached input per 1M tokens, so cost budgets see usage of the default model instead of recording it at $0. --- docsgpt/core/models/docsgpt.yaml | 3 +++ tests/test_pricing.py | 11 ++++++++++- 2 files changed, 13 insertions(+), 1 deletion(-) diff --git a/docsgpt/core/models/docsgpt.yaml b/docsgpt/core/models/docsgpt.yaml index b65b4fbf..4f494434 100644 --- a/docsgpt/core/models/docsgpt.yaml +++ b/docsgpt/core/models/docsgpt.yaml @@ -7,3 +7,6 @@ models: supports_tools: true attachments: [image] context_window: 1048576 + input_cost_per_million: 0.15 + output_cost_per_million: 0.5 + cached_input_cost_per_million: 0.03 diff --git a/tests/test_pricing.py b/tests/test_pricing.py index 81aa324f..d8a6ca69 100644 --- a/tests/test_pricing.py +++ b/tests/test_pricing.py @@ -118,7 +118,7 @@ class TestCatalogFields: self._load(tmp_path, "provider: openai\nmodels:\n - id: m\n input_cost_per_million: -1\n") def test_hosted_builtin_models_are_priced(self): - hosted = {"anthropic", "deepseek", "google", "groq", "novita", "openai", "openrouter"} + hosted = {"anthropic", "deepseek", "docsgpt", "google", "groq", "novita", "openai", "openrouter"} catalogs = [ c for c in load_model_yamls([BUILTIN_MODELS_DIR]) if c.source_path.stem in hosted ] @@ -128,3 +128,12 @@ class TestCatalogFields: caps = model.capabilities assert caps.input_cost_per_million is not None, model.id assert caps.output_cost_per_million is not None, model.id + + def test_default_docsgpt_model_rates(self): + (model,) = [ + m for c in load_model_yamls([BUILTIN_MODELS_DIR]) for m in c.models if m.id == "docsgpt-local" + ] + caps = model.capabilities + assert (caps.input_cost_per_million, caps.output_cost_per_million) == (0.15, 0.5) + assert caps.cached_input_cost_per_million == 0.03 + assert caps.cache_write_cost_per_million is None From b5296df8a934e8fec7cd3f0b2aa4b341ec8eab17 Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 14:30:48 +0100 Subject: [PATCH 128/130] feat(pricing): cached-input rates for gpt-5.4-mini and gpt-5.4-nano Checked against OpenAI's pricing page: gpt-5.5 at $5 / $30 (cached $0.50) was already right. The mini and nano models declared no cached rate, so cached prompt tokens were billed at the full input rate. --- docsgpt/core/models/openai.yaml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/docsgpt/core/models/openai.yaml b/docsgpt/core/models/openai.yaml index 598e7b3c..a4b73a3a 100644 --- a/docsgpt/core/models/openai.yaml +++ b/docsgpt/core/models/openai.yaml @@ -12,6 +12,7 @@ models: context_window: 1050000 api_flavor: responses reasoning_effort: medium + # Short-context rates. Prompts over 272K tokens bill at $10 / $45 (cached $1). input_cost_per_million: 5.0 output_cost_per_million: 30.0 cached_input_cost_per_million: 0.5 @@ -20,8 +21,10 @@ models: description: Cost-efficient GPT-5.4-class model for high-volume coding, computer use, and subagent workloads input_cost_per_million: 0.75 output_cost_per_million: 4.5 + cached_input_cost_per_million: 0.075 - id: gpt-5.4-nano display_name: GPT-5.4 Nano description: Cheapest GPT-5.4-class model, optimized for simple high-volume tasks where speed and cost matter most input_cost_per_million: 0.2 output_cost_per_million: 1.25 + cached_input_cost_per_million: 0.02 From 1e14605ee7c18636ba46e5f7bd158b1ef0d91398 Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 15:26:29 +0100 Subject: [PATCH 129/130] fix(quotas): keyless agent chat bucket, keep disabled policies disabled, cached rates - check_usage treats a request through a keyless (draft) agent as agent traffic, matching how its usage rows are bucketed and the headless rule - dashboard edits carry the stored enabled flag instead of re-enabling the policy; disabled policies are labelled in the Quotas tab and the editor - quota 429s send x-should-retry: false so OpenAI SDK clients do not retry a refusal that cannot succeed before the reset - cached-input and cache-write rates for Anthropic, OpenRouter and Groq gpt-oss-120b; refresh OpenRouter deepseek-v3.2 list prices - UsageQuota reuses usagePercent; docs note that a user override needs an existing user --- docs/content/Deploying/Usage-Quotas.mdx | 2 +- docsgpt/api/answer/routes/answer.py | 4 ++- docsgpt/api/answer/routes/base.py | 13 +++++++--- docsgpt/api/answer/routes/stream.py | 4 ++- docsgpt/api/v1/routes.py | 4 ++- docsgpt/core/models/anthropic.yaml | 6 +++++ docsgpt/core/models/groq.yaml | 1 + docsgpt/core/models/openrouter.yaml | 9 ++++--- docsgpt/quotas/http.py | 2 ++ frontend/src/admin/QuotaEditor.tsx | 7 +++++- frontend/src/admin/Quotas.tsx | 10 +++++--- frontend/src/admin/quotaUtils.test.ts | 8 ++++++ frontend/src/admin/quotaUtils.ts | 8 +++++- .../src/settings/components/UsageQuota.tsx | 6 ++--- tests/quotas/test_enforcement.py | 25 ++++++++++++++++--- 15 files changed, 87 insertions(+), 22 deletions(-) diff --git a/docs/content/Deploying/Usage-Quotas.mdx b/docs/content/Deploying/Usage-Quotas.mdx index c7391625..e4ffbaf1 100644 --- a/docs/content/Deploying/Usage-Quotas.mdx +++ b/docs/content/Deploying/Usage-Quotas.mdx @@ -97,7 +97,7 @@ Every admin endpoint requires the admin role, and every change is written to the | `GET` | `/api/admin/quotas` | All policies by layer, the current window, and used models without a price. | | `PUT` `DELETE` | `/api/admin/quotas/instance` | The instance default. | | `GET` `PUT` `DELETE` | `/api/admin/quotas/teams/` | A team's per-member allowance. | -| `GET` `PUT` `DELETE` | `/api/admin/quotas/users/` | A user's override. `GET` also returns the limits the user ends up with, the layer each came from, and their usage. | +| `GET` `PUT` `DELETE` | `/api/admin/quotas/users/` | A user's override; the user must already exist (SCIM-provisioned, or signed in once), otherwise `404`. `GET` also returns the limits the user ends up with, the layer each came from, and their usage. | | `GET` | `/api/user/quota` | The caller's own limits, usage and reset time. | A `PUT` body sets, per budget, a limit or the unlimited flag; leave both out to defer to the next layer: diff --git a/docsgpt/api/answer/routes/answer.py b/docsgpt/api/answer/routes/answer.py index b5819e3e..7cdcd722 100644 --- a/docsgpt/api/answer/routes/answer.py +++ b/docsgpt/api/answer/routes/answer.py @@ -132,7 +132,9 @@ class AnswerResource(Resource, BaseAnswerResource): return make_response({"error": "Unauthorized"}, 401) if error := self.check_usage( - processor.agent_config, processor.decoded_token + processor.agent_config, + processor.decoded_token, + agent_id=processor.agent_id, ): return error diff --git a/docsgpt/api/answer/routes/base.py b/docsgpt/api/answer/routes/base.py index c2b3f929..e0cfe4cd 100644 --- a/docsgpt/api/answer/routes/base.py +++ b/docsgpt/api/answer/routes/base.py @@ -115,7 +115,10 @@ class BaseAnswerResource: return prepared def check_usage( - self, agent_config: Dict, decoded_token: Optional[Dict] = None + self, + agent_config: Dict, + decoded_token: Optional[Dict] = None, + agent_id: Optional[str] = None, ) -> Optional[Response]: """Refuse the request when a usage limit is exhausted. @@ -127,6 +130,8 @@ class BaseAnswerResource: agent_config: The config dict of agent instance decoded_token: The request's resolved identity; its ``sub`` is the billable user. + agent_id: The agent the request runs through. A draft agent has no + key, but its usage rows carry the agent id, so it is agent traffic. Returns: None or Response if either of limits exceeded. @@ -134,7 +139,7 @@ class BaseAnswerResource: """ api_key = agent_config.get("user_api_key") user_id = (decoded_token or {}).get("sub") or agent_config.get("user_id") - exceeded = QuotaService.check(user_id, "agent" if api_key else "direct") + exceeded = QuotaService.check(user_id, "agent" if api_key or agent_id else "direct") if exceeded is not None: return quota_exceeded_response(exceeded) if not api_key: @@ -227,7 +232,9 @@ class BaseAnswerResource: Returns: None, or the refusal Response. """ - error = self.check_usage(processor.agent_config, processor.decoded_token) + error = self.check_usage( + processor.agent_config, processor.decoded_token, agent_id=processor.agent_id + ) if error is None or not conversation_id: return error user = processor.initial_user_id or (processor.decoded_token or {}).get("sub") diff --git a/docsgpt/api/answer/routes/stream.py b/docsgpt/api/answer/routes/stream.py index 54b62ec8..cdee274b 100644 --- a/docsgpt/api/answer/routes/stream.py +++ b/docsgpt/api/answer/routes/stream.py @@ -154,7 +154,9 @@ class StreamResource(Resource, BaseAnswerResource): ) if error := self.check_usage( - processor.agent_config, processor.decoded_token + processor.agent_config, + processor.decoded_token, + agent_id=processor.agent_id, ): return error should_persist, visibility = resolve_persistence( diff --git a/docsgpt/api/v1/routes.py b/docsgpt/api/v1/routes.py index a9902a4a..fa915d24 100644 --- a/docsgpt/api/v1/routes.py +++ b/docsgpt/api/v1/routes.py @@ -344,7 +344,9 @@ def chat_completions(): if claimed_conversation_id: usage_error = helper.check_usage_on_resume(processor, claimed_conversation_id) else: - usage_error = helper.check_usage(processor.agent_config, processor.decoded_token) + usage_error = helper.check_usage( + processor.agent_config, processor.decoded_token, agent_id=processor.agent_id + ) if usage_error: return usage_error diff --git a/docsgpt/core/models/anthropic.yaml b/docsgpt/core/models/anthropic.yaml index 784ab8e3..34a2cedc 100644 --- a/docsgpt/core/models/anthropic.yaml +++ b/docsgpt/core/models/anthropic.yaml @@ -12,6 +12,8 @@ models: supports_structured_output: true input_cost_per_million: 5.0 output_cost_per_million: 25.0 + cached_input_cost_per_million: 0.5 + cache_write_cost_per_million: 6.25 - id: claude-sonnet-4-6 display_name: Claude Sonnet 4.6 @@ -20,6 +22,8 @@ models: supports_structured_output: true input_cost_per_million: 3.0 output_cost_per_million: 15.0 + cached_input_cost_per_million: 0.3 + cache_write_cost_per_million: 3.75 - id: claude-haiku-4-5 display_name: Claude Haiku 4.5 @@ -27,3 +31,5 @@ models: supports_structured_output: true input_cost_per_million: 1.0 output_cost_per_million: 5.0 + cached_input_cost_per_million: 0.1 + cache_write_cost_per_million: 1.25 diff --git a/docsgpt/core/models/groq.yaml b/docsgpt/core/models/groq.yaml index c6e28d7a..a4d7edfd 100644 --- a/docsgpt/core/models/groq.yaml +++ b/docsgpt/core/models/groq.yaml @@ -10,6 +10,7 @@ models: supports_structured_output: true input_cost_per_million: 0.15 output_cost_per_million: 0.6 + cached_input_cost_per_million: 0.075 - id: llama-3.3-70b-versatile display_name: Llama 3.3 70B Versatile description: Meta's Llama 3.3 70B for general-purpose chat with parallel tool use diff --git a/docsgpt/core/models/openrouter.yaml b/docsgpt/core/models/openrouter.yaml index f0b2cffd..2957fa98 100644 --- a/docsgpt/core/models/openrouter.yaml +++ b/docsgpt/core/models/openrouter.yaml @@ -15,12 +15,13 @@ models: - id: deepseek/deepseek-v3.2 display_name: DeepSeek V3.2 - description: Open-weights reasoning model, very low cost (~$0.23 in / $0.34 out per 1M) + description: Open-weights reasoning model, very low cost (~$0.27 in / $0.40 out per 1M) context_window: 131072 attachments: [] supports_structured_output: true - input_cost_per_million: 0.23 - output_cost_per_million: 0.34 + input_cost_per_million: 0.269 + output_cost_per_million: 0.4 + cached_input_cost_per_million: 0.1345 - id: anthropic/claude-sonnet-4.6 display_name: Claude Sonnet 4.6 (via OpenRouter) @@ -29,3 +30,5 @@ models: supports_structured_output: true input_cost_per_million: 3.0 output_cost_per_million: 15.0 + cached_input_cost_per_million: 0.3 + cache_write_cost_per_million: 3.75 diff --git a/docsgpt/quotas/http.py b/docsgpt/quotas/http.py index 494b87cf..c9a1a180 100644 --- a/docsgpt/quotas/http.py +++ b/docsgpt/quotas/http.py @@ -11,4 +11,6 @@ def quota_exceeded_response(exceeded: QuotaExceeded) -> Response: """Return the 429 for an exhausted quota, with ``Retry-After`` set to the reset.""" response = make_response(jsonify(exceeded.to_payload()), 429) response.headers["Retry-After"] = str(exceeded.retry_after_seconds) + # The reset can be weeks away; the OpenAI SDKs would otherwise retry with backoff. + response.headers["x-should-retry"] = "false" return response diff --git a/frontend/src/admin/QuotaEditor.tsx b/frontend/src/admin/QuotaEditor.tsx index e962ed46..7b9de108 100644 --- a/frontend/src/admin/QuotaEditor.tsx +++ b/frontend/src/admin/QuotaEditor.tsx @@ -185,7 +185,7 @@ export default function QuotaEditor({ if (policy) remove(); return; } - const result = formToPolicy(form); + const result = formToPolicy(form, policy); if (!result.ok) { setError(result.error); return; @@ -196,6 +196,11 @@ export default function QuotaEditor({ return (

{inheritHint}

+ {policy && !policy.enabled ? ( +

+ This policy is disabled and is not enforced. Saving keeps it disabled. +

+ ) : null} ; if (!data?.success) return ; - const bucketPill = (policy: QuotaPolicy) => - isAll(policy) ? null : {policy.bucket} traffic; + const bucketPill = (policy: QuotaPolicy) => ( + <> + {isAll(policy) ? null : {policy.bucket} traffic} + {policy.enabled ? null : Disabled} + + ); return (
@@ -163,7 +167,7 @@ export default function Quotas() {

{instancePolicy - ? `Tokens: ${describeBudget(instancePolicy.token_limit, instancePolicy.token_unlimited, 'tokens')} · Cost: ${describeBudget(instancePolicy.cost_limit_usd, instancePolicy.cost_unlimited, 'cost')}` + ? `${instancePolicy.enabled ? '' : 'Disabled · '}Tokens: ${describeBudget(instancePolicy.token_limit, instancePolicy.token_unlimited, 'tokens')} · Cost: ${describeBudget(instancePolicy.cost_limit_usd, instancePolicy.cost_unlimited, 'cost')}` : 'No default: users without a team allowance or override are unlimited.'}

diff --git a/frontend/src/admin/quotaUtils.test.ts b/frontend/src/admin/quotaUtils.test.ts index 60e4f468..2afad959 100644 --- a/frontend/src/admin/quotaUtils.test.ts +++ b/frontend/src/admin/quotaUtils.test.ts @@ -54,6 +54,7 @@ describe('formToPolicy', () => { ok: true, policy: { bucket: 'all', + enabled: true, token_limit: 5000, token_unlimited: false, cost_limit_usd: 2.5, @@ -63,6 +64,13 @@ describe('formToPolicy', () => { }); }); + it('keeps a disabled policy disabled', () => { + const form = { ...base, tokenMode: 'limit' as const, tokenLimit: '10' }; + const stored = policy({ enabled: false }); + const result = formToPolicy(form, stored); + expect(result.ok && result.policy.enabled).toBe(false); + }); + it('sends unlimited without a limit', () => { const result = formToPolicy({ ...base, diff --git a/frontend/src/admin/quotaUtils.ts b/frontend/src/admin/quotaUtils.ts index 10038574..02bc2f6b 100644 --- a/frontend/src/admin/quotaUtils.ts +++ b/frontend/src/admin/quotaUtils.ts @@ -66,9 +66,15 @@ export function isEmptyForm(form: QuotaForm): boolean { return form.tokenMode === 'inherit' && form.costMode === 'inherit'; } -export function formToPolicy(form: QuotaForm): FormResult { +// ``existing`` carries the stored ``enabled`` flag through an edit: the form has +// no control for it, and a body without it would switch the policy back on. +export function formToPolicy( + form: QuotaForm, + existing?: QuotaPolicy | null, +): FormResult { const policy: Record = { bucket: 'all', + enabled: existing?.enabled ?? true, token_limit: null, token_unlimited: form.tokenMode === 'unlimited', cost_limit_usd: null, diff --git a/frontend/src/settings/components/UsageQuota.tsx b/frontend/src/settings/components/UsageQuota.tsx index 30511081..a2147beb 100644 --- a/frontend/src/settings/components/UsageQuota.tsx +++ b/frontend/src/settings/components/UsageQuota.tsx @@ -2,6 +2,7 @@ import { useEffect, useState } from 'react'; import { useTranslation } from 'react-i18next'; import { useSelector } from 'react-redux'; +import { usagePercent } from '../../admin/quotaUtils'; import userService from '../../api/services/userService'; import { selectToken } from '../../preferences/preferenceSlice'; @@ -24,10 +25,7 @@ function Meter({ }) { const { t } = useTranslation(); if (budget.limit === null) return null; - const percent = - budget.limit <= 0 - ? 100 - : Math.min(100, Math.max(0, (budget.used / budget.limit) * 100)); + const percent = usagePercent(budget.used, budget.limit); const tone = percent >= 100 ? 'bg-red-500' diff --git a/tests/quotas/test_enforcement.py b/tests/quotas/test_enforcement.py index c6300b22..0a54974c 100644 --- a/tests/quotas/test_enforcement.py +++ b/tests/quotas/test_enforcement.py @@ -30,11 +30,11 @@ def _spend(conn, user_id, tokens, api_key=None): TokenUsageRepository(conn).insert(user_id=user_id, api_key=api_key, prompt_tokens=tokens) -def _check(flask_app, agent_config, decoded_token=None): +def _check(flask_app, agent_config, decoded_token=None, agent_id=None): from docsgpt.api.answer.routes.base import BaseAnswerResource with flask_app.app_context(): - return BaseAnswerResource().check_usage(agent_config, decoded_token) + return BaseAnswerResource().check_usage(agent_config, decoded_token, agent_id=agent_id) class TestCheckUsage: @@ -90,6 +90,23 @@ class TestCheckUsage: _spend(db, "u1", 500) assert _check(flask_app, {}, {"sub": "u1"}) is None + def test_a_keyless_agent_chat_is_agent_traffic(self, db, flask_app): + agent_id = str(AgentsRepository(db).create("u1", "draft", "draft")["id"]) + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", bucket="agent", token_limit=10) + TokenUsageRepository(db).insert(user_id="u1", agent_id=agent_id, prompt_tokens=10) + + response = _check(flask_app, {"user_api_key": None}, {"sub": "u1"}, agent_id=agent_id) + + assert response.status_code == 429 + assert json.loads(response.data)["bucket"] == "agent" + # The same spend leaves chat without an agent alone. + assert _check(flask_app, {}, {"sub": "u1"}) is None + + def test_a_refusal_tells_sdk_clients_not_to_retry(self, db, flask_app): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + response = _check(flask_app, {}, {"sub": "u1"}) + assert response.headers["x-should-retry"] == "false" + class TestHeadless: def test_exhausted_owner_is_refused_before_the_run(self, db): @@ -125,7 +142,9 @@ class TestResumeRefusal: def _processor(self, user_id="u1"): from types import SimpleNamespace - return SimpleNamespace(agent_config={}, decoded_token={"sub": user_id}, initial_user_id=user_id) + return SimpleNamespace( + agent_config={}, decoded_token={"sub": user_id}, initial_user_id=user_id, agent_id=None + ) def test_a_refused_resume_releases_its_claim(self, db, flask_app): from docsgpt.api.answer.routes.base import BaseAnswerResource From 53facb460cbf614d8c1f74d65f4e7408df548a79 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Mon, 21 Sep 2026 16:10:31 +0100 Subject: [PATCH 130/130] chore: add Zed project config, .editorconfig and shared pyright settings - .zed/settings.json: ruff + basedpyright for Python (no format on save, the tree is not ruff-format clean), ESLint fixes then Prettier for the frontend, scan exclusions for caches and build outputs, .jwt_secret_key as private - .zed/tasks.json: dev services, API, worker, frontend, pytest/vitest for the current file or test, linting, uv lock + requirements export - .zed/debug.json: debugpy targets matching .vscode/launch.json - .editorconfig: whitespace rules for every editor - [tool.pyright] in pyproject.toml: venv, import root and excludes shared by pyright, basedpyright and Pylance - .gitignore: track only the shared files under .zed/ - CONTRIBUTING: editor setup section; fix the stale ESLint config path --- .editorconfig | 33 +++++++++++++++ .gitignore | 5 +++ .zed/debug.json | 44 ++++++++++++++++++++ .zed/settings.json | 89 ++++++++++++++++++++++++++++++++++++++++ .zed/tasks.json | 100 +++++++++++++++++++++++++++++++++++++++++++++ CONTRIBUTING.md | 28 ++++++++++++- pyproject.toml | 22 ++++++++++ 7 files changed, 320 insertions(+), 1 deletion(-) create mode 100644 .editorconfig create mode 100644 .zed/debug.json create mode 100644 .zed/settings.json create mode 100644 .zed/tasks.json diff --git a/.editorconfig b/.editorconfig new file mode 100644 index 00000000..9203830f --- /dev/null +++ b/.editorconfig @@ -0,0 +1,33 @@ +# https://editorconfig.org — shared whitespace rules for every editor. +root = true + +[*] +charset = utf-8 +end_of_line = lf +insert_final_newline = true +trim_trailing_whitespace = true +indent_style = space +indent_size = 2 + +[*.{py,pyi}] +indent_size = 4 +max_line_length = 120 + +[*.{sh,ps1,ini,toml}] +indent_size = 4 + +[{Dockerfile,Dockerfile.*,*.dockerfile}] +indent_size = 4 + +# Two trailing spaces are a hard line break in Markdown. +[*.{md,mdx}] +trim_trailing_whitespace = false + +[Makefile] +indent_style = tab + +# Generated or vendored; leave as produced. +[{uv.lock,package-lock.json,docsgpt/requirements*.txt}] +indent_size = unset +insert_final_newline = unset +trim_trailing_whitespace = unset diff --git a/.gitignore b/.gitignore index 27013d3b..5bb2d60b 100644 --- a/.gitignore +++ b/.gitignore @@ -195,6 +195,11 @@ docsgpt/static/ node_modules/ .vscode/settings.json .vscode/sftp.json +# Zed: the shared project config is tracked, anything else under .zed/ is local +.zed/* +!.zed/settings.json +!.zed/tasks.json +!.zed/debug.json /models/ model/ diff --git a/.zed/debug.json b/.zed/debug.json new file mode 100644 index 00000000..6c116182 --- /dev/null +++ b/.zed/debug.json @@ -0,0 +1,44 @@ +// Debug configurations for the Zed editor (https://zed.dev/docs/debugger), +// the counterparts of .vscode/launch.json. Start one with `debugger: start`. +// Zed picks the interpreter from the checkout's `.venv`. +[ + { + "label": "API (uvicorn)", + "adapter": "Debugpy", + "request": "launch", + "module": "uvicorn", + "args": ["docsgpt.asgi:asgi_app", "--host", "127.0.0.1", "--port", "7091"], + "cwd": "$ZED_WORKTREE_ROOT", + "env": { "PYTHONPATH": "$ZED_WORKTREE_ROOT" }, + "justMyCode": true + }, + { + // The solo pool keeps tasks in the debugged process, so breakpoints hit. + "label": "Celery worker (solo pool)", + "adapter": "Debugpy", + "request": "launch", + "module": "celery", + "args": ["-A", "docsgpt.app.celery", "worker", "-l", "INFO", "--pool=solo"], + "cwd": "$ZED_WORKTREE_ROOT", + "env": { "PYTHONPATH": "$ZED_WORKTREE_ROOT" }, + "justMyCode": true + }, + { + "label": "pytest: this file", + "adapter": "Debugpy", + "request": "launch", + "module": "pytest", + "args": ["--no-cov", "$ZED_RELATIVE_FILE"], + "cwd": "$ZED_WORKTREE_ROOT", + "env": { "PYTHONPATH": "$ZED_WORKTREE_ROOT" }, + "justMyCode": false + }, + { + "label": "Python: this file", + "adapter": "Debugpy", + "request": "launch", + "program": "$ZED_FILE", + "cwd": "$ZED_WORKTREE_ROOT", + "env": { "PYTHONPATH": "$ZED_WORKTREE_ROOT" } + } +] diff --git a/.zed/settings.json b/.zed/settings.json new file mode 100644 index 00000000..2a470819 --- /dev/null +++ b/.zed/settings.json @@ -0,0 +1,89 @@ +// Project settings for the Zed editor (https://zed.dev/docs/configuring-zed). +// They mirror what CI and the pre-commit hook enforce; whitespace rules live in +// .editorconfig and the Python analysis config in pyproject.toml +// ([tool.pyright]) so other editors share them. Personal preferences belong in +// your user settings, not here. +{ + // Zed replaces its defaults when this key is set, so they are repeated first. + // Build outputs, caches and local runtime data only add noise to the file + // finder and project search. + "file_scan_exclusions": [ + "**/.git", + "**/.svn", + "**/.hg", + "**/.jj", + "**/.sl", + "**/.repo", + "**/CVS", + "**/.DS_Store", + "**/Thumbs.db", + "**/.classpath", + "**/.settings", + "**/__pycache__", + "**/.ruff_cache", + "**/.pytest_cache", + "**/.mypy_cache", + "**/htmlcov", + "**/.next", + "frontend/dist", + "docsgpt/static", + "**/indexes", + "**/inputs", + "**/vectors", + "models" + ], + // Never shared with collaborators or sent to an AI assistant. The first six + // are Zed's defaults, which this key also replaces. + "private_files": [ + "**/.env*", + "**/*.pem", + "**/*.key", + "**/*.cert", + "**/*.crt", + "**/secrets.yml", + "**/.jwt_secret_key" + ], + "file_types": { + "Shell Script": [".env-template"], + "Dockerfile": ["Dockerfile*"] + }, + "languages": { + "Python": { + "language_servers": ["basedpyright", "ruff", "..."], + "formatter": { "language_server": { "name": "ruff" } }, + // CI runs `ruff check` only and most of the tree is not `ruff format` + // clean, so formatting on save would bury a change in unrelated diffs. + "format_on_save": "off", + "preferred_line_length": 120, + "wrap_guides": [120] + }, + // Same order as the lint-staged hook: ESLint fixes, then Prettier. + "TypeScript": { + "formatter": "prettier", + "format_on_save": "on", + "code_actions_on_format": { "source.fixAll.eslint": true } + }, + "TSX": { + "formatter": "prettier", + "format_on_save": "on", + "code_actions_on_format": { "source.fixAll.eslint": true } + }, + "JavaScript": { + "formatter": "prettier", + "format_on_save": "on", + "code_actions_on_format": { "source.fixAll.eslint": true } + }, + // Docs prose is reviewed by Vale, not reflowed by a formatter. + "Markdown": { "format_on_save": "off" }, + "MDX": { "format_on_save": "off" } + }, + "lsp": { + // The ESLint config is frontend/eslint.config.js, not at the root. + "eslint": { + "settings": { "workingDirectory": { "mode": "auto" } } + }, + "tailwindcss-language-server": { + "settings": { "classFunctions": ["cn", "cva", "clsx", "twMerge"] } + } + } +} diff --git a/.zed/tasks.json b/.zed/tasks.json new file mode 100644 index 00000000..3daf7eae --- /dev/null +++ b/.zed/tasks.json @@ -0,0 +1,100 @@ +// Project tasks for the Zed editor (https://zed.dev/docs/tasks). Run one with +// `task: spawn`. Python commands go through `uv run --no-sync`, which uses the +// checkout's `.venv` without changing what is installed in it; create it first +// with `uv sync`. See AGENTS.md for what each command does. +[ + { + "label": "services: Postgres + Redis", + "command": "docker compose -f deployment/docker-compose-dev.yaml up", + "cwd": "$ZED_WORKTREE_ROOT", + "use_new_terminal": true + }, + { + "label": "dev: API + worker + frontend", + "command": "uv run --no-sync docsgpt dev --ui", + "cwd": "$ZED_WORKTREE_ROOT", + "use_new_terminal": true + }, + { + "label": "dev: API + worker + frontend (mock LLM, no API key)", + "command": "uv run --no-sync docsgpt dev --ui --mock-llm", + "cwd": "$ZED_WORKTREE_ROOT", + "use_new_terminal": true + }, + { + "label": "backend: API", + "command": "uv run --no-sync docsgpt api --reload", + "cwd": "$ZED_WORKTREE_ROOT", + "use_new_terminal": true + }, + { + "label": "backend: Celery worker", + "command": "uv run --no-sync docsgpt worker", + "cwd": "$ZED_WORKTREE_ROOT", + "use_new_terminal": true + }, + { + "label": "backend: run migrations", + "command": "uv run --no-sync docsgpt migrate", + "cwd": "$ZED_WORKTREE_ROOT" + }, + { + "label": "frontend: dev server", + "command": "npm run dev", + "cwd": "$ZED_WORKTREE_ROOT/frontend", + "use_new_terminal": true + }, + { + "label": "frontend: build into docsgpt/static", + "command": "bash scripts/build_frontend.sh", + "cwd": "$ZED_WORKTREE_ROOT" + }, + // Coverage is switched off for partial runs: pytest.ini turns it on, and a + // report for one file is slow and misleading. + { + "label": "pytest: all", + "command": "uv run --no-sync python -m pytest", + "cwd": "$ZED_WORKTREE_ROOT" + }, + { + "label": "pytest: this file", + "command": "uv run --no-sync python -m pytest --no-cov \"$ZED_RELATIVE_FILE\"", + "cwd": "$ZED_WORKTREE_ROOT" + }, + { + "label": "pytest: test under cursor ($ZED_SYMBOL)", + "command": "uv run --no-sync python -m pytest --no-cov \"$ZED_RELATIVE_FILE\" -k \"$ZED_SYMBOL\"", + "cwd": "$ZED_WORKTREE_ROOT" + }, + { + "label": "pytest: last failed", + "command": "uv run --no-sync python -m pytest --no-cov --lf", + "cwd": "$ZED_WORKTREE_ROOT" + }, + { + "label": "vitest: all", + "command": "npm run test", + "cwd": "$ZED_WORKTREE_ROOT/frontend" + }, + { + "label": "vitest: this file", + "command": "npx vitest run \"$ZED_FILE\"", + "cwd": "$ZED_WORKTREE_ROOT/frontend" + }, + { + "label": "lint: ruff check --fix", + "command": "uv run --no-sync ruff check --fix .", + "cwd": "$ZED_WORKTREE_ROOT" + }, + { + "label": "lint: frontend (eslint --fix + prettier)", + "command": "npm run lint-fix && npm run format", + "cwd": "$ZED_WORKTREE_ROOT/frontend" + }, + // CI fails when docsgpt/requirements*.txt are stale against uv.lock. + { + "label": "deps: uv lock + export requirements", + "command": "uv lock && bash scripts/export_requirements.sh", + "cwd": "$ZED_WORKTREE_ROOT" + } +] diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 2510c7e2..55b48499 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -43,7 +43,7 @@ Tech Stack Overview: ### 🌐 Frontend Contributions (⚛️ React, Vite) * The updated Figma design can be found [here](https://www.figma.com/file/OXLtrl1EAy885to6S69554/DocsGPT?node-id=0%3A1&t=hjWVuxRg9yi5YkJ9-1). Please try to follow the guidelines. -* **Coding Style:** We follow a strict coding style enforced by ESLint and Prettier. Please ensure your code adheres to the configuration provided in our repository's `fronetend/.eslintrc.js` file. We recommend configuring your editor with ESLint and Prettier to help with this. +* **Coding Style:** We follow a strict coding style enforced by ESLint and Prettier. Please ensure your code adheres to the configuration provided in our repository's `frontend/eslint.config.js` and `frontend/prettier.config.cjs` files. We recommend configuring your editor with ESLint and Prettier to help with this. * **Component Structure:** Strive for small, reusable components. Favor functional components and hooks over class components where possible. * **State Management** If you need to add stores, please use Redux. @@ -75,6 +75,32 @@ Tech Stack Overview: ... ``` +### Editor setup + +Some configuration is shared by every editor, so you rarely need to set anything up by hand: + +- [`.editorconfig`](https://editorconfig.org) holds the whitespace rules (4 spaces for Python, 2 for TypeScript/JSON/YAML, LF line endings, final newline). Most editors read it natively or through a plugin. +- `[tool.pyright]` in `pyproject.toml` points Pyright, basedpyright and Pylance at the `.venv` created by `uv sync` and at the repository root for imports. +- `.ruff.toml`, `frontend/eslint.config.js` and `frontend/prettier.config.cjs` are picked up by the matching editor integrations. + +Editor-specific configuration that is tracked: + +- **VS Code:** `.vscode/launch.json` has debug targets for the API, the Celery worker and the frontend. +- **Zed:** open the repository root (not `frontend/`). `.zed/settings.json` configures the language servers and formatters, `.zed/tasks.json` adds tasks (`task: spawn`) for the dev services, the API, the worker, the frontend, tests and linting, and `.zed/debug.json` adds debug targets (`debugger: start`). Python files are not formatted on save because most of the tree is not `ruff format` clean; frontend files are, with ESLint fixes followed by Prettier, as in the pre-commit hook. Project settings cannot install extensions, so if you want the matching syntax support add this to your own Zed settings: + + ```json + { + "auto_install_extensions": { + "dockerfile": true, + "docker-compose": true, + "toml": true, + "mdx": true + } + } + ``` + +Personal preferences belong in your user settings; `.vscode/settings.json` and any other file under `.zed/` are ignored by git. + ### Testing To run unit tests from the root of the repository, execute: diff --git a/pyproject.toml b/pyproject.toml index 30f07df6..5e24bbaa 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -213,6 +213,28 @@ exclude = [ "docsgpt/vectors/", ] +# Shared by every editor's Python language server (pyright, basedpyright, +# Pylance) and by `pyright` on the command line. Imports are rooted at the +# checkout (`docsgpt.…`, `tests.…`), and the interpreter is the uv-managed +# `.venv`. "standard" keeps basedpyright from defaulting to its much stricter +# "recommended" mode; this is editor feedback, not a CI gate. +[tool.pyright] +pythonVersion = "3.12" +venvPath = "." +venv = ".venv" +extraPaths = ["."] +typeCheckingMode = "standard" +include = ["docsgpt", "application", "tests", "scripts"] +exclude = [ + "**/__pycache__", + "**/node_modules", + ".venv", + "docsgpt/static", + "docsgpt/indexes", + "docsgpt/inputs", + "docsgpt/vectors", +] + [[tool.uv.index]] name = "pytorch-cpu" url = "https://download.pytorch.org/whl/cpu"