diff --git a/lib/Controller/ServiceController.php b/lib/Controller/ServiceController.php index c86e82e2..6477a087 100644 --- a/lib/Controller/ServiceController.php +++ b/lib/Controller/ServiceController.php @@ -25,7 +25,7 @@ */ class ServiceController extends Controller { /** Properties that may only be changed with a confirmed password */ - private const SENSITIVE_PROPERTIES = ['url', 'api_key', 'basic_user', 'basic_password']; + private const SENSITIVE_PROPERTIES = ['url', 'api_key', 'basic_user', 'basic_password', 'extra_headers']; public function __construct( string $appName, @@ -59,7 +59,7 @@ public function create(): DataResponse { /** * Update the given properties of a service * - * The URL and the credentials can only be set through + * The URL, the credentials and the extra headers can only be set through * {@see self::updateSensitive()}. * * @param string $id ID of the service @@ -76,7 +76,7 @@ public function update(string $id, array $values): DataResponse { } /** - * Update the URL and the credentials of a service + * Update the URL, the credentials and the extra headers of a service * * Secrets that are sent back unchanged (as the placeholder the frontend * received) are kept. diff --git a/lib/Service/OpenAiAPIService.php b/lib/Service/OpenAiAPIService.php index a52e6242..0a66c994 100644 --- a/lib/Service/OpenAiAPIService.php +++ b/lib/Service/OpenAiAPIService.php @@ -408,6 +408,7 @@ public function createStreamedChatCompletion( ?string $toolMessage = null, ?array $tools = null, ?array $files = null, + ?string $conversationId = null, ): \Generator { if ($this->isQuotaExceeded($userId, Application::QUOTA_TYPE_TEXT, $service)) { throw new Exception($this->l10n->t('Text generation quota exceeded'), Http::STATUS_TOO_MANY_REQUESTS); @@ -439,6 +440,7 @@ public function createStreamedChatCompletion( true, 0, true, + $conversationId, ); $streamResult = yield from $this->streamingService->parseStreamChatResponse($response); @@ -468,11 +470,12 @@ public function createChatCompletion( ?string $toolMessage = null, ?array $tools = null, ?array $files = null, + ?string $conversationId = null, ): array { $response = $this->requestChatCompletion( $userId, $service, $model, $userPrompt, $systemPrompt, $history, $n, $maxTokens, $extraParams, $toolMessage, $tools, $files, - false, + false, $conversationId, ); if (isset($response['usage'], $response['usage']['total_tokens'])) { @@ -517,6 +520,7 @@ public function requestChatCompletion( ?array $tools = null, ?array $files = null, bool $stream = false, + ?string $conversationId = null, ): array { if ($this->isQuotaExceeded($userId, Application::QUOTA_TYPE_TEXT, $service)) { throw new Exception($this->l10n->t('Text generation quota exceeded'), Http::STATUS_TOO_MANY_REQUESTS); @@ -538,7 +542,7 @@ public function requestChatCompletion( $stream, ); - return $this->request($userId, $service, 'chat/completions', $params, 'POST'); + return $this->request($userId, $service, 'chat/completions', $params, 'POST', conversationId: $conversationId); } /** @@ -1171,6 +1175,30 @@ private function requestLocalAiImageEdit( return $this->request($userId, $service, 'images/generations', $params, 'POST'); } + /** + * Merge the extra headers configured on the service into the request + * options. They are applied before the authentication and content-type + * headers, so the service's own credentials always win over a configured + * Authorization header. + * + * Headers using the {$conversation_id} variable are dropped when the + * request has no conversation ID, see {@see ServiceConfig::expandHeaderValue()}. + * + * @param string|null $conversationId + * @param array $options + * @return array + */ + private function addExtraHeaders(ServiceConfig $service, array $options, ?string $conversationId = null): array { + foreach ($service->getExtraHeaders() as $header) { + $value = ServiceConfig::expandHeaderValue($header['value'], $conversationId); + if ($value === null) { + continue; + } + $options['headers'][$header['name']] = $value; + } + return $options; + } + /** * @param string|null $userId * @return array @@ -1183,6 +1211,7 @@ public function getImageRequestOptions(?string $userId, ServiceConfig $service): 'User-Agent' => Application::USER_AGENT, ], ]; + $requestOptions = $this->addExtraHeaders($service, $requestOptions); if ($service->getImageRequestAuth()) { if ($service->usesBasicAuth()) { @@ -1317,6 +1346,7 @@ private function updateExpProcessingTime(ServiceConfig $service, string $key, in * @param string|null $contentType * @param bool $logErrors if set to false error logs will be suppressed * @param int $retryCount number of retries that have been attempted so far + * @param string|null $conversationId the assistant conversation ID used to expand the {$conversation_id} token of the extra headers * @return array decoded request result or error * @throws Exception|UserFacingProcessingException */ @@ -1325,6 +1355,7 @@ public function request( ?string $contentType = null, bool $logErrors = true, int $retryCount = 0, bool $stream = false, + ?string $conversationId = null, ): array { try { // the user's own credentials take precedence over the admin ones @@ -1342,6 +1373,7 @@ public function request( 'User-Agent' => Application::USER_AGENT, ], ]; + $options = $this->addExtraHeaders($service, $options, $conversationId); if ($serviceUrl === Application::OPENAI_API_BASE_URL && $apiKey === '') { return ['error' => 'An API key is required for api.openai.com']; @@ -1474,7 +1506,7 @@ public function request( } $this->logger->warning("Rate limit exceeded, retrying in $sleep seconds", ['retry_count' => $retryCount]); sleep($sleep); - return $this->request($userId, $service, $endPoint, $params, $method, $contentType, $logErrors, $retryCount + 1, $stream); + return $this->request($userId, $service, $endPoint, $params, $method, $contentType, $logErrors, $retryCount + 1, $stream, $conversationId); } else { $this->logger->warning('Rate limit exceeded, maximum retries reached', ['retry_count' => $retryCount]); } diff --git a/lib/Service/ServiceConfig.php b/lib/Service/ServiceConfig.php index 0bd8b7bf..211c7124 100644 --- a/lib/Service/ServiceConfig.php +++ b/lib/Service/ServiceConfig.php @@ -56,11 +56,21 @@ class ServiceConfig implements JsonSerializable { 'stt_models' => 'array', 'tts_models' => 'array', 'quotas' => 'array', + 'extra_headers' => 'array', ]; /** Properties that are stored encrypted and never sent to the frontend */ public const SECRET_PROPERTIES = ['api_key', 'basic_password']; + /** Characters an HTTP header name is made of (the RFC 7230 token) */ + public const HEADER_NAME_PATTERN = '/^[a-zA-Z0-9!#$%&\'*+.^_`|~-]+$/'; + + /** The extra header value variable holding the assistant conversation ID of a chat request */ + public const CONVERSATION_ID_VARIABLE = '{$conversation_id}'; + + /** The variables extra header values may use */ + public const SUPPORTED_HEADER_VARIABLES = [self::CONVERSATION_ID_VARIABLE]; + /** * @param array $quotas * @param list $ttsVoices @@ -68,6 +78,7 @@ class ServiceConfig implements JsonSerializable { * @param list $imageModels * @param list $sttModels * @param list $ttsModels + * @param list $extraHeaders */ public function __construct( private string $id, @@ -104,6 +115,7 @@ public function __construct( private array $sttModels = [], private array $ttsModels = [], private array $quotas = Application::DEFAULT_QUOTAS, + private array $extraHeaders = [], ) { } @@ -186,6 +198,8 @@ public function with(array $values): self { break; case 'quotas': $new->quotas = self::normalizeQuotas($value, $new->quotas); break; + case 'extra_headers': $new->extraHeaders = self::normalizeHeaders($value); + break; } } return $new; @@ -226,6 +240,32 @@ private static function normalizeQuotas(mixed $quotas, array $base): array { return $normalized; } + /** + * Keep the name/value pairs that have a name, trim the rest. A row + * without a name is one the admin has not filled in yet, and sending it + * along in requests would only be confusing. + * + * @param mixed $headers + * @return list + */ + private static function normalizeHeaders(mixed $headers): array { + if (!is_array($headers)) { + return []; + } + $normalized = []; + foreach ($headers as $header) { + if (!is_array($header)) { + continue; + } + $name = trim((string)($header['name'] ?? '')); + if ($name === '') { + continue; + } + $normalized[] = ['name' => $name, 'value' => trim((string)($header['value'] ?? ''))]; + } + return $normalized; + } + public function getId(): string { return $this->id; } @@ -446,6 +486,38 @@ public function getQuota(int $quotaType): int { return $this->quotas[$quotaType] ?? 0; } + /** + * The extra headers to send with every request to this service + * + * @return list + */ + public function getExtraHeaders(): array { + return $this->extraHeaders; + } + + /** + * Replace the {$conversation_id} variable of an extra header value with + * the ID of the conversation of the request being sent. + * + * A value using the variable of a request that has no conversation is not + * meant to be sent at all: only chat requests know the conversation, so + * the header is dropped for the other requests, which is signalled by + * returning null. Anything else in the value, including text that merely + * looks like a variable, travels literally. + */ + public static function expandHeaderValue(string $value, ?string $conversationId): ?string { + if (!str_contains($value, self::CONVERSATION_ID_VARIABLE)) { + return $value; + } + if ($conversationId === null || $conversationId === '') { + return null; + } + $expanded = str_replace(self::CONVERSATION_ID_VARIABLE, $conversationId, $value); + // the ID comes from the task input of another app, so line breaks in + // it must not be able to forge headers + return preg_match('/[\r\n]/', $expanded) === 1 ? null : $expanded; + } + /** * Full representation, including secrets, as stored in app config * @@ -484,6 +556,7 @@ public function jsonSerialize(): array { 'stt_models' => $this->sttModels, 'tts_models' => $this->ttsModels, 'quotas' => $this->quotas, + 'extra_headers' => $this->extraHeaders, ]; } diff --git a/lib/Service/ServicesService.php b/lib/Service/ServicesService.php index 57959b82..366039e3 100644 --- a/lib/Service/ServicesService.php +++ b/lib/Service/ServicesService.php @@ -429,6 +429,28 @@ private function validate(array $values): void { && preg_match('/^\d+x\d+$/', (string)$values['default_image_size']) !== 1) { throw new Exception('Invalid image size value. Expected the format x'); } + if (isset($values['extra_headers'])) { + foreach ($values['extra_headers'] as $header) { + if (!is_array($header) || !isset($header['name'], $header['value']) + || !is_string($header['name']) || !is_string($header['value'])) { + throw new Exception('Invalid extra header. Expected a name and a value'); + } + $name = trim($header['name']); + if ($name !== '' && preg_match(ServiceConfig::HEADER_NAME_PATTERN, $name) !== 1) { + throw new Exception('Invalid extra header name. Only the characters allowed in an HTTP header name are accepted: ' . $name); + } + if (preg_match('/[\r\n]/', $header['value']) === 1) { + throw new Exception('Invalid extra header value. Line breaks are not allowed: ' . $name); + } + if (preg_match_all('/\{\$[^}]*\}/', $header['value'], $variables) > 0) { + foreach ($variables[0] as $variable) { + if (!in_array($variable, ServiceConfig::SUPPORTED_HEADER_VARIABLES, true)) { + throw new Exception('Unknown variable in an extra header value: ' . $variable . '. Supported variables: ' . implode(', ', ServiceConfig::SUPPORTED_HEADER_VARIABLES)); + } + } + } + } + } } /** diff --git a/lib/TaskProcessing/AudioToAudioChatProvider.php b/lib/TaskProcessing/AudioToAudioChatProvider.php index d051057c..8f35ba4b 100644 --- a/lib/TaskProcessing/AudioToAudioChatProvider.php +++ b/lib/TaskProcessing/AudioToAudioChatProvider.php @@ -89,6 +89,11 @@ public function getOptionalInputShape(): array { : $this->l->t('Speech speed modifier'), EShapeType::Number ), + 'conversation_id' => new ShapeDescriptor( + $this->l->t('Conversation ID'), + $this->l->t('The ID of the conversation this request belongs to.'), + EShapeType::Text + ), ]; } @@ -175,9 +180,13 @@ public function process(?string $userId, array $input, callable $reportProgress) 'audio' => ['voice' => $outputVoice, 'format' => 'mp3'], ]; $systemPrompt .= ' Producing text responses will break the user interface. Important: You have multimodal voice capability, and you use voice exclusively to respond.'; + $conversationId = isset($input['conversation_id']) && is_string($input['conversation_id']) + ? $input['conversation_id'] + : null; $completion = $this->openAiAPIService->createChatCompletion( $userId, $this->service, $this->model, null, $systemPrompt, $history, 1, 1000, - $extraParams, null, null, [$inputFile] + $extraParams, null, null, [$inputFile], + conversationId: $conversationId, ); $message = array_pop($completion['audio_messages']); // TODO find a way to force the model to answer with audio when there is only text in the history diff --git a/lib/TaskProcessing/MultimodalChatWithToolsProvider.php b/lib/TaskProcessing/MultimodalChatWithToolsProvider.php index f460bdd2..f98af6a0 100644 --- a/lib/TaskProcessing/MultimodalChatWithToolsProvider.php +++ b/lib/TaskProcessing/MultimodalChatWithToolsProvider.php @@ -69,6 +69,11 @@ public function getOptionalInputShape(): array { $this->l->t('The maximum number of words/tokens that can be generated in the completion.'), EShapeType::Number ), + 'conversation_id' => new ShapeDescriptor( + $this->l->t('Conversation ID'), + $this->l->t('The ID of the conversation this request belongs to.'), + EShapeType::Text + ), ]; } @@ -156,10 +161,16 @@ public function process( if (isset($input['max_tokens']) && is_int($input['max_tokens'])) { $maxTokens = $input['max_tokens']; } + + $conversationId = isset($input['conversation_id']) && is_string($input['conversation_id']) + ? $input['conversation_id'] + : null; + try { if ($preferStreaming) { $chunks = $this->openAiAPIService->createStreamedChatCompletion( - $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools, $inputAttachments + $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools, $inputAttachments, + conversationId: $conversationId, ); $time = microtime(true); $streamedOutput = ''; @@ -197,7 +208,8 @@ public function process( $returnValue = $chunks->getReturn(); } else { $returnValue = $this->openAiAPIService->createChatCompletion( - $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools, $inputAttachments + $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools, $inputAttachments, + conversationId: $conversationId, ); } } catch (UserFacingProcessingException $e) { diff --git a/lib/TaskProcessing/TextToTextChatProvider.php b/lib/TaskProcessing/TextToTextChatProvider.php index ad64b833..7ad10a3e 100644 --- a/lib/TaskProcessing/TextToTextChatProvider.php +++ b/lib/TaskProcessing/TextToTextChatProvider.php @@ -68,6 +68,11 @@ public function getOptionalInputShape(): array { $this->l->t('The memories to be injected into the chat session.'), EShapeType::ListOfTexts ), + 'conversation_id' => new ShapeDescriptor( + $this->l->t('Conversation ID'), + $this->l->t('The ID of the conversation this request belongs to.'), + EShapeType::Text + ), ]; } @@ -129,9 +134,13 @@ public function process( $maxTokens = $input['max_tokens']; } + $conversationId = isset($input['conversation_id']) && is_string($input['conversation_id']) + ? $input['conversation_id'] + : null; + try { if ($preferStreaming) { - $chunks = $this->openAiAPIService->createStreamedChatCompletion($userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens); + $chunks = $this->openAiAPIService->createStreamedChatCompletion($userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, conversationId: $conversationId); $time = microtime(true); $streamedOutput = ''; $streamedReasoning = ''; @@ -169,7 +178,7 @@ public function process( $completion = $returnValue['messages']; $reasoning = $returnValue['reasoning_messages']; } else { - $returnValue = $this->openAiAPIService->createChatCompletion($userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens); + $returnValue = $this->openAiAPIService->createChatCompletion($userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, conversationId: $conversationId); $completion = $returnValue['messages']; $reasoning = $returnValue['reasoning_messages']; } diff --git a/lib/TaskProcessing/TextToTextChatWithToolsProvider.php b/lib/TaskProcessing/TextToTextChatWithToolsProvider.php index d90cd00a..3b1b6fe2 100644 --- a/lib/TaskProcessing/TextToTextChatWithToolsProvider.php +++ b/lib/TaskProcessing/TextToTextChatWithToolsProvider.php @@ -63,6 +63,11 @@ public function getOptionalInputShape(): array { $this->l->t('The maximum number of words/tokens that can be generated in the completion.'), EShapeType::Number ), + 'conversation_id' => new ShapeDescriptor( + $this->l->t('Conversation ID'), + $this->l->t('The ID of the conversation this request belongs to.'), + EShapeType::Text + ), ]; } @@ -138,10 +143,15 @@ public function process( $maxTokens = $input['max_tokens']; } + $conversationId = isset($input['conversation_id']) && is_string($input['conversation_id']) + ? $input['conversation_id'] + : null; + try { if ($preferStreaming) { $chunks = $this->openAiAPIService->createStreamedChatCompletion( - $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools + $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools, + conversationId: $conversationId, ); $time = microtime(true); $streamedOutput = ''; @@ -179,7 +189,8 @@ public function process( $returnValue = $chunks->getReturn(); } else { $returnValue = $this->openAiAPIService->createChatCompletion( - $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools + $userId, $this->service, $this->model, $userPrompt, $systemPrompt, $history, 1, $maxTokens, null, $toolMessage, $tools, + conversationId: $conversationId, ); } } catch (UserFacingProcessingException $e) { diff --git a/src/components/AdminSettings.vue b/src/components/AdminSettings.vue index 91f9c0b6..1bb2914a 100644 --- a/src/components/AdminSettings.vue +++ b/src/components/AdminSettings.vue @@ -362,7 +362,8 @@ export default { }, /** * @param {string} serviceId ID of the service to save - * @param {boolean} sensitive whether to send the URL and the credentials + * @param {boolean} sensitive whether to send the URL, the credentials + * and the extra headers */ async putService(serviceId, sensitive) { const service = this.services.find(s => s.id === serviceId) @@ -375,6 +376,7 @@ export default { basic_user: (service.basic_user ?? '').trim(), api_key: (service.api_key ?? '').trim(), basic_password: (service.basic_password ?? '').trim(), + extra_headers: (service.extra_headers ?? []).map(header => ({ ...header })), } : { name: service.name, diff --git a/src/components/ServiceForm.vue b/src/components/ServiceForm.vue index 6ccd22b5..3f6dcd60 100644 --- a/src/components/ServiceForm.vue +++ b/src/components/ServiceForm.vue @@ -153,6 +153,46 @@ + +

{{ t('integration_openai', 'Extra request headers') }}

+ + {{ t('integration_openai', 'Headers sent with every request to this service. The service\'s own authentication always takes precedence over a configured Authorization header.') }} +
+ {{ t('integration_openai', 'The only supported variable is {example}: it is replaced with the ID of the current assistant conversation, and such a header is only sent with the chat requests that know that ID.', { example: '{$conversation_id}' }) }} +
+
+ + + + + +
+
+ + + {{ t('integration_openai', 'Add a header') }} + +
+

{{ t('integration_openai', 'Exposed models') }}

@@ -402,6 +442,7 @@ import DeleteOutlineIcon from 'vue-material-design-icons/DeleteOutline.vue' import EarthIcon from 'vue-material-design-icons/Earth.vue' import HelpCircleOutlineIcon from 'vue-material-design-icons/HelpCircleOutline.vue' import KeyOutlineIcon from 'vue-material-design-icons/KeyOutline.vue' +import PlusIcon from 'vue-material-design-icons/Plus.vue' import RefreshIcon from 'vue-material-design-icons/Refresh.vue' import UnfoldLessHorizontalIcon from 'vue-material-design-icons/UnfoldLessHorizontal.vue' import UnfoldMoreHorizontalIcon from 'vue-material-design-icons/UnfoldMoreHorizontal.vue' @@ -430,6 +471,7 @@ export default { EarthIcon, HelpCircleOutlineIcon, KeyOutlineIcon, + PlusIcon, RefreshIcon, UnfoldLessHorizontalIcon, UnfoldMoreHorizontalIcon, @@ -470,6 +512,11 @@ export default { // other property of the same request llmExtraParams: this.service.llm_extra_params ?? '', defaultImageSize: this.service.default_image_size ?? '', + // edited locally like the two above: rows that are still being + // typed, or whose name is not a valid header name yet, are not + // sent to the backend, which would reject them and with them the + // rest of the sensitive payload + extraHeaders: (this.service.extra_headers ?? []).map(header => ({ ...header })), // to prevent some browsers from filling fields with remembered passwords readonly: true, models: null, @@ -498,6 +545,12 @@ export default { const size = this.defaultImageSize.trim() return size === '' || /^\d+x\d+$/.test(size) }, + /** Whether any row is new or half-typed, so the backend must not overwrite the list */ + extraHeadersPending() { + return this.extraHeaders.some( + header => header.name.trim() === '' || this.headerDirty(header), + ) + }, modalities() { return [ { @@ -586,6 +639,11 @@ export default { this.defaultImageSize = value ?? '' } }, + 'service.extra_headers'(value) { + if (!this.extraHeadersPending) { + this.extraHeaders = (value ?? []).map(header => ({ ...header })) + } + }, }, mounted() { @@ -619,6 +677,42 @@ export default { quotas[index] = isNaN(parsed) || parsed < 0 ? 0 : parsed this.onInput({ quotas }) }, + headerNameInvalid(name) { + const trimmed = name.trim() + return trimmed !== '' && !/^[a-zA-Z0-9!#$%&'*+.^_`|~-]+$/.test(trimmed) + }, + headerValueInvalid(value) { + for (const variable of value.matchAll(/\{\$[^}]*\}/g)) { + if (variable[0] !== '{$conversation_id}') { + return true + } + } + return false + }, + headerDirty(header) { + return header.name !== header.name.trim() + || header.value !== header.value.trim() + || this.headerNameInvalid(header.name) + || this.headerValueInvalid(header.value) + }, + addExtraHeader() { + this.extraHeaders.push({ name: '', value: '' }) + }, + removeExtraHeader(index) { + this.extraHeaders.splice(index, 1) + this.onExtraHeadersInput() + }, + onExtraHeadersInput() { + // a half-typed row must not reject the rest of the sensitive + // payload it travels with, so hold it back until it is complete + if (this.extraHeaders.some(header => this.headerDirty(header))) { + return + } + const headers = this.extraHeaders + .filter(header => header.name.trim() !== '') + .map(header => ({ name: header.name.trim(), value: header.value.trim() })) + this.onSensitiveInput({ extra_headers: headers }) + }, async loadModels() { this.loadingModels = true try { @@ -705,6 +799,10 @@ export default { align-items: start; } + &.align-bottom { + align-items: flex-end; + } + .input { width: 300px; } diff --git a/tests/unit/Service/MultiServiceTest.php b/tests/unit/Service/MultiServiceTest.php index f1cb0228..4e03f946 100644 --- a/tests/unit/Service/MultiServiceTest.php +++ b/tests/unit/Service/MultiServiceTest.php @@ -27,6 +27,7 @@ use OCA\OpenAi\TaskProcessing\ProviderFactory; use OCA\OpenAi\TaskProcessing\TextToImageProvider; use OCA\OpenAi\TaskProcessing\TextToSpeechProvider; +use OCA\OpenAi\TaskProcessing\TextToTextChatProvider; use OCP\Http\Client\IClient; use OCP\Http\Client\IClientService; use OCP\IAppConfig; @@ -53,6 +54,8 @@ class MultiServiceTest extends TestCase { public const APIKEY_TRANSCRIPTION = 'This is a transcription PHPUnit test API key'; public const REQUEST_TIMEOUT_TRANSCRIPTION = 14; public const TRANSCRIPTION_MODEL = 'my-whisper-model'; + public const TEXT_BASE = 'https://text-generator.ai/v1'; + public const APIKEY_TEXT = 'This is a text PHPUnit test API key'; public const TEXT_MODEL = 'my/text-model'; private OpenAiAPIService $openAiApiService; @@ -456,6 +459,275 @@ public function testAudioToTextProvider(): void { $audioToTextProvider->process(self::TEST_USER1, ['input' => $file], fn () => null); } + public function testExtraHeadersAreStoredNormalized(): void { + $service = $this->addService([ + 'url' => self::SPEECH_BASE, + 'extra_headers' => [ + ['name' => ' X-Tenant ', 'value' => ' acme '], + ['name' => '', 'value' => 'dropped along with its name'], + ], + ]); + + $this->assertSame( + [['name' => 'X-Tenant', 'value' => 'acme']], + $service->getExtraHeaders(), + ); + } + + /** + * @dataProvider invalidExtraHeadersProvider + */ + public function testInvalidExtraHeadersAreRejected(array $extraHeaders): void { + $this->expectException(\Exception::class); + $this->servicesService->addService([ + 'url' => self::SPEECH_BASE, + 'extra_headers' => $extraHeaders, + ]); + } + + public function invalidExtraHeadersProvider(): array { + return [ + 'a name that is not an HTTP token' => [[['name' => 'X Api Key', 'value' => 'secret']]], + 'a name with an injected line break' => [[['name' => "X-Tenant\r\nX-Evil", 'value' => 'a']]], + 'a value with an injected line break' => [[['name' => 'X-Tenant', 'value' => "a\r\nb"]]], + 'a row without a value' => [[['name' => 'X-Tenant']]], + 'a row that is not a pair' => [[['nope']]], + 'a value with an unknown variable' => [[['name' => 'X-Session', 'value' => '{$conversationid}']]], + 'a value with a mistyped variable' => [[['name' => 'X-Session', 'value' => '{$Conversation_ID}']]], + 'a value with an empty variable' => [[['name' => 'X-Session', 'value' => '{$}']]], + ]; + } + + public function testExtraHeadersAreSentAndCannotOverrideTheApiKey(): void { + $service = $this->addService([ + 'url' => self::SPEECH_BASE, + 'api_key' => self::APIKEY_SPEECH, + 'request_timeout' => self::REQUEST_TIMEOUT_SPEECH, + 'tts_models' => [self::SPEECH_MODEL], + 'extra_headers' => [ + ['name' => 'X-Tenant', 'value' => 'acme'], + ['name' => 'Authorization', 'value' => 'Bearer injected'], + ], + ]); + + $ttsProvider = new TextToSpeechProvider( + $this->openAiApiService, + $this->createMock(\OCP\IL10N::class), + $this->createMock(\Psr\Log\LoggerInterface::class), + \OCP\Server::get(WatermarkingService::class), + $service, + self::SPEECH_MODEL, + ); + + $inputText = 'This is a test prompt'; + + $response = file_get_contents(__DIR__ . '/../../res/speech.mp3'); + + if (!$response) { + throw new \RuntimeException('Could not read test resourcce `speech.mp3`'); + } + + $url = self::SPEECH_BASE . '/audio/speech'; + + $options = ['timeout' => self::REQUEST_TIMEOUT_SPEECH, 'headers' => ['User-Agent' => Application::USER_AGENT, 'X-Tenant' => 'acme', 'Authorization' => 'Bearer ' . self::APIKEY_SPEECH, 'Content-Type' => 'application/json'], 'nextcloud' => ['allow_local_address' => true]]; + $options['body'] = json_encode([ + 'input' => $inputText, + 'voice' => Application::DEFAULT_SPEECH_VOICE, + 'model' => self::SPEECH_MODEL, + 'response_format' => 'mp3', + 'speed' => 1, + ]); + + $iResponse = $this->createMock(\OCP\Http\Client\IResponse::class); + $iResponse->method('getBody')->willReturn($response); + $iResponse->method('getStatusCode')->willReturn(200); + + $this->iClient->expects($this->once())->method('post')->with($url, $options)->willReturn($iResponse); + + $ttsProvider->process(self::TEST_USER1, ['input' => $inputText], fn () => null, includeWatermark: false); + } + + public function testExtraHeadersInImageRequestOptions(): void { + $service = $this->addService([ + 'url' => self::IMAGE_BASE, + 'api_key' => self::APIKEY_IMAGE, + 'image_request_auth' => false, + 'extra_headers' => [['name' => 'X-Tenant', 'value' => 'acme']], + ]); + + $options = $this->openAiApiService->getImageRequestOptions(self::TEST_USER1, $service); + + $this->assertSame([ + 'timeout' => Application::OPENAI_DEFAULT_REQUEST_TIMEOUT, + 'headers' => ['User-Agent' => Application::USER_AGENT, 'X-Tenant' => 'acme'], + ], $options); + } + + public function testConversationIdHeaderIsExpandedOnChatRequests(): void { + $service = $this->addService([ + 'url' => self::TEXT_BASE, + 'api_key' => self::APIKEY_TEXT, + 'text_models' => [self::TEXT_MODEL], + 'extra_headers' => [ + ['name' => 'X-Tenant', 'value' => 'acme'], + ['name' => 'X-Session', 'value' => '{$conversation_id}'], + ], + ]); + + $chatProvider = new TextToTextChatProvider( + $this->openAiApiService, + $this->createMock(\OCP\IL10N::class), + $service, + self::TEXT_MODEL, + ); + + $systemPrompt = 'You are a helpful assistant'; + $userPrompt = 'Hello'; + + $response = json_encode([ + 'choices' => [ + ['message' => ['role' => 'assistant', 'content' => 'Chat answer']], + ], + ]); + + $url = self::TEXT_BASE . '/chat/completions'; + $options = ['timeout' => Application::OPENAI_DEFAULT_REQUEST_TIMEOUT, 'headers' => ['User-Agent' => Application::USER_AGENT, 'X-Tenant' => 'acme', 'X-Session' => '4242', 'Authorization' => 'Bearer ' . self::APIKEY_TEXT, 'Content-Type' => 'application/json'], 'nextcloud' => ['allow_local_address' => true]]; + $options['body'] = json_encode([ + 'model' => self::TEXT_MODEL, + 'messages' => [ + ['role' => 'system', 'content' => $systemPrompt], + ['role' => 'user', 'content' => $userPrompt], + ], + 'n' => 1, + 'stream' => false, + 'max_tokens' => Application::DEFAULT_MAX_NUM_OF_TOKENS, + ]); + + $iResponse = $this->createMock(\OCP\Http\Client\IResponse::class); + $iResponse->method('getHeader')->with('Content-Type')->willReturn('application/json'); + $iResponse->method('getBody')->willReturn($response); + $iResponse->method('getStatusCode')->willReturn(200); + + $this->iClient->expects($this->once())->method('post')->with($url, $options)->willReturn($iResponse); + + $result = $chatProvider->process(self::TEST_USER1, [ + 'input' => $userPrompt, + 'system_prompt' => $systemPrompt, + 'history' => [], + 'conversation_id' => '4242', + ], fn () => null); + + $this->assertSame('Chat answer', $result['output']); + } + + public function testConversationIdHeaderIsDroppedWhenTheConversationIsUnknown(): void { + $service = $this->addService([ + 'url' => self::TEXT_BASE, + 'api_key' => self::APIKEY_TEXT, + 'text_models' => [self::TEXT_MODEL], + 'extra_headers' => [ + ['name' => 'X-Tenant', 'value' => 'acme'], + ['name' => 'X-Session', 'value' => 'conv-{$conversation_id}'], + ], + ]); + + $chatProvider = new TextToTextChatProvider( + $this->openAiApiService, + $this->createMock(\OCP\IL10N::class), + $service, + self::TEXT_MODEL, + ); + + $systemPrompt = 'You are a helpful assistant'; + $userPrompt = 'Hello'; + + $response = json_encode([ + 'choices' => [ + ['message' => ['role' => 'assistant', 'content' => 'Chat answer']], + ], + ]); + + // the same request without the conversation_id input: only the header + // that does not reference it survives + $url = self::TEXT_BASE . '/chat/completions'; + $options = ['timeout' => Application::OPENAI_DEFAULT_REQUEST_TIMEOUT, 'headers' => ['User-Agent' => Application::USER_AGENT, 'X-Tenant' => 'acme', 'Authorization' => 'Bearer ' . self::APIKEY_TEXT, 'Content-Type' => 'application/json'], 'nextcloud' => ['allow_local_address' => true]]; + $options['body'] = json_encode([ + 'model' => self::TEXT_MODEL, + 'messages' => [ + ['role' => 'system', 'content' => $systemPrompt], + ['role' => 'user', 'content' => $userPrompt], + ], + 'n' => 1, + 'stream' => false, + 'max_tokens' => Application::DEFAULT_MAX_NUM_OF_TOKENS, + ]); + + $iResponse = $this->createMock(\OCP\Http\Client\IResponse::class); + $iResponse->method('getHeader')->with('Content-Type')->willReturn('application/json'); + $iResponse->method('getBody')->willReturn($response); + $iResponse->method('getStatusCode')->willReturn(200); + + $this->iClient->expects($this->once())->method('post')->with($url, $options)->willReturn($iResponse); + + $result = $chatProvider->process(self::TEST_USER1, [ + 'input' => $userPrompt, + 'system_prompt' => $systemPrompt, + 'history' => [], + ], fn () => null); + + $this->assertSame('Chat answer', $result['output']); + } + + public function testConversationIdHeaderIsDroppedOnNonChatRequests(): void { + $service = $this->addService([ + 'url' => self::SPEECH_BASE, + 'api_key' => self::APIKEY_SPEECH, + 'request_timeout' => self::REQUEST_TIMEOUT_SPEECH, + 'tts_models' => [self::SPEECH_MODEL], + 'extra_headers' => [ + ['name' => 'X-Tenant', 'value' => 'acme'], + ['name' => 'X-Session', 'value' => '{$conversation_id}'], + ], + ]); + + $ttsProvider = new TextToSpeechProvider( + $this->openAiApiService, + $this->createMock(\OCP\IL10N::class), + $this->createMock(\Psr\Log\LoggerInterface::class), + \OCP\Server::get(WatermarkingService::class), + $service, + self::SPEECH_MODEL, + ); + + $inputText = 'This is a test prompt'; + + $response = file_get_contents(__DIR__ . '/../../res/speech.mp3'); + + if (!$response) { + throw new \RuntimeException('Could not read test resourcce `speech.mp3`'); + } + + // speech is not a chat, so the header referencing the conversation is + // not sent, while the static one is + $url = self::SPEECH_BASE . '/audio/speech'; + $options = ['timeout' => self::REQUEST_TIMEOUT_SPEECH, 'headers' => ['User-Agent' => Application::USER_AGENT, 'X-Tenant' => 'acme', 'Authorization' => 'Bearer ' . self::APIKEY_SPEECH, 'Content-Type' => 'application/json'], 'nextcloud' => ['allow_local_address' => true]]; + $options['body'] = json_encode([ + 'input' => $inputText, + 'voice' => Application::DEFAULT_SPEECH_VOICE, + 'model' => self::SPEECH_MODEL, + 'response_format' => 'mp3', + 'speed' => 1, + ]); + + $iResponse = $this->createMock(\OCP\Http\Client\IResponse::class); + $iResponse->method('getBody')->willReturn($response); + $iResponse->method('getStatusCode')->willReturn(200); + + $this->iClient->expects($this->once())->method('post')->with($url, $options)->willReturn($iResponse); + + $ttsProvider->process(self::TEST_USER1, ['input' => $inputText], fn () => null, includeWatermark: false); + } + public function testProvidersOfDifferentServicesHaveDifferentIds(): void { $first = $this->addService(['url' => self::IMAGE_BASE, 'image_models' => [self::IMAGE_MODEL]]); $second = $this->addService(['url' => self::SPEECH_BASE, 'image_models' => [self::IMAGE_MODEL]]); diff --git a/tests/unit/Service/ServiceConfigTest.php b/tests/unit/Service/ServiceConfigTest.php new file mode 100644 index 00000000..2c912c15 --- /dev/null +++ b/tests/unit/Service/ServiceConfigTest.php @@ -0,0 +1,60 @@ + ['acme', '42', 'acme'], + 'the variable alone' => ['{$conversation_id}', '42', '42'], + 'the variable in a sentence' => ['conv-{$conversation_id}-v2', '42', 'conv-42-v2'], + 'the variable twice' => ['{$conversation_id}/{$conversation_id}', '42', '42/42'], + 'a request without a conversation' => ['{$conversation_id}', null, null], + 'an empty conversation ID' => ['{$conversation_id}', '', null], + 'a mistyped variable travels literally' => ['{$Conversation_ID}', '42', '{$Conversation_ID}'], + 'a conversation ID smuggling line breaks' => ['{$conversation_id}', "42\r\nX-Evil: yes", null], + ]; + } + + /** + * @dataProvider expandHeaderValueProvider + */ + public function testExpandHeaderValue(string $value, ?string $conversationId, ?string $expected): void { + $this->assertSame($expected, ServiceConfig::expandHeaderValue($value, $conversationId)); + } + + public function testStoredServiceWithoutExtraHeadersDefaultsToNone(): void { + // a service stored before extra headers existed + $service = ServiceConfig::fromArray('s1', ['url' => 'https://example.com/v1']); + $this->assertSame([], $service->getExtraHeaders()); + } + + public function testExtraHeadersAreNormalizedOnStorage(): void { + $service = ServiceConfig::fromArray('s1', [ + 'extra_headers' => [ + ['name' => ' X-Tenant ', 'value' => ' acme '], + ['name' => '', 'value' => 'dropped with its name'], + ['name' => 'X-Empty-Value', 'value' => ''], + ], + ]); + + $this->assertSame([ + ['name' => 'X-Tenant', 'value' => 'acme'], + ['name' => 'X-Empty-Value', 'value' => ''], + ], $service->getExtraHeaders()); + + // what is normalized is also what the next read from storage gets back + $reread = ServiceConfig::fromArray('s1', $service->jsonSerialize()); + $this->assertSame($service->getExtraHeaders(), $reread->getExtraHeaders()); + } +}