diff --git a/.gitignore b/.gitignore index 679c41e..b3d259f 100644 --- a/.gitignore +++ b/.gitignore @@ -44,6 +44,9 @@ env.d/development/* !env.d/development/*.dist env.d/terraform +# Configuration +**/conversations/configuration/llm/dev.json + # npm node_modules diff --git a/CHANGELOG.md b/CHANGELOG.md index f7fcab3..0f829ab 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,7 @@ and this project adheres to - ✨(front) add ui kit #240 - 🧱(files) allow to use S3 storage without external access #849 +- ✨(backend) add FindRagBackend #209 ### Changed diff --git a/compose.yml b/compose.yml index eaa3275..a06ee63 100644 --- a/compose.yml +++ b/compose.yml @@ -71,6 +71,9 @@ services: - "host.docker.internal:host-gateway" ports: - "8071:8000" + networks: + - default + - lasuite volumes: - ./src/backend:/app - ./data/static:/data/static @@ -90,6 +93,9 @@ services: image: nginx:1.25 ports: - "8083:8083" + networks: + - default + - lasuite volumes: - ./docker/files/etc/nginx/conf.d:/etc/nginx/conf.d:ro depends_on: @@ -178,3 +184,8 @@ services: kc_postgresql: condition: service_healthy restart: true + +networks: + lasuite: + name: lasuite-network + driver: bridge diff --git a/docs/env.md b/docs/env.md index ad601a4..dff03ef 100644 --- a/docs/env.md +++ b/docs/env.md @@ -95,6 +95,9 @@ These are the environment variables you can set for the `conversations-backend` | CACHES_KEY_PREFIX | The prefix used to every cache keys. | conversations | | THEME_CUSTOMIZATION_FILE_PATH | full path to the file customizing the theme. An example is provided in src/backend/conversations/configuration/theme/default.json | BASE_DIR/conversations/configuration/theme/default.json | | THEME_CUSTOMIZATION_CACHE_TIMEOUT | Cache duration for the customization settings | 86400 | +| FIND_API_KEY | API key of Find | | +| FIND_API_URL | URL of Find | `https://app-find/api` | +| FIND_API_TIMEOUT | Find API timeout | 30 | ## conversations-frontend image diff --git a/docs/llm-configuration.md b/docs/llm-configuration.md index 8168e8e..55aeba9 100644 --- a/docs/llm-configuration.md +++ b/docs/llm-configuration.md @@ -244,8 +244,8 @@ For Mistral AI models using the Etalab platform: { "models": [ { - "hrid": "mistral-large", - "model_name": "mistral-large-latest", + "hrid": "mistral-medium", + "model_name": "mistral-medium-2508", "human_readable_name": "Mistral Large (Etalab)", "provider_name": "mistral-etalab", "profile": null, diff --git a/docs/tools/web_search_brave.md b/docs/tools/web_search_brave.md index 8a00021..6e39e39 100644 --- a/docs/tools/web_search_brave.md +++ b/docs/tools/web_search_brave.md @@ -357,6 +357,7 @@ The RAG backend performs semantic search to find the most relevant content: rag_results = document_store.search( query, results_count=settings.BRAVE_RAG_WEB_SEARCH_CHUNK_NUMBER, + **kwargs, # Additional search parameters like session with access_token ) ``` diff --git a/src/backend/chat/agent_rag/document_rag_backends/albert_rag_backend.py b/src/backend/chat/agent_rag/document_rag_backends/albert_rag_backend.py index 8af560b..172df52 100644 --- a/src/backend/chat/agent_rag/document_rag_backends/albert_rag_backend.py +++ b/src/backend/chat/agent_rag/document_rag_backends/albert_rag_backend.py @@ -213,13 +213,14 @@ class AlbertRagBackend(BaseRagBackend): # pylint: disable=too-many-instance-att logger.debug(response.json()) response.raise_for_status() - def search(self, query, results_count: int = 4) -> RAGWebResults: + def search(self, query: str, results_count: int = 4, **kwargs) -> RAGWebResults: """ Perform a search using the Albert API based on the provided query. Args: query (str): The search query. results_count (int): The number of results to return. + **kwargs: Additional arguments. Returns: RAGWebResults: The search results. @@ -256,7 +257,7 @@ class AlbertRagBackend(BaseRagBackend): # pylint: disable=too-many-instance-att ), ) - async def asearch(self, query, results_count: int = 4) -> RAGWebResults: + async def asearch(self, query, results_count: int = 4, **kwargs) -> RAGWebResults: """ Perform an asynchronous search using the Albert API based on the provided query. diff --git a/src/backend/chat/agent_rag/document_rag_backends/base_rag_backend.py b/src/backend/chat/agent_rag/document_rag_backends/base_rag_backend.py index 78d25cd..deee2ce 100644 --- a/src/backend/chat/agent_rag/document_rag_backends/base_rag_backend.py +++ b/src/backend/chat/agent_rag/document_rag_backends/base_rag_backend.py @@ -1,8 +1,8 @@ """Implementation of the Albert API for RAG document search.""" import logging +from abc import ABC, abstractmethod from contextlib import asynccontextmanager, contextmanager -from io import BytesIO from typing import List, Optional from asgiref.sync import sync_to_async @@ -12,7 +12,7 @@ from chat.agent_rag.constants import RAGWebResults logger = logging.getLogger(__name__) -class BaseRagBackend: +class BaseRagBackend(ABC): """Base class for RAG backends.""" def __init__( @@ -60,6 +60,7 @@ class BaseRagBackend: ) return collection_ids + @abstractmethod def create_collection(self, name: str, description: Optional[str] = None) -> str: """ Create a temporary collection for the search operation. @@ -74,7 +75,7 @@ class BaseRagBackend: """ return await sync_to_async(self.create_collection)(name=name, description=description) - def parse_document(self, name: str, content_type: str, content: BytesIO): + def parse_document(self, name: str, content_type: str, content: bytes): """ Parse the document and prepare it for the search operation. This method should handle the logic to convert the document @@ -142,17 +143,28 @@ class BaseRagBackend: """ return await sync_to_async(self.delete_collection)() - def search(self, query, results_count: int = 4) -> RAGWebResults: + @abstractmethod + def search(self, query: str, results_count: int = 4, **kwargs) -> RAGWebResults: """ Search the collection for the given query. + + Args: + query: The search query string. + results_count: Number of results to return. + **kwargs: Additional arguments. Expected: 'session' for OIDC authentication. """ raise NotImplementedError("Must be implemented in subclass.") - async def asearch(self, query, results_count: int = 4) -> RAGWebResults: + async def asearch(self, query: str, results_count: int = 4, **kwargs) -> RAGWebResults: """ - Search the collection for the given query. + Search the collection for the given query asynchronously. + + Args: + query: The search query string. + results_count: Number of results to return. + **kwargs: Additional arguments. Expected: 'session' for OIDC authentication. """ - return await sync_to_async(self.search)(query=query, results_count=results_count) + return await sync_to_async(self.search)(query=query, results_count=results_count, **kwargs) @classmethod @contextmanager diff --git a/src/backend/chat/agent_rag/document_rag_backends/find_rag_backend.py b/src/backend/chat/agent_rag/document_rag_backends/find_rag_backend.py new file mode 100644 index 0000000..a3e3e8a --- /dev/null +++ b/src/backend/chat/agent_rag/document_rag_backends/find_rag_backend.py @@ -0,0 +1,217 @@ +"""Implementation of the Find API for RAG document search.""" + +import logging +from io import BytesIO +from multiprocessing.context import AuthenticationError +from typing import List, Optional +from urllib.parse import urljoin +from uuid import uuid4 + +from django.conf import settings +from django.core.exceptions import ImproperlyConfigured +from django.utils import timezone + +import requests + +from chat.agent_rag.constants import RAGWebResult, RAGWebResults, RAGWebUsage +from chat.agent_rag.document_converter.markitdown import DocumentConverter +from chat.agent_rag.document_rag_backends.base_rag_backend import BaseRagBackend +from utils.oicd import refresh_access_token + +logger = logging.getLogger(__name__) + + +class FindRagBackend(BaseRagBackend): # pylint: disable=too-many-instance-attributes + """ + This class is a placeholder for the Find API implementation. + It is designed to be used with the RAG (Retrieval-Augmented Generation) document search system. + + It provides methods to: + - Parse documents and convert them to Markdown format: + + Handle PDF parsing using the Albert API. + + Use the DocumentConverter (markitdown) for other formats. + - Store parsed documents in the Find index. + - Perform a search operation using the Find API. + """ + + def __init__( + self, + collection_id: Optional[str] = None, + read_only_collection_id: Optional[List[str]] = None, + ): + # Initialize any necessary parameters or configurations here + super().__init__(collection_id, read_only_collection_id) + self._pdf_parser_endpoint = urljoin(settings.ALBERT_API_URL, "/v1/parse-beta") + self.api_key = settings.FIND_API_KEY + self.search_endpoint = "api/v1.0/documents/search/" + self.indexing_endpoint = "api/v1.0/documents/index/" + + if not self.api_key: + raise ImproperlyConfigured("FIND_API_KEY must be set in Django settings.") + + def create_collection(self, name: str, description: Optional[str] = None) -> str: + """ + init collection_id + """ + self.collection_id = self.collection_id or 1 + return self.collection_id + + # TODO + def delete_collection(self) -> None: + """ + Deletion not available + """ + logger.warning("deletion of collections is not yet supported in FindRagBackend") + + # TODO: factor with albert api + def parse_pdf_document(self, name: str, content_type: str, content: BytesIO) -> str: + """ + Parse the PDF document content and return the text content. + This method should handle the logic to convert the PDF into + a format suitable for the Albert API. + """ + response = requests.post( + self._pdf_parser_endpoint, + headers={ + "Authorization": f"Bearer {settings.ALBERT_API_KEY}", + }, + files={ + "file": ( + name, + content, + content_type, + ), # Use the name as the filename in the request + "output_format": (None, "markdown"), # Specify the output format as Markdown, + }, + timeout=settings.ALBERT_API_PARSE_TIMEOUT, + ) + response.raise_for_status() + + return "\n\n".join( + document_page["content"] for document_page in response.json().get("data", []) + ) + + # TODO: factor with albert api + def parse_document(self, name: str, content_type: str, content: BytesIO): + """ + Parse the document and prepare it for the search operation. + This method should handle the logic to convert the document + into a format suitable for the Find API. + + Args: + name (str): The name of the document. + content_type (str): The MIME type of the document (e.g., "application/pdf"). + content (BytesIO): The content of the document as a BytesIO stream. + + Returns: + str: The document content in Markdown format. + """ + # Implement the parsing logic here + if content_type == "application/pdf": + # Handle PDF parsing + markdown_content = self.parse_pdf_document( + name=name, content_type=content_type, content=content + ) + else: + markdown_content = DocumentConverter().convert_raw( + name=name, content_type=content_type, content=content + ) + + return markdown_content + + def store_document(self, name: str, content: str) -> None: + """ + index document in Find + + Args: + name (str): The name of the document. + content (str): The content of the document in Markdown format. + """ + logger.debug("index document '%s' in Find", name) + + user_sub = kwargs.get("user_sub") + if not user_sub: + raise ValueError("user_sub is required to store document in FindRagBackend") + + response = requests.post( + urljoin(settings.FIND_API_URL, self.indexing_endpoint), + headers={"Authorization": f"Bearer {self.api_key}"}, + json={ + "id": str(uuid4()), + "title": str(name) or "", + "depth": 0, + "path": str(name) or "", + "numchild": 0, + "content": content or "", + "created_at": timezone.now().isoformat(), + "updated_at": timezone.now().isoformat(), + "tags": [f"collection-{self.collection_id}"], + "size": len(content.encode("utf-8")), + "users": [], # TODO + "groups": [], + "reach": "public", + "is_active": True, + }, + timeout=settings.ALBERT_API_TIMEOUT, + ) + response.raise_for_status() + + def search(self, query: str, results_count: int = 4, **kwargs) -> RAGWebResults: + """ + Perform a search using the Find API. + Uses the user's OIDC token from the request session. + + Args: + query: The search query. + results_count: Number of results to return. + **kwargs: Additional arguments. Expected: 'session' containing OIDC tokens. + + Returns: + RAGWebResults: The search results. + """ + logger.debug("search documents in Find with query '%s'", query) + # TODO: factor session auth in a decorator + session = refresh_access_token(kwargs.get("session")) + oidc_access_token = session.get("oidc_access_token") + if not oidc_access_token: + raise AuthenticationError({"error": "Not authenticated"}) + collection_ids = self.get_all_collection_ids() + + response = requests.post( + urljoin(settings.FIND_API_URL, self.search_endpoint), + headers={"Authorization": f"Bearer {oidc_access_token}"}, + json={ + "q": query, + "tags": [f"collection:{collection_id}" for collection_id in collection_ids], + "k": 10, + }, + timeout=settings.ALBERT_API_TIMEOUT, + ) + response.raise_for_status() + + return RAGWebResults( + data=[ + RAGWebResult( + url=get_language_value(result["_source"], "title"), + content=get_language_value(result["_source"], "content"), + score=result["_score"], + ) + for result in response.json() + ], + usage=RAGWebUsage( + prompt_tokens=0, + completion_tokens=0, + ), + ) + + +def get_language_value(source, language_field): + """ + extract the value of the language field with the correct language_code extension. + "title" and "content" have extensions like "title.en" or "title.fr". + get_language_value will return the value regardless of the extension. + """ + for language_code in SUPPORTED_LANGUAGE_CODES: + if f"{language_field}.{language_code}" in source: + return source[f"{language_field}.{language_code}"] + raise ValueError(f"No '{language_field}' field with any supported language code in object") diff --git a/src/backend/chat/clients/pydantic_ai.py b/src/backend/chat/clients/pydantic_ai.py index dfab669..f84b34b 100644 --- a/src/backend/chat/clients/pydantic_ai.py +++ b/src/backend/chat/clients/pydantic_ai.py @@ -92,6 +92,7 @@ class ContextDeps: conversation: models.ChatConversation user: User + session: Optional[Dict] = None web_search_enabled: bool = False @@ -106,7 +107,14 @@ def get_model_configuration(model_hrid: str): class AIAgentService: # pylint: disable=too-many-instance-attributes """Service class for AI-related operations (Pydantic-AI edition).""" - def __init__(self, conversation: models.ChatConversation, user, model_hrid=None, language=None): + def __init__( # noqa: PLR0913 # pylint: disable=too-many-arguments,too-many-positional-arguments + self, + conversation: models.ChatConversation, + user, + session=None, + model_hrid=None, + language=None, + ): """ Initialize the AI agent service. @@ -136,6 +144,7 @@ class AIAgentService: # pylint: disable=too-many-instance-attributes self._context_deps = ContextDeps( conversation=conversation, user=user, + session=session, web_search_enabled=self._is_web_search_enabled, ) diff --git a/src/backend/chat/tests/tools/test_web_search_brave.py b/src/backend/chat/tests/tools/test_web_search_brave.py index dc45761..debcfa2 100644 --- a/src/backend/chat/tests/tools/test_web_search_brave.py +++ b/src/backend/chat/tests/tools/test_web_search_brave.py @@ -1,14 +1,19 @@ """Tests for the Brave web search tool.""" -# pylint: disable=too-many-lines +from typing import Sequence +# pylint: disable=too-many-lines from unittest.mock import AsyncMock, MagicMock, Mock, patch from urllib.parse import parse_qs +from django.contrib.auth.base_user import AbstractBaseUser +from django.contrib.sessions.backends.cache import SessionStore + import httpx import pytest import respx from pydantic_ai import ModelRetry, RunContext, RunUsage +from pydantic_ai._run_context import RunContextAgentDepsT from chat.tools.exceptions import ModelCannotRetry from chat.tools.web_search_brave import ( @@ -38,9 +43,6 @@ def brave_settings(settings): settings.BRAVE_SEARCH_EXTRA_SNIPPETS = True settings.BRAVE_SUMMARIZATION_ENABLED = False settings.BRAVE_CACHE_TTL = 3600 - settings.RAG_DOCUMENT_SEARCH_BACKEND = ( - "chat.agent_rag.document_rag_backends.albert_rag_backend.AlbertRagBackend" - ) settings.BRAVE_RAG_WEB_SEARCH_CHUNK_NUMBER = 5 @@ -48,6 +50,13 @@ def brave_settings(settings): def fixture_mocked_context(): """Fixture for a mocked RunContext.""" mock_ctx = Mock(spec=RunContext) + mock_ctx.deps = Mock(spec=RunContextAgentDepsT) + user = Mock(spec=AbstractBaseUser) + user.sub = Mock(spec=Sequence) + mock_ctx.deps.user = user + session = SessionStore() + session["oidc_access_token"] = "mocked-access-token" + mock_ctx.deps.session = session mock_ctx.usage = RunUsage(input_tokens=0, output_tokens=0) mock_ctx.max_retries = 2 mock_ctx.retries = {} @@ -1011,7 +1020,12 @@ async def test_web_search_brave_with_document_backend_rag_search_params(mocked_c await web_search_brave_with_document_backend(mocked_context, "test query") # Verify RAG search was called with correct parameters - mock_document_store.asearch.assert_called_once_with("test query", results_count=5) + mock_document_store.asearch.assert_called_once_with( + query="test query", + results_count=5, + session=mocked_context.deps.session, + user_sub=mocked_context.deps.user.sub, + ) @pytest.mark.asyncio diff --git a/src/backend/chat/tests/views/chat/conftest.py b/src/backend/chat/tests/views/chat/conftest.py new file mode 100644 index 0000000..d4be3b6 --- /dev/null +++ b/src/backend/chat/tests/views/chat/conftest.py @@ -0,0 +1,17 @@ +"""Common test fixtures for chat views tests.""" + +from unittest import mock + +import pytest + + +@pytest.fixture(autouse=True) +def mock_process_request(): + """ + Mock process_request to bypass OIDC authentication in tests. + """ + with mock.patch( + "lasuite.oidc_login.decorators.RefreshOIDCAccessToken.process_request" + ) as mocked_process_request: + mocked_process_request.return_value = None + yield mocked_process_request diff --git a/src/backend/chat/tests/views/chat/conversations/test_conversation_with_document_upload.py b/src/backend/chat/tests/views/chat/conversations/test_conversation_with_document_upload.py index eea77fa..f0941ab 100644 --- a/src/backend/chat/tests/views/chat/conversations/test_conversation_with_document_upload.py +++ b/src/backend/chat/tests/views/chat/conversations/test_conversation_with_document_upload.py @@ -8,6 +8,7 @@ import logging from io import BytesIO from unittest import mock +from django.contrib.sessions.backends.cache import SessionStore from django.utils import formats, timezone import httpx @@ -41,28 +42,49 @@ from chat.tests.utils import replace_uuids_with_placeholder pytestmark = pytest.mark.django_db(transaction=True) -@pytest.fixture(autouse=True) -def ai_settings(settings): +@pytest.fixture( + autouse=True, + params=[ + "chat.agent_rag.document_rag_backends.find_rag_backend.FindRagBackend", + "chat.agent_rag.document_rag_backends.albert_rag_backend.AlbertRagBackend", + ], +) +def ai_settings(request, settings): """Fixture to set AI service URLs for testing.""" - settings.AI_BASE_URL = "https://www.external-ai-service.com/" - settings.AI_API_KEY = "test-api-key" - settings.AI_MODEL = "test-model" - settings.AI_AGENT_INSTRUCTIONS = "You are a helpful test assistant :)" - # Enable Albert API for document search - settings.RAG_DOCUMENT_SEARCH_BACKEND = ( - "chat.agent_rag.document_rag_backends.albert_rag_backend.AlbertRagBackend" - ) - settings.ALBERT_API_URL = "https://albert.api.etalab.gouv.fr" - settings.ALBERT_API_KEY = "albert-api-key" + # enable on rag document search tool + settings.RAG_DOCUMENT_SEARCH_BACKEND = request.param settings.RAG_WEB_SEARCH_PROMPT_UPDATE = ( "Based on the following document contents:\n\n{search_results}\n\n" "Please answer the user's question: {user_prompt}" ) + settings.AI_BASE_URL = "https://www.external-ai-service.com/" + settings.AI_API_KEY = "test-api-key" + settings.AI_MODEL = "test-model" + settings.AI_AGENT_INSTRUCTIONS = "You are a helpful test assistant :)" + + # Albert API settings + settings.ALBERT_API_URL = "https://albert.api.etalab.gouv.fr" + settings.ALBERT_API_KEY = "albert-api-key" + + # Find API settings + settings.FIND_API_URL = "https://find.api.example.com" + settings.FIND_API_KEY = "find-api-key" + return settings +@pytest.fixture(autouse=True) +def mock_refresh_access_token(): + """Mock refresh_access_token to bypass token refresh in tests.""" + with mock.patch("utils.oicd.refresh_access_token") as mocked_refresh_access_token: + mock_session = Mock(spec=httpx.Client) + mock_session.access_token = "mocked-access-token" + mocked_refresh_access_token.return_value = mock_session + yield mocked_refresh_access_token + + @pytest.fixture(name="sample_pdf_content") def fixture_sample_pdf_content(): """Create a dummy PDF content as BytesIO.""" @@ -134,6 +156,32 @@ def fixture_mock_albert_api(): ) +@pytest.fixture(name="mock_find_api") +def fixture_mock_find_api(): + """Fixture to mock the Find API endpoints.""" + # Mock document indexing (Find API) + responses.post( + "https://find.api.example.com/api/v1.0/documents/index/", + json={"id": "456", "status": "indexed"}, + status=status.HTTP_200_OK, + ) + + # Mock document search (Find API) + responses.post( + "https://find.api.example.com/api/v1.0/documents/search/", + json=[ + { + "_source": { + "title.fr": "sample.pdf", + "content.fr": "This is the content of the PDF.", + }, + "_score": 0.9, + } + ], + status=status.HTTP_200_OK, + ) + + @pytest.fixture(name="mock_summarization_agent") def fixture_mock_summarization_agent(): """Mock the SummarizationAgent to return a fixed summary.""" @@ -220,6 +268,7 @@ def test_post_conversation_with_document_upload( # pylint: disable=too-many-arguments,too-many-positional-arguments api_client, mock_albert_api, # pylint: disable=unused-argument + mock_find_api, # pylint: disable=unused-argument sample_pdf_content, today_promt_date, mock_ai_agent_service, @@ -549,6 +598,7 @@ def test_post_conversation_with_document_upload_feature_disabled( def test_post_conversation_with_document_upload_summarize( # pylint: disable=too-many-arguments,too-many-positional-arguments # noqa: PLR0913 api_client, mock_albert_api, # pylint: disable=unused-argument + mock_find_api, # pylint: disable=unused-argument sample_pdf_content, today_promt_date, mock_ai_agent_service, diff --git a/src/backend/chat/tools/document_search_rag.py b/src/backend/chat/tools/document_search_rag.py index 46827fa..92e311a 100644 --- a/src/backend/chat/tools/document_search_rag.py +++ b/src/backend/chat/tools/document_search_rag.py @@ -26,7 +26,7 @@ def add_document_rag_search_tool(agent: Agent) -> None: document_store = document_store_backend(ctx.deps.conversation.collection_id) - rag_results = document_store.search(query) + rag_results = document_store.search(query, session=ctx.deps.session) ctx.usage += RunUsage( input_tokens=rag_results.usage.prompt_tokens, diff --git a/src/backend/chat/views.py b/src/backend/chat/views.py index 1fd67e5..e2ec47d 100644 --- a/src/backend/chat/views.py +++ b/src/backend/chat/views.py @@ -8,11 +8,13 @@ from django.conf import settings from django.core.cache import cache from django.core.files.storage import default_storage from django.http import Http404, StreamingHttpResponse +from django.utils.decorators import method_decorator import langfuse import magic import posthog from lasuite.malware_detection import malware_detection +from lasuite.oidc_login.decorators import refresh_oidc_access_token from rest_framework import decorators, filters, mixins, permissions, status, viewsets from rest_framework.exceptions import MethodNotAllowed, PermissionDenied, ValidationError from rest_framework.response import Response @@ -126,6 +128,7 @@ class ChatViewSet( # pylint: disable=too-many-ancestors, abstract-method self.permission_classes = [] return super().get_permissions() + @method_decorator(refresh_oidc_access_token) @decorators.action( methods=["post"], detail=True, @@ -177,6 +180,7 @@ class ChatViewSet( # pylint: disable=too-many-ancestors, abstract-method ai_service = AIAgentService( conversation=conversation, user=self.request.user, + session=request.session, model_hrid=model_hrid, language=( self.request.user.language diff --git a/src/backend/conftest.py b/src/backend/conftest.py index 81fa0bc..371b69c 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -22,7 +22,7 @@ def no_http_requests(monkeypatch): Credits: https://blog.jerrycodes.com/no-http-requests/ """ - allowed_hosts = {"localhost", "minio", "minio:9000"} + allowed_hosts = {"localhost", "127.0.0.1", "minio", "minio:9000"} original_urlopen = HTTPConnectionPool.urlopen def urlopen_mock(self, method, url, *args, **kwargs): diff --git a/src/backend/conversations/settings.py b/src/backend/conversations/settings.py index 2298492..c416f59 100755 --- a/src/backend/conversations/settings.py +++ b/src/backend/conversations/settings.py @@ -867,6 +867,23 @@ USER QUESTION: environ_prefix=None, ) + # Find + FIND_API_KEY = values.Value( + None, + environ_name="FIND_API_KEY", + environ_prefix=None, + ) + FIND_API_URL = values.Value( + "https://app-find/api", + environ_name="FIND_API_URL", + environ_prefix=None, + ) + FIND_API_TIMEOUT = values.PositiveIntegerValue( + default=30, # seconds + environ_name="FIND_API_TIMEOUT", + environ_prefix=None, + ) + # Logging # We want to make it easy to log to console but by default we log production # to Sentry and don't want to log to console. diff --git a/src/backend/utils/oidc.py b/src/backend/utils/oidc.py new file mode 100644 index 0000000..aed4901 --- /dev/null +++ b/src/backend/utils/oidc.py @@ -0,0 +1,54 @@ +"""Utility functions for OIDC token management.""" + +from functools import wraps + +from django.conf import settings + +import requests +from lasuite.oidc_login.backends import get_oidc_refresh_token, store_tokens +from rest_framework.exceptions import AuthenticationFailed + + +def refresh_access_token(session): + """Refresh the OIDC access token using the refresh token.""" + refresh_token = get_oidc_refresh_token(session) + if not refresh_token: + raise AuthenticationFailed({"error": "Refresh token is missing from session"}) + + response = requests.post( + settings.OIDC_OP_TOKEN_ENDPOINT, + data={ + "grant_type": "refresh_token", + "client_id": settings.OIDC_RP_CLIENT_ID, + "client_secret": settings.OIDC_RP_CLIENT_SECRET, + "refresh_token": refresh_token, + }, + timeout=5, + ) + response.raise_for_status() + token_info = response.json() + + store_tokens( + session, + access_token=token_info.get("access_token"), + id_token=None, + refresh_token=token_info.get("refresh_token"), + ) + return session + + +def with_fresh_access_token(func): + """ + Decorator to handle OIDC token refresh and extraction. + Expects 'session' in kwargs and update it with the fresh token. + """ + + @wraps(func) + def wrapper(*args, **kwargs): + session = kwargs.pop("session", None) + if session is None: + raise AuthenticationFailed({"error": "Session is required but not provided"}) + refreshed_session = refresh_access_token(session) + return func(*args, session=refreshed_session, **kwargs) + + return wrapper