diff --git a/.env.gui b/.env.gui index e990ca75..d3c30835 100644 --- a/.env.gui +++ b/.env.gui @@ -1,4 +1,4 @@ RELEASE=gui VERSION=1 BUILD=1 -FIX=1 +FIX=7 diff --git a/.env.llm_orchestration_service b/.env.llm_orchestration_service index 0493ed05..46fe9c27 100644 --- a/.env.llm_orchestration_service +++ b/.env.llm_orchestration_service @@ -1,4 +1,4 @@ RELEASE=orchestration VERSION=1 BUILD=1 -FIX=1 +FIX=10 diff --git a/.env.notification b/.env.notification new file mode 100644 index 00000000..89662377 --- /dev/null +++ b/.env.notification @@ -0,0 +1,5 @@ +RELEASE=notification-server +VERSION=1 +BUILD=1 +FIX=8 + diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 00000000..f7e2def6 --- /dev/null +++ b/.gitattributes @@ -0,0 +1,2 @@ +* text=auto eol=lf +*.sh text eol=lf diff --git a/.github/workflows/check-version.yml b/.github/workflows/check-version.yml new file mode 100644 index 00000000..3b038bdd --- /dev/null +++ b/.github/workflows/check-version.yml @@ -0,0 +1,78 @@ +name: Check Version + +on: + push: + branches: ["main"] + workflow_dispatch: + +env: + BRANCH: ${{ github.head_ref || github.ref_name }} + +jobs: + build: + if: "!contains(github.event.head_commit.message, '[skip ci]')" + runs-on: ubuntu-latest + steps: + - name: Checkout + uses: actions/checkout@v3 + with: + fetch-depth: 0 + persist-credentials: false + + - name: Docker Setup BuildX + uses: docker/setup-buildx-action@v2 + + - name: Generate Changelog + run: npm run changelog + + - name: Load environment variables + run: | + awk -v branch="${{ env.BRANCH }}" ' /^[0-9a-zA-Z]+$/ { current_branch = $0; } current_branch == branch && /^[A-Z_]+=/{ print $0; }' release.env >> $GITHUB_ENV + + - name: Set repo + run: | + LOWER_CASE_GITHUB_REPOSITORY=$(echo $GITHUB_REPOSITORY | tr '[:upper:]' '[:lower:]') + echo "DOCKER_TAG_CUSTOM=ghcr.io/${LOWER_CASE_GITHUB_REPOSITORY}:v${{ env.MAJOR }}.${{ env.MINOR }}.${{ env.PATCH }}" >> $GITHUB_ENV + echo "$GITHUB_ENV" + + - name: Build llm_orchestration_service image + run: | + echo "Building Docker image for branch: ${{ env.BRANCH }} major: ${{ env.MAJOR }} minor: ${{ env.MINOR }} patch: ${{ env.PATCH }}" + docker image build --tag $DOCKER_TAG_CUSTOM -f Dockerfile.llm_orchestration_service --no-cache . + + - name: Build GUI image + run: | + echo "Building Docker image for branch: ${{ env.BRANCH }} major: ${{ env.MAJOR }} minor: ${{ env.MINOR }} patch: ${{ env.PATCH }}" + cd GUI && docker image build --tag $DOCKER_TAG_CUSTOM-gui -f Dockerfile.dev --no-cache . + + - name: Build notification-server image + run: | + echo "Building Docker image for branch: ${{ env.BRANCH }} major: ${{ env.MAJOR }} minor: ${{ env.MINOR }} patch: ${{ env.PATCH }}" + cd notification-server && docker image build --tag $DOCKER_TAG_CUSTOM-notification-server --no-cache . + + - name: Build vault-init image + run: | + echo "Building Docker image for branch: ${{ env.BRANCH }} major: ${{ env.MAJOR }} minor: ${{ env.MINOR }} patch: ${{ env.PATCH }}" + docker image build --tag $DOCKER_TAG_CUSTOM-vault-init -f Dockerfile.vault-init --no-cache . + + - name: Log in to GitHub container registry + run: echo "${{ secrets.GITHUB_TOKEN }}" | docker login ghcr.io -u $ --password-stdin + + - name: Push llm_orchestration_service image to GitHub Packages + run: docker push $DOCKER_TAG_CUSTOM + + - name: Push GUI image to GitHub Packages + run: docker push $DOCKER_TAG_CUSTOM-gui + + - name: Push notification-server image to GitHub Packages + run: docker push $DOCKER_TAG_CUSTOM-notification-server + + - name: Push vault-init image to GitHub Packages + run: docker push $DOCKER_TAG_CUSTOM-vault-init + + - name: Create Release + uses: softprops/action-gh-release@v1 + with: + tag_name: v${{ env.MAJOR }}.${{ env.MINOR }}.${{ env.PATCH }} + generate_release_notes: true + body_path: ${{ github.workspace }}/CHANGELOG.md diff --git a/.github/workflows/ci-build-image-llm-orchestration-service.yml b/.github/workflows/ci-build-image-llm-orchestration-service.yml index 77cf95c2..29c971bc 100644 --- a/.github/workflows/ci-build-image-llm-orchestration-service.yml +++ b/.github/workflows/ci-build-image-llm-orchestration-service.yml @@ -3,7 +3,7 @@ name: Build and publish llm_orchestration_service on: push: branches: - - wip + - dev paths: - '.env.llm_orchestration_service' @@ -33,7 +33,7 @@ jobs: echo "$GITHUB_ENV" - name: Docker Build run: | - docker image build --tag $DOCKER_TAG_CUSTOM -f Dockerfile.llm_orchestration_service . + docker image build --tag $DOCKER_TAG_CUSTOM -f Dockerfile.llm_orchestration_service --no-cache . - name: Log in to GitHub container registry run: echo "${{ secrets.GITHUB_TOKEN }}" | docker login ghcr.io -u $ --password-stdin diff --git a/.github/workflows/ci-build-image-notification.yml b/.github/workflows/ci-build-image-notification.yml new file mode 100644 index 00000000..2e0463f0 --- /dev/null +++ b/.github/workflows/ci-build-image-notification.yml @@ -0,0 +1,42 @@ +name: Build and publish notification + +on: + push: + branches: + - dev + paths: + - '.env.notification' + +jobs: + PackageDeploy: + runs-on: ubuntu-22.04 + + steps: + - uses: actions/checkout@v2 + + - name: Docker Setup BuildX + uses: docker/setup-buildx-action@v2 + + - name: Load environment variables and set them + run: | + if [ -f .env.notification ]; then + export $(cat .env.notification | grep -v '^#' | xargs) + fi + echo "RELEASE=$RELEASE" >> $GITHUB_ENV + echo "VERSION=$VERSION" >> $GITHUB_ENV + echo "BUILD=$BUILD" >> $GITHUB_ENV + echo "FIX=$FIX" >> $GITHUB_ENV + - name: Set repo + run: | + LOWER_CASE_GITHUB_REPOSITORY=$(echo $GITHUB_REPOSITORY | tr '[:upper:]' '[:lower:]') + echo "DOCKER_TAG_CUSTOM=ghcr.io/${LOWER_CASE_GITHUB_REPOSITORY}:$RELEASE-$VERSION.$BUILD.$FIX" >> $GITHUB_ENV + echo "$GITHUB_ENV" + - name: Docker Build + run: | + cd notification-server && docker image build --tag $DOCKER_TAG_CUSTOM -f Dockerfile . + + - name: Log in to GitHub container registry + run: echo "${{ secrets.GITHUB_TOKEN }}" | docker login ghcr.io -u $ --password-stdin + + - name: Push Docker image to ghcr + run: docker push $DOCKER_TAG_CUSTOM diff --git a/.github/workflows/ci-build-image.yml b/.github/workflows/ci-build-image.yml index b8588cbf..0fa30b0d 100644 --- a/.github/workflows/ci-build-image.yml +++ b/.github/workflows/ci-build-image.yml @@ -3,7 +3,7 @@ name: Build and publish GUI on: push: branches: - - wip + - dev paths: - '.env.gui' diff --git a/.github/workflows/deepeval-tests.yml b/.github/workflows/deepeval-tests.yml index 5da84df9..5ed338ac 100644 --- a/.github/workflows/deepeval-tests.yml +++ b/.github/workflows/deepeval-tests.yml @@ -3,63 +3,265 @@ name: DeepEval RAG System Tests on: pull_request: types: [opened, synchronize, reopened] + branches: ["wip-eval"] paths: - 'src/**' - 'tests/**' + - 'data/**' + - 'docker-compose-eval.yml' + - 'Dockerfile.llm_orchestration_service' - '.github/workflows/deepeval-tests.yml' jobs: deepeval-tests: runs-on: ubuntu-latest - timeout-minutes: 40 + timeout-minutes: 80 steps: - name: Checkout code uses: actions/checkout@v4 - + + - name: Validate required secrets + id: validate_secrets + run: | + echo "validating required environment variables..." + MISSING_SECRETS=() + + # Check Azure OpenAI secrets + if [ -z "${{ secrets.AZURE_OPENAI_ENDPOINT }}" ]; then + MISSING_SECRETS+=("AZURE_OPENAI_ENDPOINT") + fi + + if [ -z "${{ secrets.AZURE_OPENAI_API_KEY }}" ]; then + MISSING_SECRETS+=("AZURE_OPENAI_API_KEY") + fi + + if [ -z "${{ secrets.AZURE_OPENAI_DEPLOYMENT }}" ]; then + MISSING_SECRETS+=("AZURE_OPENAI_DEPLOYMENT") + fi + + if [ -z "${{ secrets.AZURE_OPENAI_EMBEDDING_DEPLOYMENT }}" ]; then + MISSING_SECRETS+=("AZURE_OPENAI_EMBEDDING_DEPLOYMENT") + fi + + if [ -z "${{ secrets.AZURE_OPENAI_DEEPEVAL_DEPLOYMENT }}" ]; then + MISSING_SECRETS+=("AZURE_OPENAI_DEEPEVAL_DEPLOYMENT") + fi + + + + if [ -z "${{ secrets.AZURE_STORAGE_CONNECTION_STRING }}" ]; then + MISSING_SECRETS+=("AZURE_STORAGE_CONNECTION_STRING") + fi + + if [ -z "${{ secrets.AZURE_STORAGE_CONTAINER_NAME }}" ]; then + MISSING_SECRETS+=("AZURE_STORAGE_CONTAINER_NAME") + fi + + if [ -z "${{ secrets.AZURE_STORAGE_BLOB_NAME }}" ]; then + MISSING_SECRETS+=("AZURE_STORAGE_BLOB_NAME") + fi + + if [ -z "${{ secrets.NEXTAUTH_SECRET }}" ]; then + MISSING_SECRETS+=("NEXTAUTH_SECRET") + fi + + if [ -z "${{ secrets.ENCRYPTION_KEY }}" ]; then + MISSING_SECRETS+=("ENCRYPTION_KEY") + fi + + if [ -z "${{ secrets.SALT }}" ]; then + MISSING_SECRETS+=("SALT") + fi + + + # If any secrets are missing, fail + if [ ${#MISSING_SECRETS[@]} -gt 0 ]; then + echo "missing=true" >> $GITHUB_OUTPUT + echo "secrets_list=${MISSING_SECRETS[*]}" >> $GITHUB_OUTPUT + echo " Missing required secrets: ${MISSING_SECRETS[*]}" + exit 1 + else + echo "missing=false" >> $GITHUB_OUTPUT + echo " All required secrets are configured" + fi + + - name: Comment PR with missing secrets error + if: failure() && steps.validate_secrets.outputs.missing == 'true' + uses: actions/github-script@v7 + with: + script: | + const missingSecrets = '${{ steps.validate_secrets.outputs.secrets_list }}'.split(' '); + const secretsList = missingSecrets.map(s => `- \`${s}\``).join('\n'); + + const comment = `## DeepEval Tests: Missing Required Secrets + + The DeepEval RAG system tests cannot run because the following GitHub secrets are not configured: + + ${secretsList} + + ### How to Fix + + 1. Go to **Settings** → **Secrets and variables** → **Actions** + 2. Add the missing secrets with the appropriate values: + + **Azure OpenAI Configuration:** + - \`AZURE_OPENAI_ENDPOINT\` - Your Azure OpenAI endpoint (e.g., \`https://your-resource.openai.azure.com/\`) + - \`AZURE_OPENAI_API_KEY\` - Your Azure OpenAI API key + - \`AZURE_OPENAI_DEPLOYMENT\` - Chat model deployment name (e.g., \`gpt-4o-mini\`) + - \`AZURE_OPENAI_EMBEDDING_DEPLOYMENT\` - Embedding model deployment name (e.g., \`text-embedding-3-large\`) + - \`AZURE_STORAGE_CONNECTION_STRING\` - Connection string for Azure Blob Storage + - \`AZURE_STORAGE_CONTAINER_NAME\` - Container name in Azure Blob Storage + - \`AZURE_STORAGE_BLOB_NAME\` - Blob name for dataset in Azure + - \`AZURE_OPENAI_DEEPEVAL_DEPLOYMENT\` - DeepEval model deployment name (e.g., \`gpt-4.1\`) + + 3. Re-run the workflow after adding the secrets + + ### Note + Tests will not run until all required secrets are configured. + + --- + *Workflow: ${context.workflow} | Run: [#${context.runNumber}](${context.payload.repository.html_url}/actions/runs/${context.runId})*`; + + // Find existing comment + const comments = await github.rest.issues.listComments({ + owner: context.repo.owner, + repo: context.repo.repo, + issue_number: context.issue.number + }); + + const existingComment = comments.data.find( + comment => comment.user.login === 'github-actions[bot]' && + comment.body.includes('DeepEval Tests: Missing Required Secrets') + ); + + if (existingComment) { + await github.rest.issues.updateComment({ + owner: context.repo.owner, + repo: context.repo.repo, + comment_id: existingComment.id, + body: comment + }); + } else { + await github.rest.issues.createComment({ + owner: context.repo.owner, + repo: context.repo.repo, + issue_number: context.issue.number, + body: comment + }); + } + - name: Set up Python + if: success() uses: actions/setup-python@v5 with: python-version-file: '.python-version' - + - name: Set up uv + if: success() uses: astral-sh/setup-uv@v6 - + - name: Install dependencies (locked) + if: success() run: uv sync --frozen - - - name: Run DeepEval tests + + - name: Create test directories with proper permissions + if: success() + run: | + mkdir -p test-vault/agents/llm + mkdir -p test-vault/agent-out + # Set ownership to current user and make writable + sudo chown -R $(id -u):$(id -g) test-vault + chmod -R 777 test-vault + # Ensure the agent-out directory is world-readable after writes + sudo chmod -R a+rwX test-vault/agent-out + + - name: Set up Deepeval with azure + if: success() + run: | + uv run deepeval set-azure-openai \ + --openai-endpoint "${{ secrets.AZURE_OPENAI_ENDPOINT }}" \ + --openai-api-key "${{ secrets.AZURE_OPENAI_API_KEY }}" \ + --deployment-name "${{ secrets.AZURE_OPENAI_DEPLOYMENT }}" \ + --openai-model-name "${{ secrets.AZURE_OPENAI_DEEPEVAL_DEPLOYMENT }}" \ + --openai-api-version="2024-12-01-preview" + + - name: Run DeepEval tests with testcontainers + if: success() id: run_tests + continue-on-error: true env: - ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} - OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} - run: uv run python -m pytest tests/deepeval_tests/standard_tests.py -v --tb=short - + # LLM API Keys + AZURE_OPENAI_DEEPEVAL_DEPLOYMENT: ${{ secrets.AZURE_OPENAI_DEEPEVAL_DEPLOYMENT }} + # Azure OpenAI - Chat Model + AZURE_OPENAI_API_KEY: ${{ secrets.AZURE_OPENAI_API_KEY }} + AZURE_OPENAI_ENDPOINT: ${{ secrets.AZURE_OPENAI_ENDPOINT }} + AZURE_OPENAI_DEPLOYMENT: ${{ secrets.AZURE_OPENAI_DEPLOYMENT }} + # Azure OpenAI - Embedding Model + AZURE_OPENAI_EMBEDDING_DEPLOYMENT: ${{ secrets.AZURE_OPENAI_EMBEDDING_DEPLOYMENT }} + # Evaluation mode + AZURE_STORAGE_CONNECTION_STRING: ${{ secrets.AZURE_STORAGE_CONNECTION_STRING }} + AZURE_STORAGE_CONTAINER_NAME: ${{ secrets.AZURE_STORAGE_CONTAINER_NAME }} + AZURE_STORAGE_BLOB_NAME: ${{ secrets.AZURE_STORAGE_BLOB_NAME }} + EVAL_MODE: "true" + # Langfuse auth secrets (required by docker-compose-eval.yml) + NEXTAUTH_SECRET: ${{ secrets.NEXTAUTH_SECRET }} + ENCRYPTION_KEY: ${{ secrets.ENCRYPTION_KEY }} + SALT: ${{ secrets.SALT }} + run: | + # Run tests sequentially (one at a time) to avoid rate limiting. + # standard_tests.py — RAG quality (DeepEval metrics on /orchestrate-eval) + # api_tool_tests.py — API Tool Calling scenarios (issue #447) + uv run python -m pytest \ + tests/deepeval_tests/standard_tests.py \ + tests/deepeval_tests/api_tool_tests.py \ + -v --tb=short --log-cli-level=INFO -n 0 + + - name: Fix permissions on test artifacts + if: always() + run: | + sudo chown -R $(id -u):$(id -g) test-vault || true + sudo chmod -R a+rX test-vault || true + - name: Generate evaluation report if: always() - run: python tests/deepeval_tests/report_generator.py - + run: uv run python tests/deepeval_tests/report_generator.py + + - name: Generate API tool evaluation report + if: always() + run: uv run python tests/deepeval_tests/api_tool_report_generator.py + + - name: Save test artifacts + if: always() + uses: actions/upload-artifact@v4 + with: + name: test-results + path: | + pytest_captured_results.json + test_report.md + api_tool_test_results.json + api_tool_test_report.md + retention-days: 30 + - name: Comment PR with test results if: always() && github.event_name == 'pull_request' uses: actions/github-script@v7 with: script: | const fs = require('fs'); - try { const reportContent = fs.readFileSync('test_report.md', 'utf8'); - const comments = await github.rest.issues.listComments({ owner: context.repo.owner, repo: context.repo.repo, issue_number: context.issue.number }); - + const existingComment = comments.data.find( comment => comment.user.login === 'github-actions[bot]' && - comment.body.includes('RAG System Evaluation Report') + comment.body.includes('RAG System Evaluation Report') ); - + if (existingComment) { await github.rest.issues.updateComment({ owner: context.repo.owner, @@ -75,10 +277,8 @@ jobs: body: reportContent }); } - } catch (error) { console.error('Failed to post test results:', error); - await github.rest.issues.createComment({ issue_number: context.issue.number, owner: context.repo.owner, @@ -86,25 +286,70 @@ jobs: body: `## RAG System Evaluation Report\n\n**Error generating test report**\n\nFailed to read or post test results. Check workflow logs for details.\n\nError: ${error.message}` }); } - + + - name: Comment PR with API tool test results + if: always() && github.event_name == 'pull_request' + uses: actions/github-script@v7 + with: + script: | + const fs = require('fs'); + try { + const reportContent = fs.readFileSync('api_tool_test_report.md', 'utf8'); + const comments = await github.rest.issues.listComments({ + owner: context.repo.owner, + repo: context.repo.repo, + issue_number: context.issue.number + }); + + const existingComment = comments.data.find( + comment => comment.user.login === 'github-actions[bot]' && + comment.body.includes('API Tool Calling Evaluation Report') + ); + + if (existingComment) { + await github.rest.issues.updateComment({ + owner: context.repo.owner, + repo: context.repo.repo, + comment_id: existingComment.id, + body: reportContent + }); + } else { + await github.rest.issues.createComment({ + owner: context.repo.owner, + repo: context.repo.repo, + issue_number: context.issue.number, + body: reportContent + }); + } + } catch (error) { + console.error('Failed to post API tool test results:', error); + await github.rest.issues.createComment({ + issue_number: context.issue.number, + owner: context.repo.owner, + repo: context.repo.repo, + body: `## API Tool Calling Evaluation Report\n\n**Error generating report**\n\nFailed to read or post API tool test results. Check workflow logs.\n\nError: ${error.message}` + }); + } + - name: Check test results and fail if needed if: always() run: | - # Check if pytest ran (look at step output) - if [ "${{ steps.run_tests.outcome }}" == "failure" ]; then + # Check if pytest ran (look at step output) + if [ "${{ steps.run_tests.outcome }}" == "failure" ]; then echo "Tests ran but failed - this is expected if RAG performance is below threshold" - fi - if [ -f "pytest_captured_results.json" ]; then + fi + + if [ -f "pytest_captured_results.json" ]; then total_tests=$(jq '.total_tests // 0' pytest_captured_results.json) passed_tests=$(jq '.passed_tests // 0' pytest_captured_results.json) - + if [ "$total_tests" -eq 0 ]; then echo "ERROR: No tests were executed" exit 1 fi - + pass_rate=$(awk "BEGIN {print ($passed_tests / $total_tests) * 100}") - + echo "DeepEval Test Results:" echo "Total Tests: $total_tests" echo "Passed Tests: $passed_tests" @@ -117,7 +362,13 @@ jobs: else echo "TEST SUCCESS: Pass rate $pass_rate% meets threshold 70%" fi - else + else echo "ERROR: No test results file found" exit 1 - fi \ No newline at end of file + fi + + - name: Cleanup Docker resources + if: always() + run: | + docker compose -f docker-compose-eval.yml down -v --remove-orphans || true + docker system prune -f || true \ No newline at end of file diff --git a/.github/workflows/deepteam-red-team-tests.yml b/.github/workflows/deepteam-red-team-tests.yml index ba0861be..08e499a8 100644 --- a/.github/workflows/deepteam-red-team-tests.yml +++ b/.github/workflows/deepteam-red-team-tests.yml @@ -3,12 +3,13 @@ name: DeepTeam Red Team Security Tests on: pull_request: types: [opened, synchronize, reopened] + branches: ["wip-eval"] paths: - 'src/**' - 'tests/**' - 'mocks/**' - 'data/**' - - '.github/workflows/deepeval-red-team-tests.yml' + - '.github/workflows/deepteam-red-team-tests.yml' workflow_dispatch: inputs: attack_intensity: @@ -41,12 +42,24 @@ jobs: - name: Install dependencies (locked) run: uv sync --frozen + - name: Set up DeepTeam with Azure OpenAI + run: | + uv run deepeval set-azure-openai \ + --openai-endpoint "${{ secrets.AZURE_OPENAI_ENDPOINT }}" \ + --openai-api-key "${{ secrets.AZURE_OPENAI_API_KEY }}" \ + --deployment-name "${{ secrets.AZURE_OPENAI_DEPLOYMENT }}" \ + --openai-model-name "${{ secrets.AZURE_OPENAI_DEEPEVAL_DEPLOYMENT }}" \ + --openai-api-version="2024-12-01-preview" + - name: Run Complete Security Assessment id: run_tests continue-on-error: true env: ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} - OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} + AZURE_OPENAI_API_KEY: ${{ secrets.AZURE_OPENAI_API_KEY }} + AZURE_OPENAI_ENDPOINT: ${{ secrets.AZURE_OPENAI_ENDPOINT }} + AZURE_OPENAI_DEPLOYMENT: ${{ secrets.AZURE_OPENAI_DEPLOYMENT }} + AZURE_OPENAI_DEEPEVAL_DEPLOYMENT: ${{ secrets.AZURE_OPENAI_DEEPEVAL_DEPLOYMENT }} run: | # Run all security tests in one comprehensive session uv run python -m pytest tests/deepeval_tests/red_team_tests.py::TestRAGSystemRedTeaming -v --tb=short diff --git a/.github/workflows/helm-dependency.yaml b/.github/workflows/helm-dependency.yaml index 92de5c51..39ba544c 100644 --- a/.github/workflows/helm-dependency.yaml +++ b/.github/workflows/helm-dependency.yaml @@ -35,7 +35,7 @@ jobs: git config user.email "github-actions[bot]@users.noreply.github.com" git add kubernetes/Chart.lock kubernetes/charts/ if ! git diff --cached --quiet; then - git commit -m "chore: update Helm dependencies (charts/ and Chart.lock)" + git commit -m "chore: update Helm dependencies (charts/ and Chart.lock) [skip ci]" git push else echo "No changes to Helm dependencies." diff --git a/.husky/commit-msg b/.husky/commit-msg new file mode 100644 index 00000000..8cab56ab --- /dev/null +++ b/.husky/commit-msg @@ -0,0 +1,24 @@ +# !/bin/bash + +message="$(head -1 $1)" + +message_pattern="^(feat|fix|chore|docs|refactor|style|test)\([0-9]+\):\ .+$" + +if ! [[ $message =~ $message_pattern ]]; +then + echo "---" + echo "Violation of commit message format!" + echo "The commit message must follow the Conventional Commits standard:" + echo "[type(scope): description]" + echo "Example: feat(100): Implement automated pipeline for code commits" + echo "Accepted types: feat, fix, chore, docs, refactor, style, test" + echo "(1) feat: Added a new feature" + echo "(2) fix: Fixed a bug" + echo "(3) chore: Added changes that do not relate to a fix or feature and don't modify src or test files (for example updating dependencies)" + echo "(4) docs: Added updates to documentation such as a the README or other markdown files" + echo "(5) refactor: Refactored code that neither fixes a bug nor adds a feature" + echo "(6) style: Added Changes that do not affect the meaning of the code, likely related to code formatting such as white-space, missing semi-colons, and so on" + echo "(7) test: Included new or corrected previous tests" + echo "---" + exit 1 +fi diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index a7a1de1e..2d3b29d5 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -212,10 +212,12 @@ All FastAPI route handlers use Pydantic models for request/response validation: from pydantic import BaseModel from fastapi import FastAPI + class UserRequest(BaseModel): name: str age: int + @app.post("/users") async def create_user(user: UserRequest): # Pydantic validates name is string, age is int diff --git a/DSL/CronManager/DSL/data_resync.yml b/DSL/CronManager/DSL/data_resync.yml index c5fb58d2..a232ba39 100644 --- a/DSL/CronManager/DSL/data_resync.yml +++ b/DSL/CronManager/DSL/data_resync.yml @@ -1,5 +1,5 @@ agency_data_resync: - # trigger: "0 0/1 * * * ?" - trigger: off + trigger: "0 0 0/1 * * ?" + # trigger: off type: exec - command: "../app/scripts/agency_data_resync.sh -s 10" \ No newline at end of file + command: "/app/scripts/agency_data_resync.sh -s 10" \ No newline at end of file diff --git a/DSL/CronManager/DSL/delete_from_vault.yml b/DSL/CronManager/DSL/delete_from_vault.yml index be209617..cde1df27 100644 --- a/DSL/CronManager/DSL/delete_from_vault.yml +++ b/DSL/CronManager/DSL/delete_from_vault.yml @@ -2,4 +2,4 @@ delete_secrets: trigger: off type: exec command: "/app/scripts/delete_secrets_from_vault.sh" - allowedEnvs: ['cookie', 'connectionId','llmPlatform', 'llmModel','embeddingModel','embeddingPlatform','deploymentEnvironment'] + allowedEnvs: ['cookie','vaultUuid','llmPlatform', 'llmModel','embeddingModel','embeddingPlatform', 'vaultAgentUrl'] diff --git a/DSL/CronManager/DSL/store_in_vault.yml b/DSL/CronManager/DSL/store_in_vault.yml index 30522190..46f861e6 100644 --- a/DSL/CronManager/DSL/store_in_vault.yml +++ b/DSL/CronManager/DSL/store_in_vault.yml @@ -2,4 +2,4 @@ store_secrets: trigger: off type: exec command: "/app/scripts/store_secrets_in_vault.sh" - allowedEnvs: ['cookie', 'connectionId','llmPlatform', 'llmModel','secretKey','accessKey','deploymentName','targetUrl','apiKey','embeddingModel','embeddingPlatform','embeddingAccessKey','embeddingSecretKey','embeddingDeploymentName','embeddingTargetUri','embeddingAzureApiKey','deploymentEnvironment'] \ No newline at end of file + allowedEnvs: ['cookie','vaultUuid','llmPlatform', 'llmModel','secretKey','accessKey','deploymentName','targetUrl','apiKey','embeddingModel','embeddingPlatform','embeddingAccessKey','embeddingSecretKey','embeddingDeploymentName','embeddingTargetUri','embeddingAzureApiKey','deploymentEnvironment', 'vaultAgentUrl'] \ No newline at end of file diff --git a/DSL/CronManager/script/api_tool_indexer.sh b/DSL/CronManager/script/api_tool_indexer.sh index ed1a27e2..5c1507c1 100644 --- a/DSL/CronManager/script/api_tool_indexer.sh +++ b/DSL/CronManager/script/api_tool_indexer.sh @@ -36,7 +36,7 @@ echo "[PACKAGES] Installing required packages..." "$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "httpx>=0.27.0" || exit 1 "$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "pydantic>=2.11.7" || exit 1 "$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "qdrant-client>=1.15.1" || exit 1 -"$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "loguru>=0.7.3" || exit 1 +"$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "requests>=2.32.5" || exit 1 echo "[PACKAGES] All packages installed successfully" @@ -64,6 +64,10 @@ if [ -n "$params" ]; then url_decode "$params" > "$PARAMS_FILE" fi +# Decode URL to restore path parameter templates that were +# URL-encoded by Ruuter to prevent Spring URI template expansion errors. +DECODED_URL=$(url_decode "$url") + # Build Python command arguments array PYTHON_ARGS=( "$PYTHON_SCRIPT" @@ -71,7 +75,7 @@ PYTHON_ARGS=( --service-id "${service_id:-""}" --name "$name" --description "$description" - --url "$url" + --url "$DECODED_URL" --method "${method:-"GET"}" --visibility "${visibility:-"public"}" --type "${type:-"custom_endpoint"}" diff --git a/DSL/CronManager/script/delete_secrets_from_vault.sh b/DSL/CronManager/script/delete_secrets_from_vault.sh index 056a634f..0e457e41 100644 --- a/DSL/CronManager/script/delete_secrets_from_vault.sh +++ b/DSL/CronManager/script/delete_secrets_from_vault.sh @@ -6,9 +6,18 @@ set -e # Exit on any error # Configuration -# Use VAULT_AGENT_URL which points to vault-agent-cron proxy -# The agent automatically injects the authentication token -VAULT_ADDR="${VAULT_AGENT_URL:-http://vault-agent-cron:8203}" +# Resolve Vault Agent URL: +# 1. Use vaultAgentUrl env var if set (from container env or CronManager request) +# 2. Auto-detect Kubernetes via KUBERNETES_SERVICE_HOST (injected by kubelet, cannot be disabled) +# 3. Auto-detect Kubernetes via service account token (mounted by default in every pod) +# 4. Fallback to Docker Compose hostname +if [ -n "$vaultAgentUrl" ]; then + VAULT_ADDR="$vaultAgentUrl" +elif [ -n "$KUBERNETES_SERVICE_HOST" ] || [ -f "/var/run/secrets/kubernetes.io/serviceaccount/token" ]; then + VAULT_ADDR="http://localhost:8203" +else + VAULT_ADDR="http://vault-agent-cron:8203" +fi # Logging function log() { @@ -19,14 +28,19 @@ log "=== Starting Vault Secrets Deletion ===" # Debug: Print received parameters log "Received parameters:" -log " connectionId: $connectionId" +log " vaultUuid: $vaultUuid" log " llmPlatform: $llmPlatform" log " llmModel: $llmModel" log " embeddingModel: $embeddingModel" log " embeddingPlatform: $embeddingPlatform" -log " deploymentEnvironment: $deploymentEnvironment" log " Vault Address: $VAULT_ADDR" +# Validate required vaultUuid parameter +if [ -z "$vaultUuid" ]; then + log "ERROR: vaultUuid is required but not provided" + exit 1 +fi + # Note: No token required - vault agent proxy automatically injects authentication # Function to determine platform name @@ -49,17 +63,13 @@ get_model_name() { echo "$model_array" | sed 's/\[//g' | sed 's/\]//g' | sed 's/"//g' | cut -d',' -f1 | xargs } -# Function to build vault path +# Function to build vault path (uses vaultUuid as stable path terminal) build_vault_path() { local secret_type=$1 # "llm" or "embeddings" local platform_name=$2 - local model_name=$3 - if [ "$deploymentEnvironment" = "testing" ]; then - echo "secret/$secret_type/connections/$platform_name/$deploymentEnvironment/$connectionId" - else - echo "secret/$secret_type/connections/$platform_name/$deploymentEnvironment/$model_name" - fi + # UUID-based path: no environment in path, swap is DB-only + echo "secret/$secret_type/connections/$platform_name/$vaultUuid" } # Function to delete vault secret (both data and metadata) @@ -125,27 +135,26 @@ delete_vault_secret() { # Function to delete LLM secrets delete_llm_secrets() { - if [ -z "$llmPlatform" ] || [ -z "$llmModel" ]; then - log "No LLM platform or model specified, skipping LLM secrets deletion" + if [ -z "$llmPlatform" ]; then + log "No LLM platform specified, skipping LLM secrets deletion" return 0 fi local platform_name=$(get_platform_name "$llmPlatform") - local model_name=$(get_model_name "$llmModel") - local vault_path=$(build_vault_path "llm" "$platform_name" "$model_name") + local vault_path=$(build_vault_path "llm" "$platform_name") delete_vault_secret "$vault_path" "LLM secrets" } # Function to delete embedding secrets delete_embedding_secrets() { - if [ -z "$embeddingPlatform" ] || [ -z "$embeddingModel" ]; then - log "No embedding platform or model specified, skipping embedding secrets deletion" + if [ -z "$embeddingPlatform" ]; then + log "No embedding platform specified, skipping embedding secrets deletion" return 0 fi local platform_name=$(get_platform_name "$embeddingPlatform") - local vault_path=$(build_vault_path "embeddings" "$platform_name" "$embeddingModel") + local vault_path=$(build_vault_path "embeddings" "$platform_name") delete_vault_secret "$vault_path" "Embedding secrets" } @@ -169,4 +178,4 @@ delete_llm_secrets # Delete embedding secrets delete_embedding_secrets -log "=== Vault secrets deletion completed ===" +log "=== Vault secrets deletion completed ===" \ No newline at end of file diff --git a/DSL/CronManager/script/service_enrichment.sh b/DSL/CronManager/script/service_enrichment.sh index c50a490a..c4bdb9d1 100644 --- a/DSL/CronManager/script/service_enrichment.sh +++ b/DSL/CronManager/script/service_enrichment.sh @@ -37,7 +37,7 @@ echo "[PACKAGES] Installing required packages..." "$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "httpx>=0.27.0" || exit 1 "$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "pydantic>=2.11.7" || exit 1 "$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "qdrant-client>=1.15.1" || exit 1 -"$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "loguru>=0.7.3" || exit 1 +"$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "requests>=2.32.5" || exit 1 echo "[PACKAGES] All packages installed successfully" diff --git a/DSL/CronManager/script/store_secrets_in_vault.sh b/DSL/CronManager/script/store_secrets_in_vault.sh index 2ec78d6f..b47b8a02 100644 --- a/DSL/CronManager/script/store_secrets_in_vault.sh +++ b/DSL/CronManager/script/store_secrets_in_vault.sh @@ -6,9 +6,20 @@ set -e # Exit on any error # Configuration -# Use VAULT_AGENT_URL which points to vault-agent-cron proxy -# The agent automatically injects the authentication token -VAULT_ADDR="${VAULT_AGENT_URL:-http://vault-agent-cron:8203}" +# Resolve Vault Agent URL: +# 1. Use vaultAgentUrl env var if set (from container env or CronManager request) +# 2. Auto-detect Kubernetes via KUBERNETES_SERVICE_HOST (injected by kubelet, cannot be disabled) +# 3. Auto-detect Kubernetes via service account token (mounted by default in every pod) +# 4. Fallback to Docker Compose hostname +if [ -n "$vaultAgentUrl" ]; then + VAULT_ADDR="$vaultAgentUrl" +elif [ -n "$KUBERNETES_SERVICE_HOST" ] || [ -f "/var/run/secrets/kubernetes.io/serviceaccount/token" ]; then + VAULT_ADDR="http://localhost:8203" +else + VAULT_ADDR="http://vault-agent-cron:8203" +fi + +echo "DEBUG: VAULT_ADDR=$VAULT_ADDR vaultAgentUrl=$vaultAgentUrl KUBERNETES_SERVICE_HOST=$KUBERNETES_SERVICE_HOST" # Decryption Configuration PRIVATE_KEY_CACHE="" @@ -154,7 +165,14 @@ setup_python_environment() { echo "[$(date '+%Y-%m-%d %H:%M:%S')] ERROR: Failed to install loguru" >&2 return 1 } - + + # Install requests, required by LokiLogger which decrypt_vault_secrets.py imports + echo "[$(date '+%Y-%m-%d %H:%M:%S')] Installing requests library..." + "$UV_BIN" pip install --python "$VENV_PATH/bin/python3" "requests>=2.32" 2>&1 || { + echo "[$(date '+%Y-%m-%d %H:%M:%S')] ERROR: Failed to install requests" >&2 + return 1 + } + # Mark setup as complete touch "$VENV_PATH/.setup_complete" @@ -164,12 +182,18 @@ setup_python_environment() { log "=== Starting Vault Secrets Storage ===" log "Received parameters:" -log " connectionId: $connectionId" +log " vaultUuid: $vaultUuid" log " llmPlatform: $llmPlatform" log " llmModel: $llmModel" log " deploymentEnvironment: $deploymentEnvironment" log " Vault Address: $VAULT_ADDR" +# Validate required vaultUuid parameter +if [ -z "$vaultUuid" ]; then + log "ERROR: vaultUuid is required but not provided" + exit 1 +fi + # Redirect stderr to stdout so cron-manager can capture all logs exec 2>&1 @@ -179,6 +203,9 @@ if ! setup_python_environment; then exit 1 fi +# Set Python path +export PYTHONPATH="/app:/app/src:/app/src/vector_indexer:$PYTHONPATH" + # Function to determine platform name get_platform_name() { case "$llmPlatform" in @@ -197,24 +224,13 @@ get_model_name() { echo "$llmModel" | sed 's/\[//g' | sed 's/\]//g' | sed 's/"//g' | cut -d',' -f1 | xargs } -# Function to build vault path +# Function to build vault path (uses vaultUuid as stable path terminal) build_vault_path() { local secret_type=$1 # "llm" or "embeddings" local platform=$(get_platform_name) - # Use appropriate model based on secret type - local model - if [ "$secret_type" = "embeddings" ]; then - model="$embeddingModel" - else - model=$(get_model_name) - fi - - if [ "$deploymentEnvironment" = "testing" ]; then - echo "secret/$secret_type/connections/$platform/$deploymentEnvironment/$connectionId" - else - echo "secret/$secret_type/connections/$platform/$deploymentEnvironment/$model" - fi + # UUID-based path: no environment in path, swap is DB-only + echo "secret/$secret_type/connections/$platform/$vaultUuid" } # Function to store LLM secrets @@ -272,18 +288,16 @@ store_aws_llm_secrets() { # Build JSON payload using jq for proper escaping local json_payload=$(jq -n \ - --arg conn_id "$connectionId" \ + --arg conn_id "$vaultUuid" \ --arg access_key "$decrypted_access_key" \ --arg secret_key "$decrypted_secret_key" \ - --arg env "$deploymentEnvironment" \ --arg model "$model" \ '{data: { connection_id: $conn_id, access_key: $access_key, secret_key: $secret_key, - environment: $env, model: $model, - tags: "aws,bedrock,\($env),\($model)" + tags: "aws,bedrock,\($model)" }}') log "Storing secrets at path: $vault_path" @@ -333,21 +347,19 @@ store_azure_llm_secrets() { # Build JSON payload using jq for proper escaping local json_payload=$(jq -n \ - --arg conn_id "$connectionId" \ + --arg conn_id "$vaultUuid" \ --arg endpoint "$targetUrl" \ --arg api_key "$decrypted_api_key" \ --arg deploy_name "$deploymentName" \ - --arg env "$deploymentEnvironment" \ --arg model "$model" \ '{data: { connection_id: $conn_id, endpoint: $endpoint, api_key: $api_key, deployment_name: $deploy_name, - environment: $env, model: $model, api_version: "2024-05-01-preview", - tags: "azure,\($env),\($model)" + tags: "azure,\($model)" }}') log "Storing secrets at path: $vault_path" @@ -402,18 +414,16 @@ store_aws_embedding_secrets() { # Build JSON payload using jq for proper escaping local json_payload=$(jq -n \ - --arg conn_id "$connectionId" \ + --arg conn_id "$vaultUuid" \ --arg access_key "$decrypted_embedding_access_key" \ --arg secret_key "$decrypted_embedding_secret_key" \ - --arg env "$deploymentEnvironment" \ --arg model "$embeddingModel" \ '{data: { connection_id: $conn_id, access_key: $access_key, secret_key: $secret_key, - environment: $env, model: $model, - tags: "aws,bedrock,embedding,\($env),\($model)" + tags: "aws,bedrock,embedding,\($model)" }}') log "Storing secrets at path: $vault_path" @@ -462,21 +472,19 @@ store_azure_embedding_secrets() { # Build JSON payload using jq for proper escaping local json_payload=$(jq -n \ - --arg conn_id "$connectionId" \ + --arg conn_id "$vaultUuid" \ --arg endpoint "$embeddingTargetUri" \ --arg api_key "$decrypted_embedding_api_key" \ --arg deploy_name "$embeddingDeploymentName" \ - --arg env "$deploymentEnvironment" \ --arg model "$embeddingModel" \ '{data: { connection_id: $conn_id, endpoint: $endpoint, api_key: $api_key, deployment_name: $deploy_name, - environment: $env, model: $model, api_version: "2024-12-01-preview", - tags: "azure,embedding,\($env),\($model)" + tags: "azure,embedding,\($model)" }}') log "Storing secrets at path: $vault_path" diff --git a/DSL/Liquibase/changelog/rag-search-script-v8-uuid-vault-path.sql b/DSL/Liquibase/changelog/rag-search-script-v8-uuid-vault-path.sql new file mode 100644 index 00000000..f5c3ae90 --- /dev/null +++ b/DSL/Liquibase/changelog/rag-search-script-v8-uuid-vault-path.sql @@ -0,0 +1,16 @@ +-- liquibase formatted sql + +-- changeset uuid-vault-path:rag-script-v8-changeset1 +-- Add UUID column to llm_connections for stable Vault secret paths +ALTER TABLE rag_search.llm_connections + ADD COLUMN IF NOT EXISTS vault_uuid UUID DEFAULT gen_random_uuid(); + +-- Backfill existing rows with unique UUIDs +UPDATE rag_search.llm_connections SET vault_uuid = gen_random_uuid() WHERE vault_uuid IS NULL; + +-- Make it NOT NULL after backfill +ALTER TABLE rag_search.llm_connections ALTER COLUMN vault_uuid SET NOT NULL; + +-- Add unique constraint +ALTER TABLE rag_search.llm_connections ADD CONSTRAINT llm_connections_vault_uuid_unique UNIQUE (vault_uuid); +-- rollback ALTER TABLE rag_search.llm_connections DROP CONSTRAINT IF EXISTS llm_connections_vault_uuid_unique; ALTER TABLE rag_search.llm_connections DROP COLUMN IF EXISTS vault_uuid; diff --git a/DSL/Liquibase/master.yml b/DSL/Liquibase/master.yml index 89474bf9..1bcae390 100644 --- a/DSL/Liquibase/master.yml +++ b/DSL/Liquibase/master.yml @@ -12,4 +12,6 @@ databaseChangeLog: - include: file: changelog/rag-search-script-v6-endpoints.sql - include: - file: changelog/rag-search-script-v7-schema-migration.sql \ No newline at end of file + file: changelog/rag-search-script-v7-schema-migration.sql + - include: + file: changelog/rag-search-script-v8-uuid-vault-path.sql \ No newline at end of file diff --git a/DSL/Liquibase_production/changelog/rag-search-script-v8-uuid-vault-path.sql b/DSL/Liquibase_production/changelog/rag-search-script-v8-uuid-vault-path.sql new file mode 100644 index 00000000..f5c3ae90 --- /dev/null +++ b/DSL/Liquibase_production/changelog/rag-search-script-v8-uuid-vault-path.sql @@ -0,0 +1,16 @@ +-- liquibase formatted sql + +-- changeset uuid-vault-path:rag-script-v8-changeset1 +-- Add UUID column to llm_connections for stable Vault secret paths +ALTER TABLE rag_search.llm_connections + ADD COLUMN IF NOT EXISTS vault_uuid UUID DEFAULT gen_random_uuid(); + +-- Backfill existing rows with unique UUIDs +UPDATE rag_search.llm_connections SET vault_uuid = gen_random_uuid() WHERE vault_uuid IS NULL; + +-- Make it NOT NULL after backfill +ALTER TABLE rag_search.llm_connections ALTER COLUMN vault_uuid SET NOT NULL; + +-- Add unique constraint +ALTER TABLE rag_search.llm_connections ADD CONSTRAINT llm_connections_vault_uuid_unique UNIQUE (vault_uuid); +-- rollback ALTER TABLE rag_search.llm_connections DROP CONSTRAINT IF EXISTS llm_connections_vault_uuid_unique; ALTER TABLE rag_search.llm_connections DROP COLUMN IF EXISTS vault_uuid; diff --git a/DSL/Resql/rag-search/POST/deactivate-llm-connection-budget-exceed.sql b/DSL/Resql/rag-search/POST/deactivate-llm-connection-budget-exceed.sql index 809a2560..27e7d29b 100644 --- a/DSL/Resql/rag-search/POST/deactivate-llm-connection-budget-exceed.sql +++ b/DSL/Resql/rag-search/POST/deactivate-llm-connection-budget-exceed.sql @@ -1,11 +1,11 @@ UPDATE rag_search.llm_connections SET connection_status = 'inactive' -WHERE id = :connection_id +WHERE vault_uuid = :vault_uuid::uuid RETURNING - id, + vault_uuid, connection_name, connection_status, used_budget, stop_budget_threshold, - disconnect_on_budget_exceed; + disconnect_on_budget_exceed; \ No newline at end of file diff --git a/DSL/Resql/rag-search/POST/get-all-llm-connections-paginated.sql b/DSL/Resql/rag-search/POST/get-all-llm-connections-paginated.sql new file mode 100644 index 00000000..cb1e2394 --- /dev/null +++ b/DSL/Resql/rag-search/POST/get-all-llm-connections-paginated.sql @@ -0,0 +1,45 @@ +SELECT + id, + vault_uuid, + connection_name, + llm_platform, + llm_model, + embedding_platform, + embedding_model, + monthly_budget, + warn_budget_threshold, + stop_budget_threshold, + disconnect_on_budget_exceed, + used_budget, + environment, + connection_status, + created_at, + CEIL(COUNT(*) OVER() / :page_size::DECIMAL) AS totalPages, + CASE + WHEN used_budget IS NULL OR used_budget = 0 OR (used_budget::DECIMAL / monthly_budget::DECIMAL) < (warn_budget_threshold::DECIMAL / 100.0) THEN 'within_budget' + WHEN stop_budget_threshold != 0 AND (used_budget::DECIMAL / monthly_budget::DECIMAL) >= (stop_budget_threshold::DECIMAL / 100.0) THEN 'over_budget' + WHEN stop_budget_threshold = 0 AND (used_budget::DECIMAL / monthly_budget::DECIMAL) >= 1 THEN 'over_budget' + WHEN (used_budget::DECIMAL / monthly_budget::DECIMAL) >= (warn_budget_threshold::DECIMAL / 100.0) THEN 'close_to_exceed' + ELSE 'within_budget' + END AS budget_status +FROM rag_search.llm_connections +WHERE connection_status <> 'deleted' + -- AND environment = 'testing' + AND (:llm_platform IS NULL OR :llm_platform = '' OR llm_platform = :llm_platform) + AND (:llm_model IS NULL OR :llm_model = '' OR llm_model = :llm_model) + AND (:environment IS NULL OR :environment = '' OR environment = :environment) +ORDER BY + CASE WHEN :sorting = 'connection_name asc' THEN connection_name END ASC, + CASE WHEN :sorting = 'connection_name desc' THEN connection_name END DESC, + CASE WHEN :sorting = 'llm_platform asc' THEN llm_platform END ASC, + CASE WHEN :sorting = 'llm_platform desc' THEN llm_platform END DESC, + CASE WHEN :sorting = 'llm_model asc' THEN llm_model END ASC, + CASE WHEN :sorting = 'llm_model desc' THEN llm_model END DESC, + CASE WHEN :sorting = 'monthly_budget asc' THEN monthly_budget END ASC, + CASE WHEN :sorting = 'monthly_budget desc' THEN monthly_budget END DESC, + CASE WHEN :sorting = 'environment asc' THEN environment END ASC, + CASE WHEN :sorting = 'environment desc' THEN environment END DESC, + CASE WHEN :sorting = 'created_at asc' THEN created_at END ASC, + CASE WHEN :sorting = 'created_at desc' THEN created_at END DESC, + created_at DESC -- Default fallback sorting +OFFSET ((GREATEST(:page, 1) - 1) * :page_size) LIMIT :page_size; \ No newline at end of file diff --git a/DSL/Resql/rag-search/POST/get-llm-connection-by-vault-uuid.sql b/DSL/Resql/rag-search/POST/get-llm-connection-by-vault-uuid.sql new file mode 100644 index 00000000..f2d0d7bd --- /dev/null +++ b/DSL/Resql/rag-search/POST/get-llm-connection-by-vault-uuid.sql @@ -0,0 +1,13 @@ +SELECT + id, + vault_uuid, + connection_name, + llm_platform, + llm_model, + embedding_platform, + embedding_model, + environment, + connection_status +FROM rag_search.llm_connections +WHERE vault_uuid = :vault_uuid::uuid + AND connection_status <> 'deleted'; diff --git a/DSL/Resql/rag-search/POST/get-llm-connection.sql b/DSL/Resql/rag-search/POST/get-llm-connection.sql index 025ab507..ba520d16 100644 --- a/DSL/Resql/rag-search/POST/get-llm-connection.sql +++ b/DSL/Resql/rag-search/POST/get-llm-connection.sql @@ -1,5 +1,6 @@ SELECT id, + vault_uuid, connection_name, llm_platform, llm_model, diff --git a/DSL/Resql/rag-search/POST/get-llm-connections-paginated.sql b/DSL/Resql/rag-search/POST/get-llm-connections-paginated.sql index 239e7866..922c16ec 100644 --- a/DSL/Resql/rag-search/POST/get-llm-connections-paginated.sql +++ b/DSL/Resql/rag-search/POST/get-llm-connections-paginated.sql @@ -1,5 +1,6 @@ SELECT id, + vault_uuid, connection_name, llm_platform, llm_model, diff --git a/DSL/Resql/rag-search/POST/get-llm-models.sql b/DSL/Resql/rag-search/POST/get-llm-models.sql new file mode 100644 index 00000000..18df0970 --- /dev/null +++ b/DSL/Resql/rag-search/POST/get-llm-models.sql @@ -0,0 +1,7 @@ +SELECT + id, + platform_id, + model_key as value, + model_name as label +FROM rag_search.llm_models +ORDER BY model_name; \ No newline at end of file diff --git a/DSL/Resql/rag-search/POST/get-production-connection-filtered.sql b/DSL/Resql/rag-search/POST/get-production-connection-filtered.sql index 02ec251d..483fefe3 100644 --- a/DSL/Resql/rag-search/POST/get-production-connection-filtered.sql +++ b/DSL/Resql/rag-search/POST/get-production-connection-filtered.sql @@ -1,5 +1,6 @@ SELECT id, + vault_uuid, connection_name, llm_platform, llm_model, diff --git a/DSL/Resql/rag-search/POST/get-production-connection.sql b/DSL/Resql/rag-search/POST/get-production-connection.sql index f5853c4a..8564a38c 100644 --- a/DSL/Resql/rag-search/POST/get-production-connection.sql +++ b/DSL/Resql/rag-search/POST/get-production-connection.sql @@ -1,5 +1,6 @@ SELECT id, + vault_uuid, connection_name, used_budget, monthly_budget, diff --git a/DSL/Resql/rag-search/POST/get-testing-connection.sql b/DSL/Resql/rag-search/POST/get-testing-connection.sql index 3c1c6ae1..e88ab430 100644 --- a/DSL/Resql/rag-search/POST/get-testing-connection.sql +++ b/DSL/Resql/rag-search/POST/get-testing-connection.sql @@ -1,5 +1,6 @@ SELECT id, + vault_uuid, connection_name, used_budget, monthly_budget, diff --git a/DSL/Resql/rag-search/POST/insert-llm-connection.sql b/DSL/Resql/rag-search/POST/insert-llm-connection.sql index 5b606b1b..26c77acf 100644 --- a/DSL/Resql/rag-search/POST/insert-llm-connection.sql +++ b/DSL/Resql/rag-search/POST/insert-llm-connection.sql @@ -46,6 +46,7 @@ INSERT INTO rag_search.llm_connections ( :embedding_azure_api_key ) RETURNING id, + vault_uuid, connection_name, llm_platform, llm_model, diff --git a/DSL/Resql/rag-search/POST/store-inference-result.sql b/DSL/Resql/rag-search/POST/store-inference-result.sql index 110a937e..25d7e91f 100644 --- a/DSL/Resql/rag-search/POST/store-inference-result.sql +++ b/DSL/Resql/rag-search/POST/store-inference-result.sql @@ -18,7 +18,7 @@ INSERT INTO rag_search.inference_results ( :embedding_scores::JSONB, :final_answer, :environment, - :llm_connection_id, + (SELECT id FROM rag_search.llm_connections WHERE vault_uuid = :vault_uuid::uuid), :created_at::timestamp with time zone ) RETURNING id, @@ -31,4 +31,4 @@ INSERT INTO rag_search.inference_results ( final_answer, environment, llm_connection_id, - created_at; + created_at; \ No newline at end of file diff --git a/DSL/Resql/rag-search/POST/store-testing-inference-result.sql b/DSL/Resql/rag-search/POST/store-testing-inference-result.sql index bc08acba..4775b98f 100644 --- a/DSL/Resql/rag-search/POST/store-testing-inference-result.sql +++ b/DSL/Resql/rag-search/POST/store-testing-inference-result.sql @@ -5,9 +5,9 @@ INSERT INTO rag_search.inference_results ( environment, created_at ) VALUES ( - :llm_connection_id, + (SELECT id FROM rag_search.llm_connections WHERE vault_uuid = :vault_uuid::uuid), :user_question, :final_answer, :environment, :created_at::timestamp with time zone -) RETURNING id, llm_connection_id, user_question, final_answer, environment, created_at; +) RETURNING id, llm_connection_id, user_question, final_answer, environment, created_at; \ No newline at end of file diff --git a/DSL/Resql/rag-search/POST/update-llm-connection-environment.sql b/DSL/Resql/rag-search/POST/update-llm-connection-environment.sql index 1e20a1ce..dbd244a7 100644 --- a/DSL/Resql/rag-search/POST/update-llm-connection-environment.sql +++ b/DSL/Resql/rag-search/POST/update-llm-connection-environment.sql @@ -4,6 +4,7 @@ SET WHERE id = :connection_id RETURNING id, + vault_uuid, connection_name, llm_platform, llm_model, diff --git a/DSL/Resql/rag-search/POST/update-llm-connection-used-budget.sql b/DSL/Resql/rag-search/POST/update-llm-connection-used-budget.sql index 2f4bd4ec..16105f72 100644 --- a/DSL/Resql/rag-search/POST/update-llm-connection-used-budget.sql +++ b/DSL/Resql/rag-search/POST/update-llm-connection-used-budget.sql @@ -1,9 +1,9 @@ UPDATE rag_search.llm_connections SET used_budget = used_budget + :usage -WHERE id = :connection_id +WHERE vault_uuid = :vault_uuid::uuid RETURNING - id, + vault_uuid, connection_name, monthly_budget, used_budget, @@ -11,4 +11,5 @@ RETURNING warn_budget_threshold, stop_budget_threshold, disconnect_on_budget_exceed, - connection_status; \ No newline at end of file + connection_status, + (used_budget >= stop_budget_threshold) AS budget_exceeded; \ No newline at end of file diff --git a/DSL/Resql/rag-search/POST/update-llm-connection.sql b/DSL/Resql/rag-search/POST/update-llm-connection.sql index 91f0bacb..feac78ee 100644 --- a/DSL/Resql/rag-search/POST/update-llm-connection.sql +++ b/DSL/Resql/rag-search/POST/update-llm-connection.sql @@ -27,6 +27,7 @@ SET WHERE id = :connection_id RETURNING id, + vault_uuid, connection_name, llm_platform, llm_model, diff --git a/DSL/Ruuter.private/accounts/GET/user-role.yml b/DSL/Ruuter.private/accounts/GET/user-role.yml new file mode 100644 index 00000000..cd8f3830 --- /dev/null +++ b/DSL/Ruuter.private/accounts/GET/user-role.yml @@ -0,0 +1,38 @@ +declaration: + call: declare + version: 0.1 + description: "Get user roles dynamically from TIM" + method: get + accepts: json + returns: json + namespace: rag-search + allowlist: + headers: + - field: cookie + type: string + description: "Cookie field" + +get_user_info: + call: http.post + args: + url: "[#RAG_SEARCH_TIM]/jwt/custom-jwt-userinfo" + contentType: plaintext + headers: + cookie: ${incoming.headers.cookie} + plaintext: "customJwtCookie" + result: res + next: check_user_info_response + +check_user_info_response: + switch: + - condition: ${200 <= res.response.statusCodeValue && res.response.statusCodeValue < 300} + next: return_result + next: return_empty_array + +return_result: + return: ${res.response.body.authorities} + next: end + +return_empty_array: + return: success + next: end diff --git a/DSL/Ruuter.private/rag-search/GET/llm-connections/all.yml b/DSL/Ruuter.private/rag-search/GET/llm-connections/all.yml new file mode 100644 index 00000000..3b69aebd --- /dev/null +++ b/DSL/Ruuter.private/rag-search/GET/llm-connections/all.yml @@ -0,0 +1,84 @@ +declaration: + call: declare + version: 0.1 + description: "Get paginated list of LLM connections" + method: get + accepts: json + returns: json + namespace: rag-search + allowlist: + params: + - field: pageNumber + type: number + description: "Page number (1-based)" + - field: pageSize + type: number + description: "Number of items per page" + - field: sortBy + type: string + description: "Field to sort by (e.g. 'llm_platform', 'created_at')" + - field: sortOrder + type: string + description: "Sort order: 'asc' or 'desc'" + - field: llmPlatform + type: string + description: "Filter by LLM platform" + - field: llmModel + type: string + description: "Filter by LLM model" + - field: environment + type: string + description: "Filter by deployment environment" + +extract_request_data: + assign: + pageNumber: ${Number(incoming.params.pageNumber) ?? 1} + pageSize: ${Number(incoming.params.pageSize) ?? 10} + sortBy: ${incoming.params.sortBy ?? "created_at"} + sortOrder: ${incoming.params.sortOrder ?? "desc"} + sorting: ${sortBy + " " + sortOrder} + llmPlatform: ${incoming.params.llmPlatform ?? ""} + llmModel: ${incoming.params.llmModel ?? ""} + environment: ${incoming.params.environment ?? ""} + next: validate_page_params + +validate_page_params: + switch: + - condition: ${pageNumber < 1} + next: return_invalid_page + - condition: ${pageSize < 1 || pageSize > 100} + next: return_invalid_page_size + next: get_llm_connections + +get_llm_connections: + call: http.post + args: + url: "[#RAG_SEARCH_RESQL]/get-all-llm-connections-paginated" + body: + page: ${pageNumber} + page_size: ${pageSize} + sorting: ${sorting} + llm_platform: ${llmPlatform} + llm_model: ${llmModel} + environment: ${environment} + result: connections_result + next: transform_response + +transform_response: + assign: + response_data: ${connections_result.response.body} + next: return_success + +return_success: + return: ${response_data} + next: end + +return_invalid_page: + status: 400 + return: "Page number must be greater than 0" + next: end + +return_invalid_page_size: + status: 400 + return: "Page size must be between 1 and 100" + next: end \ No newline at end of file diff --git a/DSL/Ruuter.private/rag-search/GET/llm/models-list.yml b/DSL/Ruuter.private/rag-search/GET/llm/models-list.yml new file mode 100644 index 00000000..594520bb --- /dev/null +++ b/DSL/Ruuter.private/rag-search/GET/llm/models-list.yml @@ -0,0 +1,20 @@ +declaration: + call: declare + version: 0.1 + description: "Get LLM models by platform" + method: get + accepts: query + returns: json + namespace: rag-search + +get_llm_models: + call: http.post + args: + url: "[#RAG_SEARCH_RESQL]/get-llm-models" + result: models_result + next: return_success + +return_success: + return: ${models_result.response.body} + status: 200 + next: end \ No newline at end of file diff --git a/DSL/Ruuter.private/rag-search/POST/inference/results/production/store.yml b/DSL/Ruuter.private/rag-search/POST/inference/results/production/store.yml new file mode 100644 index 00000000..32c5093d --- /dev/null +++ b/DSL/Ruuter.private/rag-search/POST/inference/results/production/store.yml @@ -0,0 +1,100 @@ +declaration: + call: declare + version: 0.1 + description: "Store production inference result with comprehensive data" + method: post + accepts: json + returns: json + namespace: rag-search + allowlist: + body: + - field: chat_id + type: string + description: "Chat ID" + - field: user_question + type: string + description: "User's raw question/input" + - field: refined_questions + type: object + description: "List of refined questions (LLM-generated)" + - field: conversation_history + type: object + description: "Prior messages array of {role, content}" + - field: ranked_chunks + type: object + description: "Retrieved chunks ranked with metadata" + - field: embedding_scores + type: object + description: "Distance scores for each chunk" + - field: final_answer + type: string + description: "LLM's final generated answer" + +extract_request_data: + assign: + chat_id: ${incoming.body.chat_id} + user_question: ${incoming.body.user_question} + refined_questions: ${JSON.stringify(incoming.body.refined_questions) || null} + conversation_history: ${JSON.stringify(incoming.body.conversation_history) || null} + ranked_chunks: ${JSON.stringify(incoming.body.ranked_chunks) || null} + embedding_scores: ${JSON.stringify(incoming.body.embedding_scores) || null} + final_answer: ${incoming.body.final_answer} + created_at: ${new Date().toISOString()} + next: validate_required_fields + +validate_required_fields: + switch: + - condition: "${!user_question || !final_answer}" + next: return_bad_request + next: store_production_inference_result + +store_production_inference_result: + call: http.post + args: + url: "[#RAG_SEARCH_RESQL]/store-production-inference-result" + body: + chat_id: ${chat_id} + user_question: ${user_question} + refined_questions: ${refined_questions} + conversation_history: ${conversation_history} + ranked_chunks: ${ranked_chunks} + embedding_scores: ${embedding_scores} + final_answer: ${final_answer} + environment: "production" + created_at: ${created_at} + result: store_result + next: check_status + +check_status: + switch: + - condition: ${200 <= store_result.response.statusCodeValue && store_result.response.statusCodeValue < 300} + next: format_success_response + next: format_failed_response + +format_success_response: + assign: + data_success: { + data: '${store_result.response.body[0]}', + operationSuccess: true, + statusCode: 200 + } + next: return_success + +format_failed_response: + assign: + data_failed: { + data: '[]', + operationSuccess: false, + statusCode: 400 + } + next: return_bad_request + +return_success: + return: ${data_success} + status: 200 + next: end + +return_bad_request: + return: ${data_failed} + status: 400 + next: end \ No newline at end of file diff --git a/DSL/Ruuter.private/rag-search/POST/inference/results/test/store.yml b/DSL/Ruuter.private/rag-search/POST/inference/results/test/store.yml new file mode 100644 index 00000000..f73496d5 --- /dev/null +++ b/DSL/Ruuter.private/rag-search/POST/inference/results/test/store.yml @@ -0,0 +1,74 @@ +declaration: + call: declare + version: 0.1 + description: "Store inference result" + method: post + accepts: json + returns: json + namespace: rag-search + allowlist: + body: + - field: vault_uuid + type: string + description: "Vault UUID for the LLM connection" + - field: user_question + type: string + description: "User's question/input" + - field: final_answer + type: string + description: "LLM's final generated answer" + +extract_request_data: + assign: + vault_uuid: ${incoming.body.vault_uuid} + user_question: ${incoming.body.user_question} + final_answer: ${incoming.body.final_answer} + created_at: ${new Date().toISOString()} + next: store_inference_result + +store_inference_result: + call: http.post + args: + url: "[#RAG_SEARCH_RESQL]/store-testing-inference-result" + body: + vault_uuid: ${vault_uuid} + user_question: ${user_question} + final_answer: ${final_answer} + environment: "testing" + created_at: ${created_at} + result: store_result + next: check_status + +check_status: + switch: + - condition: ${200 <= store_result.response.statusCodeValue && store_result.response.statusCodeValue < 300} + next: format_success_response + next: format_failed_response + +format_success_response: + assign: + data_success: { + data: '${store_result.response.body[0]}', + operationSuccess: true, + statusCode: 200 + } + next: return_success + +format_failed_response: + assign: + data_failed: { + data: '[]', + operationSuccess: false, + statusCode: 400 + } + next: return_bad_request + +return_success: + return: ${data_success} + status: 200 + next: end + +return_bad_request: + return: ${data_failed} + status: 400 + next: end \ No newline at end of file diff --git a/DSL/Ruuter.private/rag-search/POST/inference/test.yml b/DSL/Ruuter.private/rag-search/POST/inference/test.yml index 4acd4632..5395847b 100644 --- a/DSL/Ruuter.private/rag-search/POST/inference/test.yml +++ b/DSL/Ruuter.private/rag-search/POST/inference/test.yml @@ -60,7 +60,7 @@ call_orchestrate_endpoint: args: url: "[#RAG_SEARCH_LLM_ORCHESTRATOR]/test" body: - connectionId: ${connectionId} + connectionId: ${connection_result.response.body[0].vaultUuid} message: ${message} environment: "testing" headers: diff --git a/DSL/Ruuter.private/rag-search/POST/llm-connections/add.yml b/DSL/Ruuter.private/rag-search/POST/llm-connections/add.yml index 5e7326aa..a43a9e10 100644 --- a/DSL/Ruuter.private/rag-search/POST/llm-connections/add.yml +++ b/DSL/Ruuter.private/rag-search/POST/llm-connections/add.yml @@ -133,6 +133,13 @@ update_production_connection: connection_id: ${existing_production_result.response.body[0].id} environment: "testing" result: update_result + next: clear_llm_cache + +clear_llm_cache: + call: http.post + args: + url: "[#RAG_SEARCH_LLM_CACHE_CLEAR]" + result: cache_clear_result next: add_llm_connection add_llm_connection: @@ -170,6 +177,7 @@ assign_connection_response: assign: response: { id: "${connection_result.response.body[0].id}", + vaultUuid: "${connection_result.response.body[0].vaultUuid}", status: 201, operationSuccess: true } diff --git a/DSL/Ruuter.private/rag-search/POST/llm-connections/edit.yml b/DSL/Ruuter.private/rag-search/POST/llm-connections/edit.yml index 6ae6f10e..09c73d69 100644 --- a/DSL/Ruuter.private/rag-search/POST/llm-connections/edit.yml +++ b/DSL/Ruuter.private/rag-search/POST/llm-connections/edit.yml @@ -149,6 +149,13 @@ update_production_connection: connection_id: ${existing_production_result.response.body[0].id} environment: "testing" result: update_result + next: clear_llm_cache + +clear_llm_cache: + call: http.post + args: + url: "[#RAG_SEARCH_LLM_CACHE_CLEAR]" + result: cache_clear_result next: update_llm_connection update_llm_connection: diff --git a/DSL/Ruuter.private/rag-search/POST/vault/secret/create.yml b/DSL/Ruuter.private/rag-search/POST/vault/secret/create.yml index b6a533d9..e26ce661 100644 --- a/DSL/Ruuter.private/rag-search/POST/vault/secret/create.yml +++ b/DSL/Ruuter.private/rag-search/POST/vault/secret/create.yml @@ -8,9 +8,9 @@ declaration: namespace: rag-search allowlist: body: - - field: connectionId + - field: vaultUuid type: string - description: "Body field 'connectionId'" + description: "Body field 'vaultUuid' - stable UUID for vault path" - field: llmPlatform type: string description: "Body field 'llmPlatform'" @@ -65,7 +65,7 @@ declaration: extract_request_data: assign: - connectionId: ${incoming.body.connectionId} + vaultUuid: ${incoming.body.vaultUuid} llmPlatform: ${incoming.body.llmPlatform} llmModel: ${incoming.body.llmModel} secretKey: ${incoming.body.secretKey} @@ -100,7 +100,7 @@ execute_aws_request: url: "[#RAG_SEARCH_CRON_MANAGER]/execute/store_in_vault/store_secrets" query: cookie: ${incoming.headers.cookie.replace('customJwtCookie=','')} #Removing the customJwtCookie phrase from payload to to send cookie token only - connectionId: ${connectionId} + vaultUuid: ${vaultUuid} llmPlatform: ${llmPlatform} llmModel: ${llmModel} secretKey: ${secretKey} @@ -120,7 +120,7 @@ execute_azure_request: url: "[#RAG_SEARCH_CRON_MANAGER]/execute/store_in_vault/store_secrets" query: cookie: ${incoming.headers.cookie.replace('customJwtCookie=','')} #Removing the customJwtCookie phrase from payload to to send cookie token only - connectionId: ${connectionId} + vaultUuid: ${vaultUuid} llmPlatform: ${llmPlatform} llmModel: ${llmModel} deploymentName: ${deploymentName} diff --git a/DSL/Ruuter.private/rag-search/POST/vault/secret/delete.yml b/DSL/Ruuter.private/rag-search/POST/vault/secret/delete.yml index f0a72200..67248c4d 100644 --- a/DSL/Ruuter.private/rag-search/POST/vault/secret/delete.yml +++ b/DSL/Ruuter.private/rag-search/POST/vault/secret/delete.yml @@ -8,9 +8,9 @@ declaration: namespace: rag-search allowlist: body: - - field: connectionId + - field: vaultUuid type: string - description: "Body field 'connectionId'" + description: "Body field 'vaultUuid' - stable UUID for vault path" - field: llmPlatform type: string description: "Body field 'llmPlatform'" @@ -33,7 +33,7 @@ declaration: extract_request_data: assign: - connectionId: ${incoming.body.connectionId} + vaultUuid: ${incoming.body.vaultUuid} llmPlatform: ${incoming.body.llmPlatform} llmModel: ${incoming.body.llmModel} embeddingModel: ${incoming.body.embeddingModel} @@ -45,9 +45,9 @@ extract_request_data: check_connection_exists: call: http.post args: - url: "[#RAG_SEARCH_RESQL]/get-llm-connection" + url: "[#RAG_SEARCH_RESQL]/get-llm-connection-by-vault-uuid" body: - connection_id: ${connectionId} + vault_uuid: ${vaultUuid} result: connection_result next: validate_connection_response @@ -63,19 +63,18 @@ execute_delete_request: url: "[#RAG_SEARCH_CRON_MANAGER]/execute/delete_from_vault/delete_secrets" query: cookie: ${incoming.headers.cookie.replace('customJwtCookie=','')} #Removing the customJwtCookie phrase from payload to to send cookie token only - connectionId: ${connectionId} + vaultUuid: ${vaultUuid} llmPlatform: ${llmPlatform} llmModel: ${llmModel} embeddingModel: ${embeddingModel} embeddingPlatform: ${embeddingPlatform} - deploymentEnvironment: ${deploymentEnvironment} result: cron_delete_res next: return_delete_ok assign_validation_error: assign: validation_error_res: { - message: 'Required fields missing: connectionId, llmPlatform, llmModel, and deploymentEnvironment are required', + message: 'Required fields missing: vaultUuid, llmPlatform, llmModel, and deploymentEnvironment are required', operationSuccessful: false, statusCode: 400 } @@ -84,7 +83,7 @@ assign_validation_error: assign_connection_not_found_error: assign: connection_not_found_res: { - message: 'Connection not found with the provided connectionId', + message: 'Connection not found with the provided vaultUuid', operationSuccessful: false, statusCode: 404 } diff --git a/DSL/Ruuter.public/rag-search/GET/accounts/user-role.yml b/DSL/Ruuter.public/rag-search/GET/accounts/user-role.yml new file mode 100644 index 00000000..d37b8911 --- /dev/null +++ b/DSL/Ruuter.public/rag-search/GET/accounts/user-role.yml @@ -0,0 +1,38 @@ +declaration: + call: declare + version: 0.1 + description: "Get user roles dynamically from TIM" + method: get + accepts: json + returns: json + namespace: backoffice + allowlist: + headers: + - field: cookie + type: string + description: "Cookie field" + +get_user_info: + call: http.post + args: + url: "[#CKB_TIM]/jwt/custom-jwt-userinfo" + contentType: plaintext + headers: + cookie: ${incoming.headers.cookie} + plaintext: "customJwtCookie" + result: res + next: check_user_info_response + +check_user_info_response: + switch: + - condition: ${200 <= res.response.statusCodeValue && res.response.statusCodeValue < 300} + next: return_result + next: return_empty_array + +return_result: + return: ${res.response.body.authorities} + next: end + +return_empty_array: + return: success + next: end diff --git a/DSL/Ruuter.public/rag-search/POST/api-tools/index.yml b/DSL/Ruuter.public/rag-search/POST/api-tools/index.yml index 98de449f..7d476d46 100644 --- a/DSL/Ruuter.public/rag-search/POST/api-tools/index.yml +++ b/DSL/Ruuter.public/rag-search/POST/api-tools/index.yml @@ -43,7 +43,7 @@ extract_request_data: name: ${incoming.body.name} description: ${incoming.body.description} method: ${incoming.body.method || 'GET'} - url: ${incoming.body.url} + url: ${encodeURIComponent(incoming.body.url)} visibility: ${incoming.body.visibility || 'public'} type: ${incoming.body.type || 'custom_endpoint'} params: ${encodeURIComponent(JSON.stringify(incoming.body.params || []))} @@ -80,13 +80,7 @@ execute_indexing: params: ${params} result: indexing_result on_error: handle_cron_error - next: check_indexing_status - -check_indexing_status: - switch: - - condition: ${200 <= indexing_result.response.statusCodeValue && indexing_result.response.statusCodeValue < 300} - next: assign_success - next: assign_cron_failure + next: assign_success handle_cron_error: log: "ERROR: Failed to queue api_tool indexing job - ${indexing_result.error || 'CronManager unreachable'}" @@ -98,7 +92,7 @@ assign_cron_failure: success: false error: "INDEXING_QUEUE_FAILED" message: "Failed to queue indexing job. CronManager may be unavailable." - details: ${indexing_result.error} + details: ${indexing_result.error || 'CronManager unreachable'} next: return_server_error assign_success: diff --git a/DSL/Ruuter.public/rag-search/POST/inference/results/store.yml b/DSL/Ruuter.public/rag-search/POST/inference/results/store.yml index 19d8adf2..34a9be4f 100644 --- a/DSL/Ruuter.public/rag-search/POST/inference/results/store.yml +++ b/DSL/Ruuter.public/rag-search/POST/inference/results/store.yml @@ -32,9 +32,9 @@ declaration: - field: environment type: string description: "Environment identifier (e.g., production, testing)" - - field: llm_connection_id + - field: vault_uuid type: string - description: "Connection identifier" + description: "Vault UUID for the LLM connection" extract_request_data: assign: @@ -46,7 +46,7 @@ extract_request_data: embedding_scores: ${JSON.stringify(incoming.body.embedding_scores) || null} final_answer: ${incoming.body.final_answer} environment: ${incoming.body.environment} - llm_connection_id: ${incoming.body.llm_connection_id} + vault_uuid: ${incoming.body.vault_uuid} created_at: ${new Date().toISOString()} next: validate_required_fields @@ -69,7 +69,7 @@ store_production_inference_result: embedding_scores: ${embedding_scores} final_answer: ${final_answer} environment: ${environment} - llm_connection_id: ${llm_connection_id} + vault_uuid: ${vault_uuid} created_at: ${created_at} result: store_result next: check_status diff --git a/DSL/Ruuter.public/rag-search/POST/services/enrich.yml b/DSL/Ruuter.public/rag-search/POST/services/enrich.yml index 5748ad59..b177f3fe 100644 --- a/DSL/Ruuter.public/rag-search/POST/services/enrich.yml +++ b/DSL/Ruuter.public/rag-search/POST/services/enrich.yml @@ -75,13 +75,7 @@ execute_enrichment: is_common: ${service_is_common} result: enrichment_result on_error: handle_cron_error - next: check_enrichment_status - -check_enrichment_status: - switch: - - condition: ${200 <= enrichment_result.response.statusCodeValue && enrichment_result.response.statusCodeValue < 300} - next: assign_success - next: assign_cron_failure + next: assign_success handle_cron_error: log: "ERROR: Failed to queue enrichment job - ${enrichment_result.error || 'CronManager unreachable'}" @@ -93,7 +87,7 @@ assign_cron_failure: success: false error: "ENRICHMENT_QUEUE_FAILED" message: "Failed to queue enrichment job. CronManager may be unavailable." - details: ${enrichment_result.error} + details: ${enrichment_result.error || 'CronManager unreachable'} next: return_server_error assign_success: diff --git a/Dockerfile.vault-init b/Dockerfile.vault-init new file mode 100644 index 00000000..7743fa6e --- /dev/null +++ b/Dockerfile.vault-init @@ -0,0 +1,10 @@ +FROM hashicorp/vault:1.20.3 + +# Bake the only CLI tools vault-init.sh actually needs (jq + openssl) so container +# startup never depends on the Alpine CDN. Previously these were installed via +# `apk add` on every boot, which failed intermittently on EC2. Retry guards the +# one-time build against a transient mirror hiccup. +RUN for i in 1 2 3; do \ + apk add --no-cache jq openssl && break; \ + echo "apk add failed (attempt $i), retrying..."; sleep 3; \ + done diff --git a/GUI/.env.development b/GUI/.env.development index ae5b1356..c15f05c6 100644 --- a/GUI/.env.development +++ b/GUI/.env.development @@ -4,4 +4,6 @@ REACT_APP_CUSTOMER_SERVICE_LOGIN=http://localhost:3004/et/dev-auth REACT_APP_SERVICE_ID=conversations,settings,monitoring REACT_APP_NOTIFICATION_NODE_URL=http://localhost:4040 REACT_APP_CSP=upgrade-insecure-requests; default-src 'self'; font-src 'self' data:; img-src 'self' data:; script-src 'self' 'unsafe-eval' 'unsafe-inline'; style-src 'self' 'unsafe-inline'; object-src 'none'; connect-src 'self' http://localhost:8086 http://localhost:8088 http://localhost:3004 http://localhost:4040 ws://localhost; -REACT_APP_ENABLE_HIDDEN_FEATURES=TRUE \ No newline at end of file +REACT_APP_ENABLE_HIDDEN_FEATURES=TRUE +REACT_APP_ENABLE_MULTI_DOMAIN=FALSE +REACT_APP_MENU_JSON= '[{"id":"conversations","label":{"et":"Vestlused","en":"Conversations"},"path":"/chat","children":[{"label":{"et":"Vastamata","en":"Unanswered"},"path":"/unanswered"},{"label":{"et":"Aktiivsed","en":"Active"},"path":"/active"},{"label":{"et":"Ootel","en":"Pending"},"path":"/pending"},{"label":{"et":"Ajalugu","en":"History"},"path":"/history"},{"label":{"et":"Valideerimised","en":"Validations"},"path":"/validations"}]},{"id":"training","label":{"et":"Treening","en":"Training"},"path":"/training","children":[{"label":{"et":"Treening","en":"Training"},"path":"/training","children":[{"label":{"et":"Teemad","en":"Themes"},"path":"/training/intents"},{"hidden":true,"label":{"et":"Avalikud teemad","en":"Public themes"},"path":"/training/common-intents"},{"label":{"et":"Teemade järeltreenimine","en":"Post training themes"},"path":"/training/intents-followup-training"},{"label":{"et":"Vastused","en":"Answers"},"path":"/training/responses"},{"label":{"et":"Reeglid","en":"Rules"},"path":"/training/rules"},{"hidden":true,"label":{"et":"Konfiguratsioon","en":"Configuration"},"path":"/training/configuration"},{"label":{"et":"Vormid","en":"Forms"},"path":"/training/forms"},{"label":{"et":"Mälukohad","en":"Slots"},"path":"/training/slots"}]},{"label":{"et":"Ajaloolised vestlused","en":"Historical conversations"},"path":"/history","children":[{"label":{"et":"Ajalugu","en":"History"},"path":"/history/history"},{"hidden":true,"label":{"et":"Pöördumised","en":"Appeals"},"path":"/history/appeal"}]},{"label":{"et":"Mudelipank ja analüütika","en":"Modelbank and analytics"},"path":"/analytics","children":[{"label":{"et":"Teemade ülevaade","en":"Overview of topics"},"path":"/analytics/overview"},{"label":{"et":"Mudelite võrdlus","en":"Comparison of models"},"path":"/analytics/models"},{"hidden":true,"label":{"et":"Testlood","en":"testTracks"},"path":"/analytics/testcases"}]},{"label":{"et":"Treeni uus mudel","en":"Train new model"},"path":"/train-new-model"}]},{"id":"analytics","label":{"et":"Analüütika","en":"Analytics"},"path":"/analytics","children":[{"label":{"et":"Ülevaade","en":"Overview"},"path":"/overview"},{"label":{"et":"Vestlused","en":"Chats"},"path":"/chats"},{"label":{"et":"Tagasiside","en":"Feedback"},"path":"/feedback"},{"label":{"et":"Avaandmed","en":"Reports"},"path":"/reports"}]},{"id":"services","hidden":true,"label":{"et":"Teenused","en":"Services"},"path":"/services","children":[{"label":{"et":"Ülevaade","en":"Overview"},"path":"/overview"},{"label":{"et":"Uus teenus","en":"New Service"},"path":"/newService"},{"label":{"et":"Probleemsed teenused","en":"Faulty Services"},"path":"/faultyServices"}]},{"id":"knowledge-center","label":{"et":"Teadmuskeskus","en":"Knowledge Center"},"path":"/knowledge-center","children":[{"id":"rag-search","label":{"et":"Mudelid ja seadistused","en":"Models management"},"path":"/rag-search","children":[{"label":{"et":"Mudelite ühendused","en":"LLM connections"},"path":"/llm-connections"},{"label":{"et":"Viiba Seaded","en":"Prompt Configurations"},"path":"/prompt-configurations"},{"label":{"et":"Testi mudelit","en":"Test LLM"},"path":"/test-llm"}]},{"id":"ckb","label":{"et":"Teadmusbaas","en":"Knowledge Base"},"path":"/ckb","children":[{"label":{"et":"Agentuur","en":"Agency"},"path":"/agency"},{"label":{"et":"Aruanded","en":"Reports"},"path":"/reports"},{"label":{"et":"API Integratsioonid","en":"API Integrations"},"path":"/api"}]}]},{"id":"settings","label":{"et":"Haldus","en":"Administration"},"path":"/settings","children":[{"label":{"et":"Kasutajad","en":"Users"},"path":"/users"},{"label":{"et":"Vestlusbot","en":"Chatbot"},"path":"/chatbot","children":[{"label":{"et":"Seaded","en":"Settings"},"path":"/chatbot/settings"},{"label":{"et":"Tervitussõnum","en":"Welcome message"},"path":"/chatbot/welcome-message"},{"label":{"et":"Välimus ja käitumine","en":"Appearance and behavior"},"path":"/chatbot/appearance"},{"label":{"et":"Erakorralised teated","en":"Emergency notices"},"path":"/chatbot/emergency-notices"},{"label":{"et":"Tagasiside","en":"Feedback"},"path":"/chatbot/feedback"}]},{"label":{"et":"Vestluste analüüs","en":"Chat analysis"},"path":"/chat-analysis"},{"label":{"et":"Asutuse tööaeg","en":"Office opening hours"},"path":"/working-time"},{"label":{"et":"Vestluse Kustutamine","en":"Delete Conversations"},"path":"/delete-conversations"},{"label":{"et":"Sessiooni pikkus","en":"Session length"},"path":"/session-length"},{"label":{"et":"SKMi konfiguratsioon","en":"SKM Configuration"},"path":"/skm-configuration"},{"label":{"et":"Anonümiseerija","en":"Anonymizer"},"path":"/anonymizer"}]},{"id":"monitoring","hidden":true,"label":{"et":"Seire","en":"Monitoring"},"path":"/monitoring","children":[{"label":{"et":"Aktiivaeg","en":"Working hours"},"path":"/uptime"}]}]' diff --git a/GUI/package-lock.json b/GUI/package-lock.json index c0f45b17..5f1f779d 100644 --- a/GUI/package-lock.json +++ b/GUI/package-lock.json @@ -8,6 +8,8 @@ "name": "byk-training-module-gui", "version": "0.0.0", "dependencies": { + "@buerokratt-ria/header": "^0.1.52", + "@buerokratt-ria/menu": "^0.2.15", "@buerokratt-ria/styles": "^0.0.1", "@fontsource/roboto": "^4.5.8", "@formkit/auto-animate": "^1.0.0-beta.5", @@ -89,7 +91,7 @@ "eslint-plugin-typescript": "^0.14.0", "mocksse": "^1.0.4", "msw": "^0.49.2", - "prettier": "^2.8.1", + "prettier": "^2.8.8", "sass": "^1.57.0", "typescript": "^4.9.3", "vite": "^4.0.0", @@ -2154,6 +2156,81 @@ "node": ">=6.9.0" } }, + "node_modules/@buerokratt-ria/header": { + "version": "0.1.52", + "resolved": "https://registry.npmjs.org/@buerokratt-ria/header/-/header-0.1.52.tgz", + "integrity": "sha512-4k9Ekym6UqQpyF2nEIIT0n3wdp6Ax9sTvWK4HeZ7xuwf9YdgYnYRARRnZxcfsJCMm54JUQ8IKV9Y5i5ULVk1Bg==", + "license": "ISC", + "dependencies": { + "@buerokratt-ria/styles": "^0.0.1", + "@types/react": "^18.2.21", + "react": "^18.2.0" + }, + "peerDependencies": { + "@fontsource/roboto": "^4.5.8", + "@formkit/auto-animate": "^0.7.0", + "@radix-ui/react-accessible-icon": "^1.0.3", + "@radix-ui/react-dialog": "^1.0.4", + "@radix-ui/react-switch": "^1.0.3", + "@radix-ui/react-toast": "^1.1.4", + "@tanstack/react-query": "^4.32.1", + "axios": "^1.4.0", + "clsx": "^1.2.1", + "howler": "^2.2.4", + "i18next": "^23.2.3", + "i18next-browser-languagedetector": "^7.1.0", + "path": "^0.12.7", + "react": "^18.2.0", + "react-cookie": "^4.1.1", + "react-dom": "^18.2.0", + "react-hook-form": "^7.45.4", + "react-i18next": "^12.1.1", + "react-icons": "^4.10.1", + "react-idle-timer": "^5.7.2", + "react-router-dom": "^6.14.2", + "react-select": "^5.7.4", + "rxjs": "^7.8.1", + "tslib": "^2.3.0", + "vite-plugin-dts": "^3.5.2", + "vite-plugin-svgr": "^3.2.0", + "zustand": "^4.4.0" + } + }, + "node_modules/@buerokratt-ria/menu": { + "version": "0.2.15", + "resolved": "https://registry.npmjs.org/@buerokratt-ria/menu/-/menu-0.2.15.tgz", + "integrity": "sha512-xtjIyD/3Mdg2UZ7A6Yv/3ZsZX8LGUFGpYtr3LaYeyfhv5FsLWfDwXNPL2ViooOdpPJaD4hmKURrBF2VgpmBfUw==", + "license": "ISC", + "dependencies": { + "@buerokratt-ria/styles": "^0.0.1", + "@types/react": "^18.2.21", + "react": "^18.2.0" + }, + "peerDependencies": { + "@radix-ui/react-accessible-icon": "^1.0.3", + "@radix-ui/react-dialog": "^1.0.4", + "@radix-ui/react-switch": "^1.0.3", + "@radix-ui/react-toast": "^1.1.4", + "@tanstack/react-query": "^4.32.1", + "clsx": "^1.2.1", + "i18next": "^23.2.3", + "i18next-browser-languagedetector": "^7.1.0", + "path": "^0.12.7", + "react": "^18.2.0", + "react-cookie": "^4.1.1", + "react-dom": "^18.2.0", + "react-hook-form": "^7.45.4", + "react-i18next": "^12.1.1", + "react-icons": "^4.10.1", + "react-idle-timer": "^5.7.2", + "react-router-dom": "^6.14.2", + "rxjs": "^7.8.1", + "tslib": "^2.3.0", + "vite-plugin-dts": "^3.5.2", + "vite-plugin-svgr": "^3.2.0", + "zustand": "^4.4.0" + } + }, "node_modules/@buerokratt-ria/styles": { "version": "0.0.1", "resolved": "https://registry.npmjs.org/@buerokratt-ria/styles/-/styles-0.0.1.tgz", @@ -14263,6 +14340,7 @@ "resolved": "https://registry.npmjs.org/prettier/-/prettier-2.8.8.tgz", "integrity": "sha512-tdN8qQGvNjw4CHbY+XXk0JgCXn9QiF21a55rBe5LJAU+kDyC4WQn4+awm2Xfk2lQMk5fKup9XgzTZtGkjBdP9Q==", "dev": true, + "license": "MIT", "bin": { "prettier": "bin-prettier.js" }, diff --git a/GUI/package.json b/GUI/package.json index ec9c4e78..e1dc0692 100644 --- a/GUI/package.json +++ b/GUI/package.json @@ -11,6 +11,8 @@ "prettier": "prettier --write \"{,!(node_modules)/**/}*.{ts,tsx,js,json,css,less,scss}\"" }, "dependencies": { + "@buerokratt-ria/header": "^0.1.52", + "@buerokratt-ria/menu": "^0.2.15", "@buerokratt-ria/styles": "^0.0.1", "@fontsource/roboto": "^4.5.8", "@formkit/auto-animate": "^1.0.0-beta.5", @@ -92,7 +94,7 @@ "eslint-plugin-typescript": "^0.14.0", "mocksse": "^1.0.4", "msw": "^0.49.2", - "prettier": "^2.8.1", + "prettier": "^2.8.8", "sass": "^1.57.0", "typescript": "^4.9.3", "vite": "^4.0.0", diff --git a/GUI/src/App.tsx b/GUI/src/App.tsx index 5839b180..43c50753 100644 --- a/GUI/src/App.tsx +++ b/GUI/src/App.tsx @@ -64,8 +64,7 @@ const App: FC = () => { } /> } /> } /> - } /> - } /> + } /> diff --git a/GUI/src/components/Layout/index.tsx b/GUI/src/components/Layout/index.tsx index c26eca42..ef2864ad 100644 --- a/GUI/src/components/Layout/index.tsx +++ b/GUI/src/components/Layout/index.tsx @@ -3,15 +3,23 @@ import { Outlet } from 'react-router-dom'; import useStore from 'store'; import './Layout.scss'; import { useToast } from '../../hooks/useToast'; -import Header from 'components/Header'; -import MainNavigation from 'components/MainNavigation'; +import { MainNavigation } from '@buerokratt-ria/menu'; +import { Header, useMenuCountConf } from '@buerokratt-ria/header'; const Layout: FC = () => { + const domainBarShowing = import.meta.env.REACT_APP_ENABLE_MULTI_DOMAIN?.toLowerCase() === 'true'; + const menuCountConf = useMenuCountConf(); + return (
- +
-
+
diff --git a/GUI/src/components/MainNavigation/index.tsx b/GUI/src/components/MainNavigation/index.tsx index 070c4b9a..8ae278ca 100644 --- a/GUI/src/components/MainNavigation/index.tsx +++ b/GUI/src/components/MainNavigation/index.tsx @@ -44,12 +44,6 @@ const MainNavigation: FC = () => { label: t('menu.testLLM'), path: '/test-llm', icon: - }, - { - id: 'testProductionLLM', - label: t('menu.testProductionLLM'), - path: '/test-production-llm', - icon: } ]; diff --git a/GUI/src/hooks/useStreamingResponse.tsx b/GUI/src/hooks/useStreamingResponse.tsx index fc7204cd..6d53a5ff 100644 --- a/GUI/src/hooks/useStreamingResponse.tsx +++ b/GUI/src/hooks/useStreamingResponse.tsx @@ -1,147 +1,336 @@ -import { useState, useRef, useCallback, useEffect } from 'react'; -import axios from 'axios'; -import { ChoiceButton } from 'services/inference'; - -const getNotificationNodeUrl = (): string => { - const value = import.meta.env.REACT_APP_NOTIFICATION_NODE_URL; - if (!value) { - throw new Error( - 'Environment variable REACT_APP_NOTIFICATION_NODE_URL is not defined. ' + - 'Please set it to the base URL of the notification service to enable streaming responses.' - ); - } - return value; -}; -const notificationNodeUrl = getNotificationNodeUrl(); -console.log(notificationNodeUrl); - -interface StreamingOptions { - authorId: string; - conversationHistory: Array<{ authorRole: string; message: string; timestamp: string }>; - url: string; -} - -interface UseStreamingResponseReturn { - startStreaming: (message: string, options: StreamingOptions, onToken: (token: string) => void, onComplete: () => void, onError: (error: string) => void, onButtons?: (buttons: ChoiceButton[]) => void) => Promise; - stopStreaming: () => void; - isStreaming: boolean; -} - -export const useStreamingResponse = (channelId: string): UseStreamingResponseReturn => { - const [isStreaming, setIsStreaming] = useState(false); - const eventSourceRef = useRef(null); - - const stopStreaming = useCallback(() => { - if (eventSourceRef.current) { - console.log('[SSE] Closing connection'); - eventSourceRef.current.close(); - eventSourceRef.current = null; - } - setIsStreaming(false); - }, []); - - // Cleanup on unmount - useEffect(() => { - return () => { - if (eventSourceRef.current) { - eventSourceRef.current.close(); - } - }; - }, []); - - const startStreaming = useCallback( - async ( - message: string, - options: StreamingOptions, - onToken: (token: string) => void, - onComplete: () => void, - onError: (error: string) => void, - onButtons?: (buttons: ChoiceButton[]) => void - ) => { - console.log('[SSE] Starting streaming for channel:', channelId); - - // Close any existing connection - stopStreaming(); - - try { - // Step 1: Open SSE connection FIRST - const sseUrl = `${notificationNodeUrl}/sse/stream/${channelId}`; - console.log('[SSE] Connecting to:', sseUrl); - - const eventSource = new EventSource(sseUrl); - eventSourceRef.current = eventSource; - - eventSource.onopen = () => { - console.log('[SSE] Connection opened'); - }; - - eventSource.onmessage = (event) => { - console.log('[SSE] Message received:', event.data); - - try { - const data = JSON.parse(event.data); - - if (data.type === 'stream_start') { - console.log('[SSE] Stream started'); - setIsStreaming(true); - } else if (data.type === 'stream_chunk' && data.content) { - console.log('[SSE] Token:', data.content); - onToken(data.content); - if (data.buttons && data.buttons.length > 0 && onButtons) { - onButtons(data.buttons); - } - } else if (data.type === 'stream_end') { - console.log('[SSE] Stream ended'); - setIsStreaming(false); - eventSource.close(); - eventSourceRef.current = null; - onComplete(); - } else if (data.type === 'stream_error') { - console.error('[SSE] Stream error:', data.error); - setIsStreaming(false); - eventSource.close(); - eventSourceRef.current = null; - onError(data.error || 'Stream error occurred'); - } - } catch (e) { - console.error('[SSE] Failed to parse message:', e); - } - }; - - eventSource.onerror = (err) => { - console.error('[SSE] Connection error:', err); - setIsStreaming(false); - eventSource.close(); - eventSourceRef.current = null; - onError('Connection error'); - }; - - // Step 2: Wait a moment for SSE connection to establish, then trigger the stream - await new Promise(resolve => setTimeout(resolve, 500)); - - // Step 3: POST to trigger streaming - const postUrl = `${notificationNodeUrl}/channels/${channelId}/orchestrate/stream`; - console.log('[API] Triggering stream:', postUrl); - - await axios.post(postUrl, { - message, - options, - }); - - console.log('[API] Stream triggered successfully'); - - } catch (err) { - console.error('[SSE] Error starting stream:', err); - stopStreaming(); - onError(err instanceof Error ? err.message : 'Failed to start streaming'); - } - }, - [channelId, stopStreaming] - ); - - return { - startStreaming, - stopStreaming, - isStreaming, - }; -}; \ No newline at end of file +import { useState, useRef, useCallback, useEffect } from 'react'; +import axios from 'axios'; +import { ChoiceButton } from 'services/inference'; + +const getNotificationNodeUrl = (): string => { + const value = import.meta.env.REACT_APP_NOTIFICATION_NODE_URL; + if (!value) { + throw new Error( + 'Environment variable REACT_APP_NOTIFICATION_NODE_URL is not defined. ' + + 'Please set it to the base URL of the notification service to enable streaming responses.' + ); + } + return value; +}; +const notificationNodeUrl = getNotificationNodeUrl(); +console.log(notificationNodeUrl); + +// The trigger POST is held open by the notification server for the full +// generation, so it needs a ceiling well above any realistic answer time. +const TRIGGER_POST_TIMEOUT_MS = 600_000; + +// How long to keep waiting on SSE after the trigger POST has failed. A failed +// POST is not proof that generation failed - it may already be streaming - but +// if nothing has arrived by now, nothing is coming: the POST failed before the +// server ever dispatched to the relay, so no stream_end or stream_error will +// ever be sent and without this the UI would wait forever. +const TRIGGER_FAILURE_GRACE_MS = 15_000; + +// Messages that prove the backend actually engaged this stream. Heartbeats are +// SSE comment frames, which EventSource never surfaces, so they cannot count. +const STREAM_EVENT_TYPES = new Set([ + 'stream_start', + 'stream_chunk', + 'stream_end', + 'stream_error', +]); + +// Typewriter pacing. +// +// Output guardrails validate the answer in blocks before releasing it, so tokens +// reach the browser in bursts (roughly 200, then 150, then the tail) rather than +// one at a time. Rendering each burst the instant it lands makes the answer snap +// onto the screen. Instead we queue arriving tokens and drain them on a timer, +// which gives a steady word-by-word effect without weakening the guardrails or +// paying for the extra validation calls that smaller server-side chunks cost. +// +// Typing speed. This is the knob to turn if the effect feels too fast or slow - +// higher is faster. Around 25/s reads like brisk typing; 40+/s starts to look +// like the text is simply appearing. +const TYPING_TOKENS_PER_SECOND = 25; +// Ceiling on how far rendering may fall behind the stream. If a burst arrives +// faster than the typing speed, the drain rate rises so the backlog still clears +// within this budget rather than typing on long after the answer is complete. +// Approximate: setInterval fires late under load, so expect ~15-20% over. +const MAX_CATCH_UP_SECONDS = 8; + +const DRAIN_INTERVAL_MS = Math.round(1000 / TYPING_TOKENS_PER_SECOND); +const DRAIN_TARGET_TICKS = Math.round( + (MAX_CATCH_UP_SECONDS * 1000) / DRAIN_INTERVAL_MS +); + +interface StreamingOptions { + authorId: string; + conversationHistory: Array<{ authorRole: string; message: string; timestamp: string }>; + url: string; +} + +interface UseStreamingResponseReturn { + startStreaming: (message: string, options: StreamingOptions, onToken: (token: string) => void, onComplete: () => void, onError: (error: string) => void, onButtons?: (buttons: ChoiceButton[]) => void) => Promise; + stopStreaming: () => void; + isStreaming: boolean; +} + +export const useStreamingResponse = (channelId: string): UseStreamingResponseReturn => { + const [isStreaming, setIsStreaming] = useState(false); + const eventSourceRef = useRef(null); + + // Typewriter state + const queueRef = useRef([]); + const drainTimerRef = useRef | null>(null); + // Tokens emitted per tick. Only ever raised during a run: recomputing it from + // the shrinking queue each tick would decay the rate geometrically and stretch + // a large backlog out well beyond MAX_CATCH_UP_SECONDS. + const drainRateRef = useRef(1); + const streamEndedRef = useRef(false); + const pendingButtonsRef = useRef(null); + // Set by any stream event; read by the trigger-POST fallback below. + const sawStreamEventRef = useRef(false); + const triggerFallbackTimerRef = useRef | null>(null); + const onTokenRef = useRef<(token: string) => void>(() => {}); + const onCompleteRef = useRef<() => void>(() => {}); + const onButtonsRef = useRef<((buttons: ChoiceButton[]) => void) | undefined>(undefined); + + const stopDrain = useCallback(() => { + if (drainTimerRef.current) { + clearInterval(drainTimerRef.current); + drainTimerRef.current = null; + } + }, []); + + const clearTriggerFallback = useCallback(() => { + if (triggerFallbackTimerRef.current) { + clearTimeout(triggerFallbackTimerRef.current); + triggerFallbackTimerRef.current = null; + } + }, []); + + // Drop anything not yet rendered. Used when a guardrail blocks the answer or + // the user cancels: text the rail rejected must never reach the screen. + const discardQueue = useCallback(() => { + queueRef.current = []; + pendingButtonsRef.current = null; + drainRateRef.current = 1; + stopDrain(); + }, [stopDrain]); + + const startDrain = useCallback(() => { + if (drainTimerRef.current) return; + + drainTimerRef.current = setInterval(() => { + const queue = queueRef.current; + + if (queue.length === 0) { + if (streamEndedRef.current) { + stopDrain(); + if (pendingButtonsRef.current?.length && onButtonsRef.current) { + onButtonsRef.current(pendingButtonsRef.current); + pendingButtonsRef.current = null; + } + setIsStreaming(false); + onCompleteRef.current(); + } + return; + } + + drainRateRef.current = Math.max( + drainRateRef.current, + Math.ceil(queue.length / DRAIN_TARGET_TICKS) + ); + onTokenRef.current(queue.splice(0, drainRateRef.current).join('')); + }, DRAIN_INTERVAL_MS); + }, [stopDrain]); + + const stopStreaming = useCallback(() => { + if (eventSourceRef.current) { + console.log('[SSE] Closing connection'); + eventSourceRef.current.close(); + eventSourceRef.current = null; + } + clearTriggerFallback(); + discardQueue(); + setIsStreaming(false); + }, [discardQueue, clearTriggerFallback]); + + // Cleanup on unmount + useEffect(() => { + return () => { + if (eventSourceRef.current) { + eventSourceRef.current.close(); + } + if (drainTimerRef.current) { + clearInterval(drainTimerRef.current); + } + if (triggerFallbackTimerRef.current) { + clearTimeout(triggerFallbackTimerRef.current); + } + }; + }, []); + + const startStreaming = useCallback( + async ( + message: string, + options: StreamingOptions, + onToken: (token: string) => void, + onComplete: () => void, + onError: (error: string) => void, + onButtons?: (buttons: ChoiceButton[]) => void + ) => { + console.log('[SSE] Starting streaming for channel:', channelId); + + // Close any existing connection + stopStreaming(); + + // Reset typewriter state for this run + queueRef.current = []; + drainRateRef.current = 1; + streamEndedRef.current = false; + pendingButtonsRef.current = null; + sawStreamEventRef.current = false; + onTokenRef.current = onToken; + onCompleteRef.current = onComplete; + onButtonsRef.current = onButtons; + + try { + // Step 1: Open SSE connection FIRST + const sseUrl = `${notificationNodeUrl}/sse/stream/${channelId}`; + console.log('[SSE] Connecting to:', sseUrl); + + const eventSource = new EventSource(sseUrl); + eventSourceRef.current = eventSource; + + eventSource.onopen = () => { + console.log('[SSE] Connection opened'); + }; + + eventSource.onmessage = (event) => { + console.log('[SSE] Message received:', event.data); + + try { + const data = JSON.parse(event.data); + + if (STREAM_EVENT_TYPES.has(data.type)) { + // The backend is talking to us, so the trigger-POST fallback must + // not fire - whatever happens next arrives over this connection. + sawStreamEventRef.current = true; + clearTriggerFallback(); + } + + if (data.type === 'stream_start') { + console.log('[SSE] Stream started'); + setIsStreaming(true); + startDrain(); + } else if (data.type === 'stream_chunk' && data.content) { + // Queue rather than render, so bursts play out word by word. + queueRef.current.push(data.content); + if (data.buttons && data.buttons.length > 0) { + // Held back until the text finishes typing, so choices do not + // appear above a half-rendered answer. + pendingButtonsRef.current = data.buttons; + } + startDrain(); + } else if (data.type === 'stream_end') { + console.log('[SSE] Stream ended'); + eventSource.close(); + eventSourceRef.current = null; + // Do not fire onComplete yet - let the queue finish rendering. + streamEndedRef.current = true; + startDrain(); + } else if (data.type === 'stream_error') { + console.error('[SSE] Stream error:', data.error); + // Discard unrendered text: a guardrail may have just rejected it. + discardQueue(); + setIsStreaming(false); + eventSource.close(); + eventSourceRef.current = null; + onError(data.error || 'Stream error occurred'); + } + } catch (e) { + console.error('[SSE] Failed to parse message:', e); + } + }; + + eventSource.onerror = (err) => { + console.error('[SSE] Connection error:', err); + // Otherwise a pending fallback would later report a second error for + // the same failed run. + clearTriggerFallback(); + discardQueue(); + setIsStreaming(false); + eventSource.close(); + eventSourceRef.current = null; + onError('Connection error'); + }; + + // Step 2: Wait a moment for SSE connection to establish, then trigger the stream + await new Promise(resolve => setTimeout(resolve, 500)); + + // Step 3: POST to trigger streaming. + // Note: this request stays open for the whole generation, so it can fail + // (gateway 504, network blip) while the SSE stream is perfectly healthy. + const postUrl = `${notificationNodeUrl}/channels/${channelId}/orchestrate/stream`; + console.log('[API] Triggering stream:', postUrl); + + try { + await axios.post( + postUrl, + { message, options }, + { timeout: TRIGGER_POST_TIMEOUT_MS } + ); + console.log('[API] Stream triggered successfully'); + } catch (postErr) { + // Do NOT tear down the EventSource here. The answer arrives over SSE, + // not in this response body; killing the stream on a POST failure is + // what turned a slow answer into a truncated one. Let stream_end / + // stream_error decide, and only surface an error if neither arrives. + console.warn( + '[API] Trigger POST failed; keeping SSE open and waiting for stream events:', + postErr + ); + + const postErrMessage = + postErr instanceof Error ? postErr.message : 'Failed to start streaming'; + + if (!eventSourceRef.current) { + // SSE already gone - nothing left to wait for. + onError(postErrMessage); + return; + } + + // Bound the wait. If the POST failed before the server dispatched to + // the relay (a 4xx, or no active connection so the request was only + // queued), no stream event will ever arrive and there is nothing to + // end the stream - so give up once the grace period expires. + clearTriggerFallback(); + triggerFallbackTimerRef.current = setTimeout(() => { + triggerFallbackTimerRef.current = null; + if (sawStreamEventRef.current) return; + + console.error( + `[SSE] No stream events ${TRIGGER_FAILURE_GRACE_MS}ms after trigger POST failed; giving up` + ); + discardQueue(); + setIsStreaming(false); + if (eventSourceRef.current) { + eventSourceRef.current.close(); + eventSourceRef.current = null; + } + onError(postErrMessage); + }, TRIGGER_FAILURE_GRACE_MS); + } + + } catch (err) { + console.error('[SSE] Error starting stream:', err); + stopStreaming(); + onError(err instanceof Error ? err.message : 'Failed to start streaming'); + } + }, + [channelId, stopStreaming, startDrain, discardQueue, clearTriggerFallback] + ); + + return { + startStreaming, + stopStreaming, + isStreaming, + }; +}; diff --git a/GUI/src/pages/LLMConnections/index.tsx b/GUI/src/pages/LLMConnections/index.tsx index 0af35601..83617ad7 100644 --- a/GUI/src/pages/LLMConnections/index.tsx +++ b/GUI/src/pages/LLMConnections/index.tsx @@ -18,6 +18,7 @@ import { llmConnectionsQueryKeys } from 'utils/queryKeys'; import { useToast } from 'hooks/useToast'; import { ToastTypes } from 'enums/commonEnums'; import useStore from 'store'; +import { getAllLLMModels, getLLMPlatforms } from 'services/llmConfigs'; const LLMConnections: FC = () => { const { t } = useTranslation(); @@ -49,6 +50,17 @@ const LLMConnections: FC = () => { }); + // Fetch platform and model options from API + const { data: llmPlatformsData = [], isLoading: llmPlatformsLoading, error: llmPlatformsError } = useQuery({ + queryKey: ['llm-platforms'], + queryFn: getLLMPlatforms + }); + + const { data: llmModels= [], isLoading: llmModelsLoading, error: llmModelsError } = useQuery({ + queryKey: ['llm-models'], + queryFn: getAllLLMModels, + }); + const llmConnections = connectionsResponse; const totalPages = connectionsResponse?.[0]?.totalPages || 1; @@ -107,21 +119,21 @@ const LLMConnections: FC = () => { setPageIndex(1); } }; - - // Platform filter options + const platformOptions = [ { label: t('dataModels.filters.allPlatforms'), value: 'all' }, - { label: t('dataModels.platforms.azure'), value: 'azure' }, - { label: t('dataModels.platforms.aws'), value: 'aws' }, + ...llmPlatformsData.map((platform) => ({ + label: platform.label, + value: platform.value, + })), ]; - // LLM Model filter options - these would ideally come from an API const llmModelOptions = [ { label: t('dataModels.filters.allModels'), value: 'all' }, - { label: t('dataModels.models.gpt4Mini'), value: 'gpt-4o-mini' }, - { label: t('dataModels.models.gpt4o'), value: 'gpt-4o' }, - { label: t('dataModels.models.claude35Sonnet'), value: 'anthropic-claude-3.5-sonnet' }, - { label: t('dataModels.models.claude37Sonnet'), value: 'anthropic-claude-3.7-sonnet' }, + ...llmModels.map((model) => ({ + label: model.label, + value: model.value, + })), ]; // Environment filter options diff --git a/GUI/src/pages/TestModel/TestLLM.scss b/GUI/src/pages/TestModel/TestLLM.scss deleted file mode 100644 index 3d0c2156..00000000 --- a/GUI/src/pages/TestModel/TestLLM.scss +++ /dev/null @@ -1,217 +0,0 @@ -.testModalFormTextArea { - margin-top: 30px; -} - -.mcq-buttons { - display: flex; - flex-wrap: wrap; - gap: 0.75rem; - margin-top: 1rem; -} - -.testModalClassifyButton { - text-align: right; - margin-top: 20px; -} - -.llm-connection-section { - width: 50%; -} - -.llm-connection-controls { - display: flex; - gap: 1rem; - align-items: center; -} - -.inference-results-container { - max-width: 100%; - background-color: #d7efff; - padding: 20px; - border-radius: 8px; - margin-top: 20px; - - .result-item { - margin-bottom: 15px; - - strong { - color: #333; - } - } - - .response-content { - margin-top: 8px; - padding: 12px; - background-color: #f5f5f5; - border-radius: 4px; - white-space: pre-wrap; - line-height: 1.5; - color: #555; - } - - .context-section { - margin-top: 20px; - - .context-list { - display: flex; - flex-direction: column; - gap: 12px; - margin-top: 8px; - } - - .context-item { - padding: 12px; - background-color: #ffffff; - border: 1px solid #e0e0e0; - border-radius: 6px; - box-shadow: 0 1px 3px rgba(0, 0, 0, 0.1); - - .context-rank { - margin-bottom: 8px; - padding-bottom: 4px; - border-bottom: 1px solid #f0f0f0; - - strong { - color: #2563eb; - font-size: 0.875rem; - font-weight: 600; - } - } - - .context-content { - color: #374151; - line-height: 1.5; - font-size: 0.9rem; - white-space: pre-wrap; - } - } - } -} - -.testModalList { - list-style: disc; - margin-left: 30px; -} - -.mt-20 { - margin-top: 20px; -} - -.classification-results { - margin-top: 1rem; - padding: 1rem; - border: 1px solid #e0e0e0; - border-radius: 8px; - background-color: #f9f9f9; - - h3 { - margin: 0 0 1rem 0; - color: #333; - } - - h4 { - margin: 0 0 0.75rem 0; - color: #555; - font-size: 1rem; - } - - .results-container { - display: flex; - flex-direction: column; - gap: 1.5rem; - } - - .top-prediction { - .prediction-card { - display: flex; - justify-content: space-between; - align-items: center; - padding: 1rem; - border-radius: 8px; - background-color: #e8f5e8; - border: 2px solid #4caf50; - - .agency-name { - font-weight: 600; - color: #2e7d32; - font-size: 1.1rem; - } - - .confidence-score { - font-weight: 700; - color: #2e7d32; - font-size: 1.2rem; - } - } - } - - .predictions-list { - display: flex; - flex-direction: column; - gap: 0.75rem; - - .prediction-item { - display: flex; - align-items: center; - gap: 1rem; - padding: 0.75rem; - background-color: white; - border-radius: 6px; - border: 1px solid #ddd; - - &.highest { - border-color: #4caf50; - background-color: #f8fff8; - } - - .rank { - font-weight: 600; - color: #666; - min-width: 2rem; - } - - .agency-info { - flex: 1; - display: flex; - flex-direction: column; - gap: 0.25rem; - - .agency-name { - font-weight: 500; - color: #333; - } - - .confidence-bar-container { - width: 100%; - height: 4px; - background-color: #e0e0e0; - border-radius: 2px; - overflow: hidden; - - .confidence-bar { - height: 100%; - background-color: #4caf50; - transition: width 0.3s ease; - } - } - } - - .confidence-percentage { - font-weight: 600; - color: #555; - min-width: 4rem; - text-align: right; - } - } - } -} - -.classification-error { - margin-top: 1rem; - padding: 1rem; - background-color: #ffebee; - border: 1px solid #f44336; - border-radius: 6px; - color: #c62828; - text-align: center; -} \ No newline at end of file diff --git a/GUI/src/pages/TestModel/index.tsx b/GUI/src/pages/TestModel/index.tsx deleted file mode 100644 index 2fc116bd..00000000 --- a/GUI/src/pages/TestModel/index.tsx +++ /dev/null @@ -1,228 +0,0 @@ -import { useMutation, useQuery } from '@tanstack/react-query'; -import { Button, FormSelect, FormTextarea, Collapsible } from 'components'; -import CircularSpinner from 'components/molecules/CircularSpinner/CircularSpinner'; -import { ComponentPropsWithoutRef, FC, useState } from 'react'; -import { useTranslation } from 'react-i18next'; -import ReactMarkdown from 'react-markdown'; -import remarkGfm from 'remark-gfm'; -import './TestLLM.scss'; -import { useDialog } from 'hooks/useDialog'; -import { fetchLLMConnectionsPaginated, LegacyLLMConnectionFilters } from 'services/llmConnections'; -import { viewInferenceResult, InferenceRequest, InferenceResponse, ChoiceButton } from 'services/inference'; -import { llmConnectionsQueryKeys } from 'utils/queryKeys'; -import { ButtonAppearanceTypes } from 'enums/commonEnums'; - -const TestLLM: FC = () => { - const { t } = useTranslation(); - const { open: openDialog, close: closeDialog } = useDialog(); - const [inferenceResult, setInferenceResult] = useState(null); - const [pendingButtons, setPendingButtons] = useState([]); - const [testLLM, setTestLLM] = useState({ - connectionId: null, - text: '', - }); - - // Sort context by rank - const sortedContext = inferenceResult?.chunks?.toSorted((a, b) => a.rank - b.rank) ?? []; - - // Fetch LLM connections for dropdown - using the working legacy endpoint for now - const { data: connections, isLoading: isLoadingConnections } = useQuery({ - queryKey: llmConnectionsQueryKeys.list({ - page: 1, - pageSize: 100, // Get all connections for dropdown - sorting: 'created_at desc', - }), - queryFn: () => fetchLLMConnectionsPaginated({ - pageNumber: 1, - pageSize: 100, - sortBy: 'created_at desc', - }), - }); - - // Transform connections data for dropdown - const connectionOptions = connections?.map((connection: any) => ({ - label: `${connection.llmPlatform} - ${connection.llmModel} (${connection.environment})`, - value: connection.id, - })) || []; - - // Inference mutation - const inferenceMutation = useMutation({ - mutationFn: (request: InferenceRequest) => viewInferenceResult(request), - onSuccess: (data: InferenceResponse) => { - setInferenceResult(data?.response); - setPendingButtons(data?.response?.buttons ?? []); - }, - onError: (error: any) => { - console.error('Error getting inference result:', error); - openDialog({ - title: t('testModels.inferenceErrorTitle') || 'Inference Error', - content:

{t('testModels.inferenceErrorMessage') || 'Failed to get inference result. Please try again.'}

, - footer: ( - - ), - }); - }, - }); - - const handleSend = () => { - if (testLLM.connectionId && testLLM.text) { - inferenceMutation.mutate({ - llmConnectionId: Number(testLLM.connectionId), - message: testLLM.text, - }); - } - }; - - const handleButtonClick = (payload: string) => { - if (!testLLM.connectionId) return; - setPendingButtons([]); - inferenceMutation.mutate({ - llmConnectionId: Number(testLLM.connectionId), - message: payload, - }); - }; - - const handleChange = (key: string, value: string | number) => { - // Prevent changes while inference is loading - if (inferenceMutation.isLoading) { - return; - } - setTestLLM((prev) => ({ - ...prev, - [key]: value, - })); - }; - - const markdownComponents = { - ol: ({children}: any) => ( -
    - {children} -
- ), - a: (props: ComponentPropsWithoutRef<"a">) => ( - - ), - }; - - return ( -
- {isLoadingConnections ? ( - - ) : ( -
-
-
{t('testModels.title') || 'Test LLM'}
-
-
-

{t('testModels.llmConnectionLabel') || 'LLM Connection'}

-
- - { - handleChange('connectionId', selection?.value as string); - }} - value={testLLM?.connectionId === null ? t('testModels.connectionNotExist') || 'Connection does not exist' : undefined} - defaultValue={testLLM?.connectionId ?? undefined} - disabled={inferenceMutation.isLoading} - /> -
-
- -
-

{t('testModels.classifyTextLabel') || 'Enter text to test'}

- handleChange('text', e.target.value)} - showMaxLength={true} - /> -
-
- -
- - {/* Inference Result */} - - {inferenceResult && !inferenceMutation.isLoading && ( -
-
- Response: -
- - {inferenceResult.content} - -
-
- - {/* MCQ Buttons */} - {pendingButtons.length > 0 && ( -
- {pendingButtons.map((btn) => ( - - ))} -
- )} - - {/* Context Section */} - { - sortedContext && sortedContext?.length > 0 && ( -
- -
- {sortedContext?.map((contextItem, index) => ( -
-
- Rank {contextItem.rank} -
-
- - {contextItem.chunkRetrieved} - -
-
- ))} -
-
-
- ) - } - -
- )} - - {/* Error State */} - {inferenceMutation.isError && ( -
-

{t('testModels.classificationFailed') || 'Inference failed. Please try again.'}

-
- )} -
- )} -
- ); -}; - -export default TestLLM; \ No newline at end of file diff --git a/GUI/src/pages/TestProductionLLM/TestProductionLLM.scss b/GUI/src/pages/TestProductionLLM/TestProductionLLM.scss index 9cb5c00c..2fa456fb 100644 --- a/GUI/src/pages/TestProductionLLM/TestProductionLLM.scss +++ b/GUI/src/pages/TestProductionLLM/TestProductionLLM.scss @@ -3,6 +3,21 @@ margin: 0 auto; padding: 2rem; + .llm-connection-section { + width: 50%; + margin-bottom: 1.5rem; + p { + margin-bottom: 0.5rem; + font-weight: 500; + color: #333; + } + } + .llm-connection-controls { + display: flex; + gap: 1rem; + align-items: center; + } + .mcq-buttons { display: flex; flex-wrap: wrap; diff --git a/GUI/src/pages/TestProductionLLM/index.tsx b/GUI/src/pages/TestProductionLLM/index.tsx index f29cfcf9..8ba70ea6 100644 --- a/GUI/src/pages/TestProductionLLM/index.tsx +++ b/GUI/src/pages/TestProductionLLM/index.tsx @@ -1,11 +1,16 @@ import { FC, useState, useRef, useEffect, useMemo } from 'react'; +import { useQuery } from '@tanstack/react-query'; import { useTranslation } from 'react-i18next'; -import { Button, FormTextarea } from 'components'; +import { Button, FormTextarea, FormSelect } from 'components'; import { useToast } from 'hooks/useToast'; import { useStreamingResponse } from 'hooks/useStreamingResponse'; import { ChoiceButton } from 'services/inference'; import './TestProductionLLM.scss'; import MessageContent from 'components/MessageContent'; +import { llmConnectionsQueryKeys } from 'utils/queryKeys'; +import { fetchAllLLMConnectionsPaginated } from 'services/llmConnections'; + + interface Message { id: string; content: string; @@ -23,11 +28,49 @@ const TestProductionLLM: FC = () => { const [messages, setMessages] = useState([]); const [isLoading, setIsLoading] = useState(false); const messagesEndRef = useRef(null); + const [testLLM, setTestLLM] = useState({ + connectionId: null, + text: '', + }); // Generate a unique channel ID for this session const channelId = useMemo(() => `channel-${Math.random().toString(36).substring(2, 15)}`, []); const { startStreaming, stopStreaming, isStreaming } = useStreamingResponse(channelId); + const [selectedConnectionId, setSelectedConnectionId] = useState(null); + + // Fetch LLM connections for dropdown - using the working legacy endpoint for now + const { data: connections, isLoading: isLoadingConnections } = useQuery({ + queryKey: llmConnectionsQueryKeys.list({ + page: 1, + pageSize: 100, // Get all connections for dropdown + sorting: 'created_at desc', + }), + queryFn: () => fetchAllLLMConnectionsPaginated({ + pageNumber: 1, + pageSize: 100, + sortBy: 'created_at desc', + }), + }); + // Transform connections data for dropdown + const connectionOptions = useMemo( + () => + connections?.map((connection: any) => ({ + label: `${connection.llmPlatform} - ${connection.llmModel} (${connection.environment})`, + value: String(connection.id), + })) || [], + [connections] + ); + const selectedConnection = useMemo(() => { + return connections?.find((conn: any) => String(conn.id) === selectedConnectionId) || null; + }, [connections, selectedConnectionId]); + + const handleConnectionChange = (value: string | number) => { + console.log('Selected connection ID:', value); + if (isLoading || isStreaming) return; + setSelectedConnectionId(value ? String(value) : null); + }; + // Auto-scroll to bottom useEffect(() => { messagesEndRef.current?.scrollIntoView({ behavior: 'smooth' }); @@ -82,6 +125,8 @@ const TestProductionLLM: FC = () => { authorId: 'test-user-456', conversationHistory, url: 'opensearch-dashboard-test', + environment: selectedConnection?.environment || 'production', + connection_id: selectedConnection?.vaultUuid || undefined, }; // Callbacks for streaming @@ -257,11 +302,27 @@ const TestProductionLLM: FC = () => {
-

{t('testProductionLLM.title')}

+

{t('testModels.title')}

+
+

{t('testModels.llmConnectionLabel') || 'LLM Connection'}

+
+ { + handleConnectionChange(selection?.value as string); + }} + defaultValue={selectedConnectionId ?? undefined} + disabled={isLoading || isStreaming} + /> +
+
@@ -313,7 +374,7 @@ const TestProductionLLM: FC = () => {
))} - {isLoading && ( + {isLoading && (messages.length === 0 || messages[messages.length - 1].isUser) && (
diff --git a/GUI/src/services/llmConfigs.ts b/GUI/src/services/llmConfigs.ts index fc582a86..3776a112 100644 --- a/GUI/src/services/llmConfigs.ts +++ b/GUI/src/services/llmConfigs.ts @@ -45,4 +45,10 @@ export async function getEmbeddingModels(platformKey?: string): Promise { + const { data } = await apiDev.get('/rag-search/llm/models-list'); + return data?.response; } \ No newline at end of file diff --git a/GUI/src/services/llmConnections.ts b/GUI/src/services/llmConnections.ts index 65417f3c..cd07324e 100644 --- a/GUI/src/services/llmConnections.ts +++ b/GUI/src/services/llmConnections.ts @@ -6,6 +6,7 @@ import { encryptLLMCredentials } from 'utils/encryption'; export interface LLMConnection { id: number; + vaultUuid?: string; connectionName: string; llmPlatform: string; llmModel: string; @@ -126,7 +127,7 @@ export interface LLMConnectionFormData { } // Vault secret service functions -async function createVaultSecret(connectionId: string, connectionData: LLMConnectionFormData): Promise { +async function createVaultSecret(vaultUuid: string, connectionData: LLMConnectionFormData): Promise { // Encrypt sensitive credentials before sending to vault const encryptedCredentials = await encryptLLMCredentials({ @@ -143,7 +144,7 @@ async function createVaultSecret(connectionId: string, connectionData: LLMConnec }); const payload = { - connectionId, + vaultUuid, llmPlatform: connectionData.llmPlatform, llmModel: connectionData.llmModel, embeddingModel: connectionData.embeddingModel, @@ -176,15 +177,14 @@ async function createVaultSecret(connectionId: string, connectionData: LLMConnec await apiDev.post(vaultEndpoints.CREATE_VAULT_SECRET(), payload); } -async function deleteVaultSecret(connectionId: string, connectionData: Partial): Promise { +async function deleteVaultSecret(vaultUuid: string, connectionData: Partial): Promise { const payload = { - connectionId, + vaultUuid, llmPlatform: connectionData.llmPlatform || '', llmModel: connectionData.llmModel || '', embeddingModel: connectionData.embeddingModel || '', embeddingPlatform: connectionData.embeddingModelPlatform || '', - deploymentEnvironment: connectionData.deploymentEnvironment?.toLowerCase() || '', }; await apiDev.post(vaultEndpoints.DELETE_VAULT_SECRET(), payload); @@ -266,8 +266,8 @@ export async function createLLMConnection(connectionData: LLMConnectionFormData) console.log('Created LLM Connection:', connection); // After successful database creation, store secrets in vault - if (connection && connection.id) { - await createVaultSecret(connection.id.toString(), connectionData); + if (connection && connection.id && connection.vaultUuid) { + await createVaultSecret(connection.vaultUuid, connectionData); } return connection; @@ -315,7 +315,10 @@ export async function updateLLMConnection( || connectionData.embeddingSecretKey && !connectionData.embeddingSecretKey?.includes('*') || connectionData.embeddingAzureApiKey && !connectionData.embeddingAzureApiKey?.includes('*'))) { try { - await createVaultSecret(id.toString(), connectionData); + const vaultUuid = connection.vaultUuid || (await getLLMConnection(id)).vaultUuid; + if (vaultUuid) { + await createVaultSecret(vaultUuid, connectionData); + } } catch (vaultError) { console.error('Failed to update secrets in vault:', vaultError); } @@ -333,27 +336,26 @@ export async function deleteLLMConnection(id: string | number): Promise { console.error('Failed to get connection data before deletion:', error); } - // Delete from database - await apiDev.post(llmConnectionsEndpoints.DELETE_LLM_CONNECTION(), { - connection_id: id, - }); - - // After successful database deletion, delete secrets from vault - if (connectionToDelete) { + // Delete secrets from vault BEFORE database deletion + // (delete.yml validates connection exists, so DB must not be soft-deleted yet) + if (connectionToDelete && connectionToDelete.vaultUuid) { try { - await deleteVaultSecret(id.toString(), { + await deleteVaultSecret(connectionToDelete.vaultUuid, { llmPlatform: connectionToDelete.llmPlatform, llmModel: connectionToDelete.llmModel, embeddingModel: connectionToDelete.embeddingModel, embeddingModelPlatform: connectionToDelete.embeddingPlatform, - deploymentEnvironment: connectionToDelete.environment, }); } catch (vaultError) { console.error('Failed to delete secrets from vault:', vaultError); - // Note: We don't throw here as the database deletion has already succeeded - // This is logged for monitoring/debugging purposes + // Continue with database deletion even if vault deletion fails } } + + // Delete from database (soft-delete sets connection_status = 'deleted') + await apiDev.post(llmConnectionsEndpoints.DELETE_LLM_CONNECTION(), { + connection_id: id, + }); } export async function checkBudgetStatus(): Promise { @@ -376,3 +378,19 @@ export async function updateLLMConnectionStatus( }); return data?.response; } + +export async function fetchAllLLMConnectionsPaginated(filters: LLMConnectionFilters): Promise { + const queryParams = new URLSearchParams(); + + if (filters.pageNumber) queryParams.append('pageNumber', filters.pageNumber.toString()); + if (filters.pageSize) queryParams.append('pageSize', filters.pageSize.toString()); + if (filters.sortBy) queryParams.append('sortBy', filters.sortBy); + if (filters.sortOrder) queryParams.append('sortOrder', filters.sortOrder); + if (filters.llmPlatform) queryParams.append('llmPlatform', filters.llmPlatform); + if (filters.llmModel) queryParams.append('llmModel', filters.llmModel); + if (filters.environment) queryParams.append('environment', filters.environment); + + const url = `${llmConnectionsEndpoints.FETCH_ALL_LLM_CONNECTIONS_PAGINATED()}?${queryParams.toString()}`; + const { data } = await apiDev.get(url); + return data?.response; +} \ No newline at end of file diff --git a/GUI/src/store/index.ts b/GUI/src/store/index.ts index c5fe37db..16db7e6e 100644 --- a/GUI/src/store/index.ts +++ b/GUI/src/store/index.ts @@ -5,7 +5,9 @@ import { LLMConnectionFilters, ProductionConnectionFilters } from 'services/llmC interface StoreState { userInfo: UserInfo | null; userId: string; + userDomains: string[]; setUserInfo: (info: UserInfo) => void; + setUserDomains: (domains: string[]) => void; llmConnectionFilters: LLMConnectionFilters; llmConnectionPageIndex: number; productionConnectionFilters: ProductionConnectionFilters; @@ -32,7 +34,9 @@ const defaultProductionConnectionFilters: ProductionConnectionFilters = { const useStore = create((set) => ({ userInfo: null, userId: '', + userDomains: [], setUserInfo: (data) => set({ userInfo: data, userId: data?.userIdCode || '' }), + setUserDomains: (domains: string[]) => set({ userDomains: domains }), llmConnectionFilters: defaultLLMConnectionFilters, llmConnectionPageIndex: 1, productionConnectionFilters: defaultProductionConnectionFilters, diff --git a/GUI/src/utils/endpoints.ts b/GUI/src/utils/endpoints.ts index 386db296..30624914 100644 --- a/GUI/src/utils/endpoints.ts +++ b/GUI/src/utils/endpoints.ts @@ -15,6 +15,7 @@ export const authEndpoints = { export const llmConnectionsEndpoints = { FETCH_LLM_CONNECTIONS_PAGINATED: (): string => `/rag-search/llm-connections/list`, + FETCH_ALL_LLM_CONNECTIONS_PAGINATED: (): string => `/rag-search/llm-connections/all`, GET_LLM_CONNECTION: (): string => `/rag-search/llm-connections/get`, GET_PRODUCTION_CONNECTION: (): string => `/rag-search/llm-connections/production`, CREATE_LLM_CONNECTION: (): string => `/rag-search/llm-connections/add`, diff --git a/GUI/translations/en/common.json b/GUI/translations/en/common.json index f3d932c7..09178417 100644 --- a/GUI/translations/en/common.json +++ b/GUI/translations/en/common.json @@ -36,6 +36,11 @@ "endDate": "End date", "preview": "Preview", "logout": "Logout", + "present": "Present", + "away": "Away", + "csaStatus": "Customer support status", + "statusClarification": "Status clarification", + "notificationErrorMsg": "An error occurred", "change": "Change", "loading": "Loading", "asc": "asc", @@ -59,6 +64,18 @@ "entries": "records", "deleteSelected": "Delete selection" }, + "chat": { + "unanswered": "Unanswered", + "forwarded": "Forwarded", + "pending": "Pending" + }, + "mainMenu": { + "menuLabel": "Main navigation", + "closeMenu": "Collapse menu", + "openMenu": "Open menu", + "openIcon": "Open menu icon", + "closeIcon": "Close menu icon" + }, "menu": { "userManagement": "User management", "testLLM": "Test LLM", @@ -155,6 +172,7 @@ "aws": "AWS Bedrock" }, "models": { + "gpt4.1": "GPT-4.1", "gpt4Mini": "GPT-4 Mini", "gpt4o": "GPT-4o", "claude35Sonnet": "Anthropic Claude 3.5 Sonnet", diff --git a/GUI/translations/et/common.json b/GUI/translations/et/common.json index 52e76621..c4e41f87 100644 --- a/GUI/translations/et/common.json +++ b/GUI/translations/et/common.json @@ -36,6 +36,11 @@ "endDate": "Lõppkuupäev", "preview": "Eelvaade", "logout": "Logi välja", + "present": "Kohal", + "away": "Eemal", + "csaStatus": "Nõustaja", + "statusClarification": "Staatuse täpsustus", + "notificationErrorMsg": "Midagi läks valesti", "change": "Muuda", "loading": "Laadimine", "asc": "kasvav", @@ -59,6 +64,18 @@ "entries": "kirjeid", "deleteSelected": "Kustuta valik" }, + "chat": { + "unanswered": "Vastamata", + "forwarded": "Suunatud", + "pending": "Ootel" + }, + "mainMenu": { + "menuLabel": "Põhinavigatsioon", + "closeMenu": "Kitsenda menüü", + "openMenu": "Ava menüü", + "openIcon": "Ava menüü ikoon", + "closeIcon": "Sulge menüü ikoon" + }, "menu": { "userManagement": "Kasutajate haldus", "testLLM": "Testi mudelit", @@ -155,6 +172,7 @@ "aws": "AWS Bedrock" }, "models": { + "gpt4.1": "GPT-4.1", "gpt4Mini": "GPT-4 Mini", "gpt4o": "GPT-4o", "claude35Sonnet": "Anthropic Claude 3.5 Sonnet", diff --git a/README.md b/README.md index ad5edce9..d50cdcac 100644 --- a/README.md +++ b/README.md @@ -1,47 +1,214 @@ -# BYK-RAG (Retrieval-Augmented Generation Module) +# LLM Module -The **BYK-RAG Module** is part of the Burokratt ecosystem, designed to provide **retrieval-augmented generation (RAG)** capabilities for Estonian government digital services. It ensures reliable, multilingual, and compliant AI-powered responses by integrating with multiple LLM providers syncing with knowledge bases, and exposing flexible configuration and monitoring features for administrators. +The **LLM Module** is the LLM orchestration component of the +[Bürokratt](https://github.com/buerokratt) ecosystem, providing reliable, multilingual, and compliant +AI-powered responses for Estonian government digital services. It is a **multi-workflow orchestrator**: +a tool classifier inspects every user query and routes it to the most appropriate workflow — answering +from the knowledge base, calling backend services and APIs, handling conversation, or declining +gracefully when a request is out of scope. ---- +## Overview -## Features +Rather than treating every request as a single retrieval problem, the LLM Module classifies intent and +dispatches to one of several specialised workflows. Retrieval-Augmented Generation (RAG) is **one** of +these workflows — alongside Service, Context, API-Tool, and Out-of-Domain handling. All workflows run +over configurable, multi-provider LLMs, are protected by safety guardrails, and are fully traced for +cost and quality. -- **Configurable LLM Providers** - - Support for AWS Bedrock, Azure AI, Google Cloud, OpenAI, Anthropic, and self-hosted open-source LLMs. - - Admins can create "connections" and switch providers/models without downtime. - - Models searchable via dropdown with cache-enabled indicators. +### Key Features -- **Enhanced Security with RSA Encryption** - - LLM credentials encrypted with RSA-2048 asymmetric encryption before storage. - - GUI encrypts using public key; CronManager decrypts with private key. - - Additional security layer beyond HashiCorp Vault's encryption. +- **Multi-workflow orchestration** — a tool classifier routes each query to the Service, Context, + **RAG**, API-Tool, or Out-of-Domain (OOD) workflow using hybrid dense + sparse (BM25) search. +- **Configurable LLM providers** — Azure OpenAI, AWS Bedrock, Google Cloud, OpenAI, Anthropic, and + self-hosted models. Admins create "connections" and switch providers/models without downtime. +- **Grounded, cited answers** — RAG responses are restricted to Central Knowledge Base content, with + clear citations and an "I don't know" fallback when confidence is low. +- **Agentic API tool calling** — decomposes multi-intent queries and executes multi-endpoint API + workflows to fulfil actionable requests. +- **Secure credential management** — provider credentials stored in HashiCorp Vault with an additional + RSA-2048 encryption layer. +- **Safety guardrails** — NeMo Guardrails check input and output content with cost tracking. +- **Observability** — Langfuse for traces/cost analytics and Grafana/Loki for logs. -- **Knowledge Base Integration** - - Continuous sync with central knowledge base (CKB). - - Last sync timestamp displayed in UI. - - LLMs restricted to answering only from CKB content. - - "I don't know" payload returned when confidence is low. +## Architecture -- **Citations & Transparency** - - All responses are accompanied with **clear citations**. +The LLM Module sits behind the Ruuter API gateway and orchestrates retrieval, generation, and tool +calling over a set of supporting data stores (Qdrant, Redis, PostgreSQL/ClickHouse, MinIO) and +HashiCorp Vault for secrets. -- **Analytics & Monitoring** - - External **Langfuse dashboard** for API usage, inference trends, cost analysis, and performance logs. - - Agencies can configure cost alerts and view alerts via LLM Alerts UI. - - Logs integrated with **Grafana Loki**. +![LLM Module — System Context](./docs/images/LLM%20Module%20Context%20Diagram%20(Current).png) -### Storing Langfuse Secrets +For the full picture — container and component (C4) diagrams plus the request lifecycle — see +**[docs/ARCHITECTURE.md](./docs/ARCHITECTURE.md)**. -1. **Generate API keys from Langfuse UI** (Settings → Project → API Keys) +## Tech Stack + +| Concern | Technology | +| --- | --- | +| Language / runtime | Python 3.12.10 | +| Package manager | [uv](https://docs.astral.sh/uv/) | +| API framework | FastAPI + Uvicorn (port `8100`) | +| LLM pipelines | DSPy | +| Safety | NeMo Guardrails | +| Vector database | Qdrant | +| Sessions & history | Redis | +| Analytics store | PostgreSQL + ClickHouse (Langfuse) | +| Object storage | MinIO (S3-compatible) | +| Secrets | HashiCorp Vault | +| Observability | Langfuse, Grafana, Loki | +| API gateway | Ruuter | + +## Quick Start + +### Prerequisites + +- Docker and Docker Compose +- [uv](https://docs.astral.sh/uv/) (for local development outside containers) + +### Run the stack -2. **Copy the script to vault container:** ```bash -docker cp store-langfuse-secrets.sh vault:/tmp/store-langfuse-secrets.sh +# Start the full stack (orchestration service + data stores + tooling) +docker compose up -d + +# Check the orchestration service health +curl http://localhost:8100/health ``` -3. **Execute the script with your API keys:** +### Environment configuration + +Configuration is supplied through environment files at the repository root: + +| File | Scope | +| --- | --- | +| `.env` | Shared infrastructure (storage, databases, Redis, Vault, feature flags) | +| `.env.llm_orchestration_service` | LLM Orchestration Service | +| `.env.gui` | Admin GUI | +| `.env.notification` | Notification server | + +Feature flags such as `TOOL_CLASSIFIER_ENABLED`, `SERVICE_WORKFLOW_ENABLED`, +`CONTEXT_WORKFLOW_ENABLED`, and `API_TOOL_CALLING_WORKFLOW_ENABLED` toggle individual workflows. + +## Components + +The orchestration service lives under `src/`. The main packages: + +| Component | Responsibility | Path | +| --- | --- | --- | +| Orchestration service & API | FastAPI app coordinating the pipeline | `src/llm_orchestration_service_api.py`, `src/llm_orchestration_service.py` | +| Tool classifier & workflows | Intent detection and per-workflow execution (Service / Context / RAG / API-Tool / OOD) | `src/tool_classifier/` | +| Contextual retrieval | Hybrid semantic + BM25 search with RRF fusion | `src/contextual_retrieval/` | +| Vector indexer | Document ingestion and embedding into Qdrant | `src/vector_indexer/` | +| API tool indexer | Indexing of API endpoints for tool calling | `src/api_tool_indexer/` | +| Intent data enrichment | LLM-generated enrichment of intent data | `src/intent_data_enrichment/` | +| Response generator | Grounded answer generation with citations | `src/response_generator/` | +| Prompt refiner | Query refinement for better retrieval | `src/prompt_refine_manager/` | +| Guardrails | NeMo Guardrails input/output safety | `src/guardrails/` | +| LLM configuration | Provider management, Vault credentials, feature flags | `src/llm_orchestrator_config/` | +| Optimization | DSPy-based tuning of guardrails, refiner, generator | `src/optimization/` | +| Utilities | Redis, sessions, rate limiting, streaming, cost, logging | `src/utils/` | + +## Core Workflows + +The tool classifier routes each query to one of: + +- **Service** — maps the query to a backend service/intent ([docs](./docs/TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md)) +- **Context** — greetings and conversation handling with Redis-backed history ([docs](./docs/CONTEXT_WORKFLOW_GREETING_DETECTION.md)) +- **RAG** — contextual retrieval then grounded generation ([docs](./docs/CONTEXTUAL_RETRIEVAL_FLOW.md)) +- **API-Tool** — agentic multi-endpoint API calling ([docs](./docs/API_TOOL_CALLING.md)) +- **OOD** — graceful fallback for out-of-domain queries + +Classification itself uses hybrid dense + sparse search — see +[docs/HYBRID_SEARCH_CLASSIFICATION.md](./docs/HYBRID_SEARCH_CLASSIFICATION.md). + +## Documentation + +The full documentation catalogue lives in **[docs/README.md](./docs/README.md)**, covering +architecture, retrieval, workflows, configuration/secrets, sessions, and the API reference. + +## Development + ```bash +# Install the pinned Python and dependencies +uv python install 3.12.10 +uv sync --frozen + +# Install pre-commit hooks +uv run pre-commit install + +# Run the test suite +uv run pytest tests/ -v +``` + +See **[CONTRIBUTING.md](./CONTRIBUTING.md)** for the full development workflow, tooling, and CI checks. + +## Deployment + +Several Docker Compose variants are provided for different environments: + +| File | Purpose | +| --- | --- | +| `docker-compose.yml` | Default full stack | +| `docker-compose-ec2.yml` | AWS EC2 deployment variant | +| `docker-compose-test.yml` | Integration testing | +| `docker-compose-eval.yml` | Evaluation / benchmarking | + +Kubernetes manifests and Helm charts are under [`kubernetes/`](./kubernetes), including +[`LANGFUSE_SETUP.md`](./kubernetes/LANGFUSE_SETUP.md) and +[`CONTAINER_REGISTRY_SETUP.md`](./kubernetes/CONTAINER_REGISTRY_SETUP.md). + +## API Reference + +HTTP endpoints for LLM connections, inference results, and chatbot inquiries are documented in +**[docs/API_REFERENCE.md](./docs/API_REFERENCE.md)**. + +## Configuration + +- **Environment files** — see [Environment configuration](#environment-configuration) above. +- **Feature flags** — `src/llm_orchestrator_config/feature_flags.py`. +- **Service configuration** — YAML configs under each module's `config/` directory (e.g. LLM + providers, contextual retrieval parameters, indexer settings, guardrails policies). +- **Secrets** — managed in HashiCorp Vault; see + [docs/LLM_CONFIG_VAULT_INTEGRATION.md](./docs/LLM_CONFIG_VAULT_INTEGRATION.md) and + [docs/VAULT_SETUP_AND_USAGE.md](./docs/VAULT_SETUP_AND_USAGE.md). + +### Storing Langfuse secrets + +Generate API keys in the Langfuse UI (**Settings → Project → API Keys**), then store them in Vault. + +For Docker Compose deployments, use the [`store-langfuse-secrets.sh`](./store-langfuse-secrets.sh) +script: + +```bash +# Copy the script into the vault container +docker cp store-langfuse-secrets.sh vault:/tmp/store-langfuse-secrets.sh + +# Run it with your Langfuse keys docker exec -e LANGFUSE_INIT_PROJECT_PUBLIC_KEY= \ -e LANGFUSE_INIT_PROJECT_SECRET_KEY= \ vault sh -c "chmod +x /tmp/store-langfuse-secrets.sh && /tmp/store-langfuse-secrets.sh" ``` + +For Kubernetes, see [kubernetes/LANGFUSE_SETUP.md](./kubernetes/LANGFUSE_SETUP.md). + +## Troubleshooting + +| Symptom | Where to look | +| --- | --- | +| Service unhealthy | `curl http://localhost:8100/health`; `docker compose ps`; `docker compose logs llm-orchestration-service` | +| Vault / credential errors | [docs/VAULT_SETUP_AND_USAGE.md](./docs/VAULT_SETUP_AND_USAGE.md) and [docs/VAULT_SECURITY_ARCHITECTURE.md](./docs/VAULT_SECURITY_ARCHITECTURE.md) | +| Missing conversation history | Redis connectivity (`REDIS_HOST`/`REDIS_AUTH`); [docs/REDIS_SESSION_STORE.md](./docs/REDIS_SESSION_STORE.md) | +| Retrieval returns nothing | Qdrant availability (port `6333`) and that indexing has run | +| Logs & traces | Grafana/Loki for logs; Langfuse for request traces and cost | + +## License + +This project is licensed under the terms in the [LICENSE](./LICENSE) file. + +## Links + +- [Bürokratt project](https://github.com/buerokratt) +- [Architecture](./docs/ARCHITECTURE.md) +- [Documentation index](./docs/README.md) +- [API reference](./docs/API_REFERENCE.md) +- [Contributing](./CONTRIBUTING.md) diff --git a/constants.ini b/constants.ini index 90c8ddcb..3d8a96d2 100644 --- a/constants.ini +++ b/constants.ini @@ -7,6 +7,7 @@ RAG_SEARCH_PROJECT_LAYER=rag-search RAG_SEARCH_TIM=http://tim:8085 RAG_SEARCH_CRON_MANAGER=http://cron-manager:9010 RAG_SEARCH_LLM_ORCHESTRATOR=http://llm-orchestration-service:8100/orchestrate +RAG_SEARCH_LLM_CACHE_CLEAR=http://llm-orchestration-service:8100/cache/clear RAG_SEARCH_PROMPT_REFRESH=http://llm-orchestration-service:8100/prompt-config/refresh DOMAIN=localhost DB_PASSWORD=dbadmin diff --git a/docker-compose-ec2.yml b/docker-compose-ec2.yml index a9052865..51f41549 100644 --- a/docker-compose-ec2.yml +++ b/docker-compose-ec2.yml @@ -133,8 +133,9 @@ services: - DEBUG_ENABLED=true - CHOKIDAR_USEPOLLING=true - PORT=3001 - - REACT_APP_SERVICE_ID=conversations,settings,monitoring + - REACT_APP_SERVICE_ID=settings,services,training,rag-search,knowledge-center - REACT_APP_ENABLE_HIDDEN_FEATURES=TRUE + - REACT_APP_MENU_JSON= [{"id":"conversations","label":{"et":"Vestlused","en":"Conversations"},"path":"/chat","children":[{"label":{"et":"Vastamata","en":"Unanswered"},"path":"/unanswered"},{"label":{"et":"Aktiivsed","en":"Active"},"path":"/active"},{"label":{"et":"Ootel","en":"Pending"},"path":"/pending"},{"label":{"et":"Ajalugu","en":"History"},"path":"/history"},{"label":{"et":"Valideerimised","en":"Validations"},"path":"/validations"}]},{"id":"training","label":{"et":"Treening","en":"Training"},"path":"/training","children":[{"label":{"et":"Treening","en":"Training"},"path":"/training","children":[{"label":{"et":"Teemad","en":"Themes"},"path":"/training/intents"},{"hidden":true,"label":{"et":"Avalikud teemad","en":"Public themes"},"path":"/training/common-intents"},{"label":{"et":"Teemade järeltreenimine","en":"Post training themes"},"path":"/training/intents-followup-training"},{"label":{"et":"Vastused","en":"Answers"},"path":"/training/responses"},{"label":{"et":"Reeglid","en":"Rules"},"path":"/training/rules"},{"hidden":true,"label":{"et":"Konfiguratsioon","en":"Configuration"},"path":"/training/configuration"},{"label":{"et":"Vormid","en":"Forms"},"path":"/training/forms"},{"label":{"et":"Mälukohad","en":"Slots"},"path":"/training/slots"}]},{"label":{"et":"Ajaloolised vestlused","en":"Historical conversations"},"path":"/history","children":[{"label":{"et":"Ajalugu","en":"History"},"path":"/history/history"},{"hidden":true,"label":{"et":"Pöördumised","en":"Appeals"},"path":"/history/appeal"}]},{"label":{"et":"Mudelipank ja analüütika","en":"Modelbank and analytics"},"path":"/analytics","children":[{"label":{"et":"Teemade ülevaade","en":"Overview of topics"},"path":"/analytics/overview"},{"label":{"et":"Mudelite võrdlus","en":"Comparison of models"},"path":"/analytics/models"},{"hidden":true,"label":{"et":"Testlood","en":"testTracks"},"path":"/analytics/testcases"}]},{"label":{"et":"Treeni uus mudel","en":"Train new model"},"path":"/train-new-model"}]},{"id":"analytics","label":{"et":"Analüütika","en":"Analytics"},"path":"/analytics","children":[{"label":{"et":"Ülevaade","en":"Overview"},"path":"/overview"},{"label":{"et":"Vestlused","en":"Chats"},"path":"/chats"},{"label":{"et":"Tagasiside","en":"Feedback"},"path":"/feedback"},{"label":{"et":"Avaandmed","en":"Reports"},"path":"/reports"}]},{"id":"services","hidden":true,"label":{"et":"Teenused","en":"Services"},"path":"/services","children":[{"label":{"et":"Ülevaade","en":"Overview"},"path":"/overview"},{"label":{"et":"Uus teenus","en":"New Service"},"path":"/newService"},{"label":{"et":"Probleemsed teenused","en":"Faulty Services"},"path":"/faultyServices"}]},{"id":"knowledge-center","label":{"et":"Teadmuskeskus","en":"Knowledge Center"},"path":"/knowledge-center","children":[{"id":"rag-search","label":{"et":"Mudelid ja seadistused","en":"Models management"},"path":"/rag-search","children":[{"label":{"et":"Mudelite ühendused","en":"LLM connections"},"path":"/llm-connections"},{"label":{"et":"Viiba Seaded","en":"Prompt Configurations"},"path":"/prompt-configurations"},{"label":{"et":"Testi mudelit","en":"Test LLM"},"path":"/test-llm"}]},{"id":"ckb","label":{"et":"Teadmusbaas","en":"Knowledge Base"},"path":"/ckb","children":[{"label":{"et":"Agentuur","en":"Agency"},"path":"/agency"},{"label":{"et":"Aruanded","en":"Reports"},"path":"/reports"},{"label":{"et":"API Integratsioonid","en":"API Integrations"},"path":"/api"}]}]},{"id":"settings","label":{"et":"Haldus","en":"Administration"},"path":"/settings","children":[{"label":{"et":"Kasutajad","en":"Users"},"path":"/users"},{"label":{"et":"Vestlusbot","en":"Chatbot"},"path":"/chatbot","children":[{"label":{"et":"Seaded","en":"Settings"},"path":"/chatbot/settings"},{"label":{"et":"Tervitussõnum","en":"Welcome message"},"path":"/chatbot/welcome-message"},{"label":{"et":"Välimus ja käitumine","en":"Appearance and behavior"},"path":"/chatbot/appearance"},{"label":{"et":"Erakorralised teated","en":"Emergency notices"},"path":"/chatbot/emergency-notices"},{"label":{"et":"Tagasiside","en":"Feedback"},"path":"/chatbot/feedback"}]},{"label":{"et":"Vestluste analüüs","en":"Chat analysis"},"path":"/chat-analysis"},{"label":{"et":"Asutuse tööaeg","en":"Office opening hours"},"path":"/working-time"},{"label":{"et":"Vestluse Kustutamine","en":"Delete Conversations"},"path":"/delete-conversations"},{"label":{"et":"Sessiooni pikkus","en":"Session length"},"path":"/session-length"},{"label":{"et":"SKMi konfiguratsioon","en":"SKM Configuration"},"path":"/skm-configuration"},{"label":{"et":"Anonümiseerija","en":"Anonymizer"},"path":"/anonymizer"}]},{"id":"monitoring","hidden":true,"label":{"et":"Seire","en":"Monitoring"},"path":"/monitoring","children":[{"label":{"et":"Aktiivaeg","en":"Working hours"},"path":"/uptime"}]}] - VITE_HOST=0.0.0.0 - VITE_PORT=3001 - HOST=0.0.0.0 @@ -183,6 +184,8 @@ services: - ./src/tool_classifier:/app/src/tool_classifier - ./src/intent_data_enrichment:/app/src/intent_data_enrichment - ./src/api_tool_indexer:/app/src/api_tool_indexer + - ./src/__init__.py:/app/src/__init__.py:ro + - ./src/loki_logger.py:/app/src/loki_logger.py:ro - ./src/utils/decrypt_vault_secrets.py:/app/src/utils/decrypt_vault_secrets.py:ro # Decryption utility (read-only) - cron_data:/app/data - shared-volume:/app/shared # Access to shared resources for cross-container coordination @@ -191,7 +194,7 @@ services: - ./.env:/app/.env:ro environment: - server.port=9010 - - PYTHONPATH=/app:/app/src/vector_indexer:/app/src/intent_data_enrichment:/app/src/api_tool_indexer + - PYTHONPATH=/app:/app/src:/app/src/vector_indexer:/app/src/intent_data_enrichment:/app/src/api_tool_indexer - VAULT_AGENT_URL=http://vault-agent-cron:8203 ports: - 9010:8080 @@ -502,7 +505,11 @@ services: - ./vault/config:/vault/config:ro - ./vault/logs:/vault/logs networks: - - vault-network # Only on vault-network for security + vault-network: # Only on vault-network for security + # Local testing: bare "vault" collides with the ckb stack on the shared + # bykstack network, so expose this Vault under a unique alias instead. + aliases: + - rag-vault restart: unless-stopped healthcheck: test: ["CMD", "sh", "-c", "wget -q -O- http://127.0.0.1:8200/v1/sys/health || exit 0"] @@ -512,14 +519,17 @@ services: start_period: 10s vault-init: - image: hashicorp/vault:1.20.3 + build: + context: . + dockerfile: Dockerfile.vault-init + image: rag-vault-init:1.20.3 container_name: vault-init user: "0" depends_on: vault: condition: service_healthy environment: - VAULT_ADDR: http://vault:8200 + VAULT_ADDR: http://rag-vault:8200 volumes: - vault-data:/vault/data - vault-agent-creds:/agent/credentials @@ -528,13 +538,12 @@ services: - vault-agent-llm-token:/agent/llm-token - ./vault-init.sh:/vault-init.sh:ro networks: - - vault-network # Access vault - - bykstack # Access to write agent tokens + # vault-network only: tokens/creds go via shared volumes, not the network. + - vault-network entrypoint: ["/bin/sh"] command: - -c - | - apk add --no-cache curl jq uuidgen openssl # Create and set permissions for all agent directories mkdir -p /agent/credentials /agent/gui-token /agent/cron-token /agent/llm-token /agent/out chown -R vault:vault /agent/credentials /agent/gui-token /agent/cron-token /agent/llm-token /agent/out @@ -638,6 +647,7 @@ services: - ./src/llm_config_module/config:/app/src/llm_config_module/config:ro - ./src/optimization/optimized_modules:/app/src/optimization/optimized_modules - llm_orchestration_logs:/app/logs + - ./grafana-configs/loki_logger.py:/app/src/loki_logger.py:ro networks: - bykstack depends_on: diff --git a/docker-compose-eval.yml b/docker-compose-eval.yml new file mode 100644 index 00000000..bbedf1ee --- /dev/null +++ b/docker-compose-eval.yml @@ -0,0 +1,291 @@ +services: + # === Core Infrastructure === + + # Shared PostgreSQL database (used by both application and Langfuse) + rag_search_db: + image: postgres:14.1 + container_name: rag_search_db + restart: always + environment: + POSTGRES_USER: postgres + POSTGRES_PASSWORD: dbadmin + POSTGRES_DB: rag-search + volumes: + - test_rag_search_db:/var/lib/postgresql/data + ports: + - "5436:5432" + networks: + - test-network + + # Vector database for RAG + qdrant: + image: qdrant/qdrant:v1.15.1 + container_name: qdrant + restart: always + ports: + - "6333:6333" + - "6334:6334" + volumes: + - test_qdrant_data:/qdrant/storage + networks: + - test-network + + # === Secret Management === + + # Vault - Secret management (dev mode) + vault: + image: hashicorp/vault:1.20.3 + container_name: vault + cap_add: + - IPC_LOCK + ports: + - "8200:8200" + environment: + VAULT_DEV_ROOT_TOKEN_ID: root + VAULT_ADDR: http://0.0.0.0:8200 + VAULT_API_ADDR: http://0.0.0.0:8200 + command: server -dev -dev-listen-address=0.0.0.0:8200 + networks: + - test-network + + # Vault Agent - Automatic token management via AppRole + vault-agent-llm: + image: hashicorp/vault:1.20.3 + container_name: vault-agent-llm + depends_on: + - vault + volumes: + - ./test-vault/agents/llm:/agent/in + - ./test-vault/agent-out:/agent/llm-token + entrypoint: ["sh", "-c"] + command: + - | + # Wait for Vault to be ready + sleep 5 + echo "Waiting for AppRole credentials..." + while [ ! -f /agent/in/role_id ] || [ ! -s /agent/in/role_id ]; do + sleep 1 + done + while [ ! -f /agent/in/secret_id ] || [ ! -s /agent/in/secret_id ]; do + sleep 1 + done + echo "Credentials found, starting Vault Agent..." + exec vault agent -config=/agent/in/agent.hcl -log-level=debug + networks: + - test-network + + # === Langfuse Observability Stack === + + # Redis - Queue and cache for Langfuse + redis: + image: redis:7 + container_name: redis + restart: always + command: --requirepass myredissecret + ports: + - "127.0.0.1:6379:6379" + networks: + - test-network + + # MinIO - S3-compatible storage for Langfuse + minio: + image: minio/minio:latest + container_name: minio + restart: always + entrypoint: sh + command: -c "mkdir -p /data/langfuse && minio server /data --address ':9000' --console-address ':9001'" + environment: + MINIO_ROOT_USER: minio + MINIO_ROOT_PASSWORD: miniosecret + ports: + - "9090:9000" + - "127.0.0.1:9091:9001" + volumes: + - test_minio_data:/data + networks: + - test-network + + # ClickHouse - Analytics database for Langfuse (REQUIRED in v3) + clickhouse: + image: clickhouse/clickhouse-server:24.3 + container_name: clickhouse + restart: always + environment: + CLICKHOUSE_DB: default + CLICKHOUSE_USER: default + CLICKHOUSE_PASSWORD: clickhouse + volumes: + - test_clickhouse_data:/var/lib/clickhouse + ports: + - "127.0.0.1:8123:8123" + - "127.0.0.1:9000:9000" + networks: + - test-network + ulimits: + nofile: + soft: 262144 + hard: 262144 + + # Langfuse Worker - Background job processor + langfuse-worker: + image: langfuse/langfuse-worker:3 + container_name: langfuse-worker + restart: always + depends_on: + - rag_search_db + - minio + - redis + - clickhouse + ports: + - "127.0.0.1:3030:3030" + environment: + # Database + DATABASE_URL: postgresql://postgres:dbadmin@rag_search_db:5432/rag-search + + # Auth & Security (TEST VALUES ONLY - NOT FOR PRODUCTION) + # gitleaks:allow - These are test-only hex strings + NEXTAUTH_URL: http://localhost:3000 + SALT: ${SALT} + ENCRYPTION_KEY: ${ENCRYPTION_KEY} + + # Features + TELEMETRY_ENABLED: "false" + LANGFUSE_ENABLE_EXPERIMENTAL_FEATURES: "false" + + # ClickHouse (REQUIRED for Langfuse v3) + CLICKHOUSE_MIGRATION_URL: clickhouse://clickhouse:9000/default + CLICKHOUSE_URL: http://clickhouse:8123 + CLICKHOUSE_USER: default + CLICKHOUSE_PASSWORD: clickhouse + CLICKHOUSE_CLUSTER_ENABLED: "false" + + # S3/MinIO Event Upload + LANGFUSE_S3_EVENT_UPLOAD_BUCKET: langfuse + LANGFUSE_S3_EVENT_UPLOAD_REGION: us-east-1 + LANGFUSE_S3_EVENT_UPLOAD_ACCESS_KEY_ID: minio + LANGFUSE_S3_EVENT_UPLOAD_SECRET_ACCESS_KEY: miniosecret + LANGFUSE_S3_EVENT_UPLOAD_ENDPOINT: http://minio:9000 + LANGFUSE_S3_EVENT_UPLOAD_FORCE_PATH_STYLE: "true" + + # S3/MinIO Media Upload + LANGFUSE_S3_MEDIA_UPLOAD_BUCKET: langfuse + LANGFUSE_S3_MEDIA_UPLOAD_REGION: us-east-1 + LANGFUSE_S3_MEDIA_UPLOAD_ACCESS_KEY_ID: minio + LANGFUSE_S3_MEDIA_UPLOAD_SECRET_ACCESS_KEY: miniosecret + LANGFUSE_S3_MEDIA_UPLOAD_ENDPOINT: http://minio:9000 + LANGFUSE_S3_MEDIA_UPLOAD_FORCE_PATH_STYLE: "true" + + # Redis + REDIS_HOST: redis + REDIS_PORT: "6379" + REDIS_AUTH: myredissecret + networks: + - test-network + + # Langfuse Web - UI and API + langfuse-web: + image: langfuse/langfuse:3 + container_name: langfuse-web + restart: always + depends_on: + - langfuse-worker + - rag_search_db + - clickhouse + ports: + - "3000:3000" + environment: + # Database + DATABASE_URL: postgresql://postgres:dbadmin@rag_search_db:5432/rag-search + + # Auth & Security (TEST VALUES ONLY - NOT FOR PRODUCTION) + # gitleaks:allow - These are test-only hex strings + NEXTAUTH_URL: http://localhost:3000 + NEXTAUTH_SECRET: ${NEXTAUTH_SECRET} + SALT: ${SALT} + ENCRYPTION_KEY: ${ENCRYPTION_KEY} + + # Features + TELEMETRY_ENABLED: "false" + LANGFUSE_ENABLE_EXPERIMENTAL_FEATURES: "false" + + # ClickHouse (REQUIRED for Langfuse v3) + CLICKHOUSE_MIGRATION_URL: clickhouse://clickhouse:9000/default + CLICKHOUSE_URL: http://clickhouse:8123 + CLICKHOUSE_USER: default + CLICKHOUSE_PASSWORD: clickhouse + CLICKHOUSE_CLUSTER_ENABLED: "false" + + # S3/MinIO Event Upload + LANGFUSE_S3_EVENT_UPLOAD_BUCKET: langfuse + LANGFUSE_S3_EVENT_UPLOAD_REGION: us-east-1 + LANGFUSE_S3_EVENT_UPLOAD_ACCESS_KEY_ID: minio + LANGFUSE_S3_EVENT_UPLOAD_SECRET_ACCESS_KEY: miniosecret + LANGFUSE_S3_EVENT_UPLOAD_ENDPOINT: http://minio:9000 + LANGFUSE_S3_EVENT_UPLOAD_FORCE_PATH_STYLE: "true" + + # S3/MinIO Media Upload + LANGFUSE_S3_MEDIA_UPLOAD_BUCKET: langfuse + LANGFUSE_S3_MEDIA_UPLOAD_REGION: us-east-1 + LANGFUSE_S3_MEDIA_UPLOAD_ACCESS_KEY_ID: minio + LANGFUSE_S3_MEDIA_UPLOAD_SECRET_ACCESS_KEY: miniosecret + LANGFUSE_S3_MEDIA_UPLOAD_ENDPOINT: http://minio:9000 + LANGFUSE_S3_MEDIA_UPLOAD_FORCE_PATH_STYLE: "true" + + # Redis + REDIS_HOST: redis + REDIS_PORT: "6379" + REDIS_AUTH: myredissecret + + # Initialize test project with known credentials + LANGFUSE_INIT_PROJECT_PUBLIC_KEY: pk-lf-test + LANGFUSE_INIT_PROJECT_SECRET_KEY: sk-lf-test + networks: + - test-network + + # === LLM Orchestration Service === + + llm-orchestration-service: + build: + context: . + dockerfile: Dockerfile.llm_orchestration_service + container_name: llm-orchestration-service + restart: always + ports: + - "8100:8100" + environment: + - VAULT_ADDR=http://vault:8200 + - VAULT_TOKEN_FILE=/agent/llm-token/token + - ENVIRONMENT=development + - QDRANT_URL=http://qdrant:6333 + - EVAL_MODE=true + volumes: + - ./src/llm_config_module/config:/app/src/llm_config_module/config:ro + - ./test-vault/agent-out:/agent/llm-token:ro + - test_llm_orchestration_logs:/app/logs + depends_on: + - qdrant + - langfuse-web + - vault-agent-llm + networks: + - test-network + +# === Networks === + +networks: + test-network: + name: test-network + driver: bridge + +# === Volumes === + +volumes: + test_rag_search_db: + name: test_rag_search_db + test_qdrant_data: + name: test_qdrant_data + test_minio_data: + name: test_minio_data + test_clickhouse_data: + name: test_clickhouse_data + test_llm_orchestration_logs: + name: test_llm_orchestration_logs \ No newline at end of file diff --git a/docker-compose-test.yml b/docker-compose-test.yml index a0c56074..645e3b61 100644 --- a/docker-compose-test.yml +++ b/docker-compose-test.yml @@ -147,6 +147,9 @@ services: volumes: - ./test-vault/agents/llm:/agent/in - ./test-vault/agent-out:/agent/out + # agent.hcl writes the token/pidfile/dummy to /agent/llm-token; map it to the + # same host dir as /agent/out so the host and other services see the token. + - ./test-vault/agent-out:/agent/llm-token entrypoint: ["sh", "-c"] command: - | @@ -345,14 +348,16 @@ services: environment: # Infrastructure connections - VAULT_ADDR=http://vault:8200 - - VAULT_TOKEN_FILE=/agent/out/token + # VaultAgentClient reads the token from /agent/llm-token/token by default + # (src/llm_orchestrator_config/vault/vault_client.py); mount must match. + - VAULT_TOKEN_FILE=/agent/llm-token/token - QDRANT_URL=http://qdrant:6333 - EVAL_MODE=true # Disable OpenTelemetry tracing in test environment - OTEL_SDK_DISABLED=true volumes: - ./src/llm_config_module/config:/app/src/llm_config_module/config:ro - - ./test-vault/agent-out:/agent/out:ro + - ./test-vault/agent-out:/agent/llm-token:ro - test_llm_orchestration_logs:/app/logs depends_on: - qdrant diff --git a/docker-compose.yml b/docker-compose.yml index 7fb5b0fb..71a0028d 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -87,7 +87,7 @@ services: volumes: - ./tim-db:/var/lib/postgresql/data ports: - - 9875:5432 + - 9876:5432 networks: - bykstack @@ -132,8 +132,10 @@ services: - DEBUG_ENABLED=true - CHOKIDAR_USEPOLLING=true - PORT=3001 - - REACT_APP_SERVICE_ID=conversations,settings,monitoring + - REACT_APP_SERVICE_ID=settings,services,training,rag-search,knowledge-center - REACT_APP_ENABLE_HIDDEN_FEATURES=TRUE + - REACT_APP_ENABLE_MULTI_DOMAIN=FALSE + - REACT_APP_MENU_JSON= [{"id":"conversations","label":{"et":"Vestlused","en":"Conversations"},"path":"/chat","children":[{"label":{"et":"Vastamata","en":"Unanswered"},"path":"/unanswered"},{"label":{"et":"Aktiivsed","en":"Active"},"path":"/active"},{"label":{"et":"Ootel","en":"Pending"},"path":"/pending"},{"label":{"et":"Ajalugu","en":"History"},"path":"/history"},{"label":{"et":"Valideerimised","en":"Validations"},"path":"/validations"}]},{"id":"training","label":{"et":"Treening","en":"Training"},"path":"/training","children":[{"label":{"et":"Treening","en":"Training"},"path":"/training","children":[{"label":{"et":"Teemad","en":"Themes"},"path":"/training/intents"},{"hidden":true,"label":{"et":"Avalikud teemad","en":"Public themes"},"path":"/training/common-intents"},{"label":{"et":"Teemade järeltreenimine","en":"Post training themes"},"path":"/training/intents-followup-training"},{"label":{"et":"Vastused","en":"Answers"},"path":"/training/responses"},{"label":{"et":"Reeglid","en":"Rules"},"path":"/training/rules"},{"hidden":true,"label":{"et":"Konfiguratsioon","en":"Configuration"},"path":"/training/configuration"},{"label":{"et":"Vormid","en":"Forms"},"path":"/training/forms"},{"label":{"et":"Mälukohad","en":"Slots"},"path":"/training/slots"}]},{"label":{"et":"Ajaloolised vestlused","en":"Historical conversations"},"path":"/history","children":[{"label":{"et":"Ajalugu","en":"History"},"path":"/history/history"},{"hidden":true,"label":{"et":"Pöördumised","en":"Appeals"},"path":"/history/appeal"}]},{"label":{"et":"Mudelipank ja analüütika","en":"Modelbank and analytics"},"path":"/analytics","children":[{"label":{"et":"Teemade ülevaade","en":"Overview of topics"},"path":"/analytics/overview"},{"label":{"et":"Mudelite võrdlus","en":"Comparison of models"},"path":"/analytics/models"},{"hidden":true,"label":{"et":"Testlood","en":"testTracks"},"path":"/analytics/testcases"}]},{"label":{"et":"Treeni uus mudel","en":"Train new model"},"path":"/train-new-model"}]},{"id":"analytics","label":{"et":"Analüütika","en":"Analytics"},"path":"/analytics","children":[{"label":{"et":"Ülevaade","en":"Overview"},"path":"/overview"},{"label":{"et":"Vestlused","en":"Chats"},"path":"/chats"},{"label":{"et":"Tagasiside","en":"Feedback"},"path":"/feedback"},{"label":{"et":"Avaandmed","en":"Reports"},"path":"/reports"}]},{"id":"services","hidden":true,"label":{"et":"Teenused","en":"Services"},"path":"/services","children":[{"label":{"et":"Ülevaade","en":"Overview"},"path":"/overview"},{"label":{"et":"Uus teenus","en":"New Service"},"path":"/newService"},{"label":{"et":"Probleemsed teenused","en":"Faulty Services"},"path":"/faultyServices"}]},{"id":"knowledge-center","label":{"et":"Teadmuskeskus","en":"Knowledge Center"},"path":"/knowledge-center","children":[{"id":"rag-search","label":{"et":"Mudelid ja seadistused","en":"Models management"},"path":"/rag-search","children":[{"label":{"et":"Mudelite ühendused","en":"LLM connections"},"path":"/llm-connections"},{"label":{"et":"Viiba Seaded","en":"Prompt Configurations"},"path":"/prompt-configurations"},{"label":{"et":"Testi mudelit","en":"Test LLM"},"path":"/test-llm"}]},{"id":"ckb","label":{"et":"Teadmusbaas","en":"Knowledge Base"},"path":"/ckb","children":[{"label":{"et":"Agentuur","en":"Agency"},"path":"/agency"},{"label":{"et":"Aruanded","en":"Reports"},"path":"/reports"},{"label":{"et":"API Integratsioonid","en":"API Integrations"},"path":"/api"}]}]},{"id":"settings","label":{"et":"Haldus","en":"Administration"},"path":"/settings","children":[{"label":{"et":"Kasutajad","en":"Users"},"path":"/users"},{"label":{"et":"Vestlusbot","en":"Chatbot"},"path":"/chatbot","children":[{"label":{"et":"Seaded","en":"Settings"},"path":"/chatbot/settings"},{"label":{"et":"Tervitussõnum","en":"Welcome message"},"path":"/chatbot/welcome-message"},{"label":{"et":"Välimus ja käitumine","en":"Appearance and behavior"},"path":"/chatbot/appearance"},{"label":{"et":"Erakorralised teated","en":"Emergency notices"},"path":"/chatbot/emergency-notices"},{"label":{"et":"Tagasiside","en":"Feedback"},"path":"/chatbot/feedback"}]},{"label":{"et":"Vestluste analüüs","en":"Chat analysis"},"path":"/chat-analysis"},{"label":{"et":"Asutuse tööaeg","en":"Office opening hours"},"path":"/working-time"},{"label":{"et":"Vestluse Kustutamine","en":"Delete Conversations"},"path":"/delete-conversations"},{"label":{"et":"Sessiooni pikkus","en":"Session length"},"path":"/session-length"},{"label":{"et":"SKMi konfiguratsioon","en":"SKM Configuration"},"path":"/skm-configuration"},{"label":{"et":"Anonümiseerija","en":"Anonymizer"},"path":"/anonymizer"}]},{"id":"monitoring","hidden":true,"label":{"et":"Seire","en":"Monitoring"},"path":"/monitoring","children":[{"label":{"et":"Aktiivaeg","en":"Working hours"},"path":"/uptime"}]}] - VITE_HOST=0.0.0.0 - VITE_PORT=3001 - HOST=0.0.0.0 @@ -182,6 +184,8 @@ services: - ./src/tool_classifier:/app/src/tool_classifier - ./src/intent_data_enrichment:/app/src/intent_data_enrichment - ./src/api_tool_indexer:/app/src/api_tool_indexer + - ./src/__init__.py:/app/src/__init__.py:ro + - ./src/loki_logger.py:/app/src/loki_logger.py:ro - ./src/utils/decrypt_vault_secrets.py:/app/src/utils/decrypt_vault_secrets.py:ro # Decryption utility (read-only) - cron_data:/app/data - shared-volume:/app/shared # Access to shared resources for cross-container coordination @@ -190,8 +194,8 @@ services: - ./.env:/app/.env:ro environment: - server.port=9010 - - PYTHONPATH=/app:/app/src/vector_indexer:/app/src/intent_data_enrichment:/app/src/api_tool_indexer - - VAULT_AGENT_URL=http://vault-agent-cron:8203 + - PYTHONPATH=/app:/app/src:/app/src/vector_indexer:/app/src/intent_data_enrichment:/app/src/api_tool_indexer + - vaultAgentUrl=http://vault-agent-cron:8203 ports: - 9010:8080 depends_on: @@ -449,7 +453,11 @@ services: - ./vault/config:/vault/config:ro - ./vault/logs:/vault/logs networks: - - vault-network # Only on vault-network for security + vault-network: # Only on vault-network for security + # Local testing: bare "vault" collides with the ckb stack on the shared + # bykstack network, so expose this Vault under a unique alias instead. + aliases: + - rag-vault restart: unless-stopped healthcheck: test: ["CMD", "sh", "-c", "wget -q -O- http://127.0.0.1:8200/v1/sys/health || exit 0"] @@ -459,14 +467,17 @@ services: start_period: 10s vault-init: - image: hashicorp/vault:1.20.3 + build: + context: . + dockerfile: Dockerfile.vault-init + image: rag-vault-init:1.20.3 container_name: vault-init user: "0" depends_on: vault: condition: service_healthy environment: - VAULT_ADDR: http://vault:8200 + VAULT_ADDR: http://rag-vault:8200 volumes: - vault-data:/vault/data - vault-agent-creds:/agent/credentials @@ -475,13 +486,12 @@ services: - vault-agent-llm-token:/agent/llm-token - ./vault-init.sh:/vault-init.sh:ro networks: - - vault-network # Access vault - - bykstack # Access to write agent tokens + # vault-network only: tokens/creds go via shared volumes, not the network. + - vault-network entrypoint: ["/bin/sh"] command: - -c - | - apk add --no-cache curl jq uuidgen openssl # Create and set permissions for all agent directories mkdir -p /agent/credentials /agent/gui-token /agent/cron-token /agent/llm-token /agent/out chown -R vault:vault /agent/credentials /agent/gui-token /agent/cron-token /agent/llm-token /agent/out @@ -586,6 +596,7 @@ services: - ./src/optimization/optimized_modules:/app/src/optimization/optimized_modules - llm_orchestration_logs:/app/logs - ./tests:/app/tests # mount tests directory (excluded from image via .dockerignore) + - ./grafana-configs/loki_logger.py:/app/src/loki_logger.py:ro networks: - bykstack depends_on: diff --git a/docs/API_REFERENCE.md b/docs/API_REFERENCE.md new file mode 100644 index 00000000..9aed52cc --- /dev/null +++ b/docs/API_REFERENCE.md @@ -0,0 +1,1760 @@ +# API Reference + +This document is the consolidated HTTP API reference for the LLM Module. It covers the **LLM +Connections** management endpoints, the **Inference Results** storage/retrieval endpoints, and the +chatbot **Inquiry** endpoint exposed to the LLM Orchestration Service. + +> Routing note: the public-facing paths below are served through the Ruuter API gateway +> (`ruuter-private` / `ruuter-public`), which proxies to the LLM Orchestration Service. See +> [ARCHITECTURE.md](./ARCHITECTURE.md) for how requests flow through the system. + +## Contents + +- [LLM Connections API](#llm-connections-api-endpoints) +- [Inference Results API](#inference-results-api-endpoints) + +--- + +## LLM Connections API Endpoints + +### Base URL +``` +/ruuter-private/llm/connections +``` + +--- + +## 1. Create LLM Connection + +### Endpoint +```http +POST /ruuter-private/llm/connections/create +``` + +### Request Body +```json +{ + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "deploymentEnvironment": "Testing", + // Azure credentials (optional) + "deploymentName": "my-deployment", + "targetUri": "https://my-endpoint.azure.com", + "apiKey": "azure-api-key", + // AWS Bedrock credentials (optional) + "secretKey": "aws-secret-key", + "accessKey": "aws-access-key", + // Embedding model credentials (optional) + "embeddingModelApiKey": "embedding-api-key" +} +``` + +### Response (201 Created) +```json +{ + "id": 1, + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "usedBudget": 0.00, + "deploymentEnvironment": "Testing", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + // Azure credentials (if provided) + "deploymentName": "my-deployment", + "targetUri": "https://my-endpoint.azure.com", + "apiKey": "azure-api-key", + // AWS Bedrock credentials (if provided) + "secretKey": "aws-secret-key", + "accessKey": "aws-access-key", + // Embedding model credentials (if provided) + "embeddingModelApiKey": "embedding-api-key" +} +``` + +--- + +## 2. Update LLM Connection + +### Endpoint +```http +POST /ruuter-private/llm/connections/update +``` + +### Request Body +```json +{ + "connectionId": 1, + "llmPlatform": "Azure AI", + "llmModel": "GPT-4o-mini", + "embeddingPlatform": "Azure AI", + "embeddingModel": "text-embedding-ada-002", + "monthlyBudget": 2000.00, + "deploymentEnvironment": "Production", + // Azure credentials (optional) + "deploymentName": "updated-deployment", + "targetUri": "https://updated-endpoint.azure.com", + "apiKey": "updated-azure-api-key", + // AWS Bedrock credentials (optional) + "secretKey": "updated-aws-secret-key", + "accessKey": "updated-aws-access-key", + // Embedding model credentials (optional) + "embeddingModelApiKey": "updated-embedding-api-key" +} +``` + +### Response (200 OK) +```json +{ + "id": 1, + "llmPlatform": "Azure AI", + "llmModel": "GPT-4o-mini", + "embeddingPlatform": "Azure AI", + "embeddingModel": "text-embedding-ada-002", + "monthlyBudget": 2000.00, + "usedBudget": 150.75, + "deploymentEnvironment": "Production", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + // Azure credentials (if provided) + "deploymentName": "updated-deployment", + "targetUri": "https://updated-endpoint.azure.com", + "apiKey": "updated-azure-api-key", + // AWS Bedrock credentials (if provided) + "secretKey": "updated-aws-secret-key", + "accessKey": "updated-aws-access-key", + // Embedding model credentials (if provided) + "embeddingModelApiKey": "updated-embedding-api-key" +} +``` + +--- + +## 3. Get LLM Connections (Paginated List) + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/list +``` + +### Request Body +```json +{ + "page": 1, + "page_size": 10, + "sorting": "created_at desc" +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | Default | +|-----------|------|----------|-------------|---------| +| `page` | number | No | Page number (1-based) | 1 | +| `page_size` | number | No | Number of items per page | 10 | +| `sorting` | string | No | Sorting criteria | "created_at desc" | + +### Sorting Options +- `llm_platform asc/desc` +- `llm_model asc/desc` +- `embedding_platform asc/desc` +- `embedding_model asc/desc` +- `monthly_budget asc/desc` +- `environment asc/desc` +- `status asc/desc` +- `created_at asc/desc` +- `updated_at asc/desc` + +### Response (200 OK) +```json +[ + { + "id": 1, + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "environment": "Testing", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + "updatedAt": "2025-09-02T10:15:30.000Z", + "totalPages": 3 + }, + { + "id": 2, + "llmPlatform": "Azure AI", + "llmModel": "GPT-4o-mini", + "embeddingPlatform": "Azure AI", + "embeddingModel": "Ada-200-1", + "monthlyBudget": 2000.00, + "environment": "Production", + "status": "active", + "createdAt": "2025-09-02T09:30:15.000Z", + "updatedAt": "2025-09-02T11:00:00.000Z", + "totalPages": 3 + } +] +``` + +--- + +## 4. Get Single LLM Connection + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/get +``` + +### Request Body +```json +{ + "connection_id": 1 +} +``` + +### Response (200 OK) +```json +{ + "id": 1, + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "environment": "Testing", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + "updatedAt": "2025-09-02T10:15:30.000Z" +} +``` + +### Response (404 Not Found) +```json +"error: connection not found" +``` + +--- + +## 5. Add New LLM Connection + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/add +``` + +### Request Body +```json +{ + "llm_platform": "OpenAI", + "llm_model": "GPT-4o", + "embedding_platform": "OpenAI", + "embedding_model": "text-embedding-3-small", + "monthly_budget": 1000.00, + "environment": "Testing" +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `llm_platform` | string | Yes | LLM platform (e.g., "Azure AI", "OpenAI") | +| `llm_model` | string | Yes | LLM model (e.g., "GPT-4o") | +| `embedding_platform` | string | Yes | Embedding platform | +| `embedding_model` | string | Yes | Embedding model | +| `monthly_budget` | number | Yes | Monthly budget amount | +| `environment` | string | Yes | "Testing" or "Production" | + +### Response (200 OK) +```json +{ + "id": 3, + "llm_platform": "OpenAI", + "llm_model": "GPT-4o", + "embedding_platform": "OpenAI", + "embedding_model": "text-embedding-3-small", + "monthly_budget": 1000.00, + "environment": "Testing", + "status": "active", + "created_at": "2025-09-02T12:00:00.000Z", + "updated_at": "2025-09-02T12:00:00.000Z" +} +``` + +### Response (400 Bad Request) +```json +"error: environment must be 'Testing' or 'Production'" +``` + +--- + +## 6. Update LLM Connection + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/edit +``` + +### Request Body +```json +{ + "connection_id": 1, + "llm_platform": "Azure AI", + "llm_model": "GPT-4o-mini", + "embedding_platform": "Azure AI", + "embedding_model": "Ada-200-1", + "monthly_budget": 2000.00, + "environment": "Production" +} +``` + +### Response (200 OK) +```json +{ + "id": 1, + "llm_platform": "Azure AI", + "llm_model": "GPT-4o-mini", + "embedding_platform": "Azure AI", + "embedding_model": "Ada-200-1", + "monthly_budget": 2000.00, + "environment": "Production", + "status": "active", + "created_at": "2025-09-02T10:15:30.000Z", + "updated_at": "2025-09-02T12:30:00.000Z" +} +``` + +### Response (404 Not Found) +```json +"error: connection not found" +``` + +--- + +## 7. Delete LLM Connection + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/delete +``` + +### Request Body +```json +{ + "connection_id": 1 +} +``` + +### Response (200 OK) +```json +"LLM connection deleted successfully" +``` + +### Response (404 Not Found) +```json +"error: connection not found" +``` + +--- + +## 4. List All LLM Connections + +### Endpoint +```http +GET /ruuter-private/llm/connections/list +``` + +### Query Parameters (Optional for filtering) +| Parameter | Type | Description | +|-----------|------|-------------| +| `llmPlatform` | `string` | Filter by LLM platform | +| `llmModel` | `string` | Filter by LLM model | +| `deploymentEnvironment` | `string` | Filter by environment (Testing / Production) | +| `pageNumber` | `number` | Page number (1-based) | +| `pageSize` | `number` | Number of items per page | +| `sortBy` | `string` | Field to sort by | +| `sortOrder` | `string` | Sort order: 'asc' or 'desc' | + +### Example Request +```http +GET /ruuter-private/llm/connections/list?llmPlatform=OpenAI&deploymentEnvironment=Testing&model=GPT4 +``` + +--- + +## 5. Get Production LLM Connection (with filters) + +### Endpoint +```http +GET /ruuter-private/llm/connections/production +``` + +### Query Parameters (Optional for filtering) +| Parameter | Type | Description | +|-----------|------|-------------| +| `llmPlatform` | `string` | Filter by LLM platform | +| `llmModel` | `string` | Filter by LLM model | +| `embeddingPlatform` | `string` | Filter by embedding platform | +| `embeddingModel` | `string` | Filter by embedding model | +| `connectionStatus` | `string` | Filter by connection status | +| `sortBy` | `string` | Field to sort by | +| `sortOrder` | `string` | Sort order: 'asc' or 'desc' | + +### Example Request +```http +GET /ruuter-private/llm/connections/production?llmPlatform=OpenAI&connectionStatus=active +``` + +### Response (200 OK) +```json +[ + { + "id": 1, + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "deploymentEnvironment": "Testing", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + "updatedAt": "2025-09-02T10:15:30.000Z" + } +] +``` + +--- + +## 5. Get Single LLM Connection + +### Endpoint +```http +GET /ruuter-private/llm/connections/overview +``` + +### Response (200 OK) +```json +{ + "id": 1, + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "deploymentEnvironment": "Testing", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + "updatedAt": "2025-09-02T10:15:30.000Z" +} +``` + +--- + +## 6. Check if LLM Connection Exists + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/exists +``` + +### Request Body +```json +{ "connection_id": 1 } +``` + +### Response (200 OK) +```json +"true" +``` +or +```json +"false" +``` + +--- + +## 7. Update LLM Connection Status + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/update-status +``` + +### Request Body +```json +{ + "connection_id": 1, + "connection_status": "inactive" +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `connection_id` | number | Yes | LLM connection ID | +| `connection_status` | string | Yes | `"active"` or `"inactive"` | + +### Response (200 OK) +Returns the updated connection object. + +### Response (400 Bad Request) +```json +"error: connection_status must be 'active' or 'inactive'" +``` + +### Response (404 Not Found) +```json +"error: connection not found" +``` + +--- + +## 8. List LLM Connections — GET (Paginated, with Filters) + +### Endpoint +```http +GET /ruuter-private/rag-search/llm-connections/list +``` + +### Query Parameters +| Parameter | Type | Required | Default | Description | +|-----------|------|----------|---------|-------------| +| `pageNumber` | number | No | `1` | Page number (1-based) | +| `pageSize` | number | No | `10` | Items per page (1–100) | +| `sortBy` | string | No | `"created_at"` | Field to sort by | +| `sortOrder` | string | No | `"desc"` | `"asc"` or `"desc"` | +| `llmPlatform` | string | No | `""` | Filter by LLM platform | +| `llmModel` | string | No | `""` | Filter by LLM model | +| `environment` | string | No | `""` | Filter by environment | + +### Example Request +```http +GET /ruuter-private/rag-search/llm-connections/list?pageNumber=1&pageSize=10&llmPlatform=OpenAI +``` + +### Response (200 OK) +```json +[ + { + "id": 1, + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "environment": "Testing", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + "updatedAt": "2025-09-02T10:15:30.000Z", + "totalPages": 3 + } +] +``` + +### Response (400 Bad Request) +```json +"Page number must be greater than 0" +``` + +--- + +## 9. List All LLM Connections — GET (Paginated, with Filters) + +### Endpoint +```http +GET /ruuter-private/rag-search/llm-connections/all +``` + +Same as endpoint 8 above but queries all connections regardless of status. Accepts the same query parameters. + +--- + +## 10. Get Production LLM Connection — GET (with Filters) + +### Endpoint +```http +GET /ruuter-private/rag-search/llm-connections/production +``` + +### Query Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `llmPlatform` | string | No | Filter by LLM platform | +| `llmModel` | string | No | Filter by LLM model | +| `embeddingPlatform` | string | No | Filter by embedding platform | +| `embeddingModel` | string | No | Filter by embedding model | +| `connectionStatus` | string | No | Filter by connection status | +| `sortBy` | string | No | Field to sort by (default: `"created_at"`) | +| `sortOrder` | string | No | `"asc"` or `"desc"` (default: `"desc"`) | + +### Example Request +```http +GET /ruuter-private/rag-search/llm-connections/production?connectionStatus=active +``` + +### Response (200 OK) +```json +[ + { + "id": 1, + "llmPlatform": "OpenAI", + "llmModel": "GPT-4o", + "embeddingPlatform": "OpenAI", + "embeddingModel": "text-embedding-3-small", + "monthlyBudget": 1000.00, + "environment": "Production", + "status": "active", + "createdAt": "2025-09-02T10:15:30.000Z", + "updatedAt": "2025-09-02T10:15:30.000Z" + } +] +``` + +--- + +## 11. Update Used Budget for a Connection + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/cost/update +``` + +Adds `usage` to the connection's current `used_budget`. If `disconnectOnBudgetExceed` is set and the stop threshold is reached, the connection is automatically deactivated. + +### Request Body +```json +{ + "connection_id": 1, + "usage": 12.50 +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `connection_id` | number | Yes | LLM connection ID | +| `usage` | number | Yes | Amount to add to `used_budget` (≥ 0) | + +### Response (200 OK) — within budget +```json +{ + "data": { "id": 1, "usedBudget": 162.50, "monthlyBudget": 1000.00 }, + "budgetExceeded": false, + "message": "Used budget updated successfully", + "operationSuccess": true, + "statusCode": 200 +} +``` + +### Response (200 OK) — budget exceeded, connection deactivated +```json +{ + "data": { "id": 1, "usedBudget": 1005.00, "status": "inactive" }, + "budgetExceeded": true, + "message": "Used budget updated successfully. Connection deactivated due to budget threshold exceeded.", + "operationSuccess": true, + "statusCode": 200 +} +``` + +### Response (400 Bad Request) +```json +"error: connection_id and usage (>= 0) are required" +``` + +### Response (404 Not Found) +```json +"error: connection not found" +``` + +--- + +## 12. Check Budget Usage + +### Endpoint +```http +POST /ruuter-private/rag-search/llm-connections/usage/check +``` + +Returns whether the connection's budget is within the stop threshold, exceeded (not disconnected), or exceeded with disconnection. + +### Request Body +```json +{ "connection_id": 1 } +``` + +### Response (200 OK) — within budget +```json +{ + "isBudgetExceed": false, + "isLLMConnectionDisconnected": false +} +``` + +### Response (200 OK) — exceeded, not disconnected +```json +{ + "isBudgetExceed": true, + "isLLMConnectionDisconnected": false +} +``` + +### Response (200 OK) — exceeded and disconnected +```json +{ + "isBudgetExceed": true, + "isLLMConnectionDisconnected": true +} +``` + +### Response (404 Not Found) +```json +"Connection not found" +``` + +--- + +## 13. Check Budget Thresholds for Production Connection + +### Endpoint +```http +GET /ruuter-private/rag-search/llm-connections/cost/check +``` + +Returns warn/stop threshold status for the active production connection. + +### Response (200 OK) +```json +{ + "data": { + "id": 1, + "monthlyBudget": 1000.00, + "usedBudget": 620.00, + "warnBudgetThreshold": 70, + "stopBudgetThreshold": 90 + }, + "used_budget_percentage": 62.0, + "exceeded_stop_budget": false, + "exceeded_warn_budget": false +} +``` + +### Response (404 Not Found) +```json +"No production LLM connection found" +``` + +--- + +## 14. Reset Used Budget for All Connections + +### Endpoint +```http +POST /ruuter-public/rag-search/llm-connections/cost/reset +``` + +Resets `used_budget` to `0` for all LLM connections. Typically called by a scheduled job at the start of each billing period. + +### Request Body +None required. + +### Response (200 OK) +```json +{ + "message": "Used budget reset to 0 successfully for all connections", + "totalConnections": "5", + "operationSuccess": true, + "statusCode": 200 +} +``` + +### Response (500 Internal Server Error) +```json +"error: failed to reset used budget" +``` + +--- +# Inference Results API Endpoints + +## Base URL +``` +/ruuter-private/inference/results +``` + +--- + +## 1. Store Test Inference Result + +### Endpoint +```http +POST /ruuter-private/inference/results/test/store +``` + +### Request Body +```json +{ + "llm_connection_id": 1, + "user_question": "What are the benefits of using LLMs?", + "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `llm_connection_id` | number | Yes | ID of the LLM connection | +| `user_question` | string | Yes | User's raw question/input | +| `final_answer` | string | Yes | LLM's final generated answer | + +### Response (200 OK) +```json +{ + "data": { + "id": 10, + "llm_connection_id": 1, + "chat_id": null, + "user_question": "What are the benefits of using LLMs?", + "refined_questions": null, + "conversation_history": null, + "ranked_chunks": null, + "embedding_scores": null, + "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", + "environment": "testing", + "created_at": "2025-09-25T12:15:00.000Z" + }, + "operationSuccess": true, + "statusCode": 200 +} +``` + +### Response (400 Bad Request) +```json +{ + "data": "[]", + "operationSuccess": false, + "statusCode": 400 +} +``` + +### Response (404 Not Found) +```json +"error: LLM connection not found" +``` + +--- + +## 2. Store Production Inference Result + +### Endpoint +```http +POST /ruuter-private/inference/results/production/store +``` + +### Request Body +```json +{ + "chat_id": "chat-12345", + "user_question": "What are the benefits of using LLMs?", + "refined_questions": [ + "How do LLMs improve productivity?", + "What are practical use cases of LLMs?" + ], + "conversation_history": [ + { "role": "user", "content": "Hello" }, + { "role": "assistant", "content": "Hi! How can I help you?" } + ], + "ranked_chunks": [ + { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, + { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } + ], + "embedding_scores": [0.92, 0.85, 0.78], + "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `chat_id` | string | No | Optional chat session ID | +| `user_question` | string | Yes | User's raw question/input | +| `refined_questions` | object | No | List of refined questions (LLM-generated) | +| `conversation_history` | object | No | Prior messages array of {role, content} | +| `ranked_chunks` | object | No | Retrieved chunks ranked with metadata | +| `embedding_scores` | object | No | Distance scores for each chunk | +| `final_answer` | string | Yes | LLM's final generated answer | + +### Response (200 OK) +```json +{ + "data": { + "id": 15, + "llm_connection_id": null, + "chat_id": "chat-12345", + "user_question": "What are the benefits of using LLMs?", + "refined_questions": [ + "How do LLMs improve productivity?", + "What are practical use cases of LLMs?" + ], + "conversation_history": [ + { "role": "user", "content": "Hello" }, + { "role": "assistant", "content": "Hi! How can I help you?" } + ], + "ranked_chunks": [ + { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, + { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } + ], + "embedding_scores": [0.92, 0.85, 0.78], + "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", + "environment": "production", + "created_at": "2025-09-25T12:15:00.000Z" + }, + "operationSuccess": true, + "statusCode": 200 +} +``` + +### Response (400 Bad Request) +```json +{ + "data": "[]", + "operationSuccess": false, + "statusCode": 400 +} +``` + +--- + +## 3. View/get Inference Result + +### Endpoint +```http +POST /ruuter-private/inference/results/test/store +``` + +### Request Body +```json +{ + "llmConnectionId": 1, + "userQuestion": "What are the benefits of using LLMs?", + "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." +} +``` + +### Response (201 Created) +```json +{ + "data": { + "id": 15, + "llmConnectionId": 1, + "userQuestion": "What are the benefits of using LLMs?", + "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", + "environment": "testing", + "createdAt": "2025-09-25T10:15:30.000Z" + }, + "operationSuccess": true, + "statusCode": 200 +} +``` + +## 4. Inquiry from chatbot to llm orchestration service + +### Endpoint +```http +POST /ruuter-private/inference/results/production/store +``` + +### Request Body +```json +{ + "llmConnectionId": 1, + "chatId": "chat-session-12345", + "userQuestion": "What are the benefits of using LLMs?", + "refinedQuestions": [ + "How do LLMs improve productivity?", + "What are practical use cases of LLMs?" + ], + "conversationHistory": [ + { "role": "user", "content": "Hello" }, + { "role": "assistant", "content": "Hi! How can I help you?" } + ], + "rankedChunks": [ + { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, + { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } + ], + "embeddingScores": { + "chunk_1": 0.92, + "chunk_2": 0.85 + }, + "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." +} +``` + +### Response (201 Created) +```json +{ + "id": 20, + "llmConnectionId": 1, + "chatId": "chat-session-12345", + "userQuestion": "What are the benefits of using LLMs?", + "refinedQuestions": [ + "How do LLMs improve productivity?", + "What are practical use cases of LLMs?" + ], + "conversationHistory": [ + { "role": "user", "content": "Hello" }, + { "role": "assistant", "content": "Hi! How can I help you?" } + ], + "rankedChunks": [ + { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, + { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } + ], + "embeddingScores": { + "chunk_1": 0.92, + "chunk_2": 0.85 + }, + "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", + "environment": "production", + "createdAt": "2025-09-25T10:15:30.000Z" +} +``` + +--- + +## 5. Production Inference + +### Endpoint +```http +POST /ruuter-private/rag-search/inference/production +``` + +Validates the production connection's budget then proxies the request to the LLM Orchestration Service. + +### Request Body +```json +{ + "chatId": "chat-session-123", + "message": "What are the benefits of using LLMs?", + "authorId": "user-456", + "conversationHistory": [ + { "role": "user", "content": "Hello" }, + { "role": "assistant", "content": "Hi! How can I help you?" } + ], + "url": "https://example.com/context" +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `chatId` | string | Yes | Chat session ID | +| `message` | string | Yes | User message | +| `authorId` | string | Yes | Author ID | +| `conversationHistory` | array | No | Prior `{role, content}` messages | +| `url` | string | No | URL reference | + +### Response (200 OK) +Proxied response from the LLM Orchestration Service. + +### Response (400 Bad Request) — connection disconnected due to budget +```json +{ + "chatId": "chat-session-123", + "content": "The LLM connection is currently unavailable. Your request couldn't be processed. Please retry shortly.", + "status": 400 +} +``` + +### Response (404 Not Found) +```json +"No production connection found" +``` + +--- + +## 6. Test Inference + +### Endpoint +```http +POST /ruuter-private/rag-search/inference/test +``` + +Validates a specific connection's budget then calls the LLM Orchestration Service `/test` endpoint. + +### Request Body +```json +{ + "connectionId": "1", + "message": "What are the benefits of using LLMs?" +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `connectionId` | string | Yes | Connection ID to test against | +| `message` | string | Yes | User message | + +### Response (200 OK) +Proxied response from the LLM Orchestration Service `/test` endpoint. + +### Response (400 Bad Request) — connection disconnected due to budget +```json +{ + "connectionId": "1", + "content": "The LLM connection is currently unavailable. Your request couldn't be processed. Please retry shortly.", + "status": 400 +} +``` + +### Response (404 Not Found) +```json +"No test connection found" +``` + +--- + +## 7. View Inference Result (Mock) + +### Endpoint +```http +POST /ruuter-private/rag-search/inference/results/view +``` + +Returns a mock inference response for testing purposes. + +### Request Body +```json +{ + "llmConnectionId": 1, + "message": "What services are available?" +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `llmConnectionId` | number | Yes | LLM connection ID | +| `message` | string | Yes | User message/question | + +### Response (200 OK) +```json +{ + "chatId": 10, + "llmServiceActive": true, + "questionOutOfLlmScope": true, + "content": "Random answer with citations\n - https://gov.ee/sample1,\n - https://gov.ee/sample1" +} +``` + +### Response (400 Bad Request) +```json +"llmConnectionId and message are required" +``` + +--- + +## 8. Store Inference Result (Public) + +### Endpoint +```http +POST /ruuter-public/rag-search/inference/results/store +``` + +Public variant of the inference result store. Accepts the same fields as the private store endpoints, plus `environment` and `vault_uuid`. + +### Request Body +```json +{ + "user_question": "What are the benefits of using LLMs?", + "final_answer": "LLMs can improve productivity...", + "chat_id": "chat-12345", + "environment": "production", + "vault_uuid": "550e8400-e29b-41d4-a716-446655440000", + "refined_questions": ["How do LLMs improve productivity?"], + "conversation_history": [{ "role": "user", "content": "Hello" }], + "ranked_chunks": [{ "id": "chunk_1", "content": "...", "rank": 1 }], + "embedding_scores": [0.92, 0.85] +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `user_question` | string | Yes | User's raw question/input | +| `final_answer` | string | Yes | LLM's final generated answer | +| `chat_id` | string | No | Chat session ID | +| `environment` | string | No | Environment identifier | +| `vault_uuid` | string | No | Vault UUID for the LLM connection | +| `refined_questions` | object | No | List of refined questions | +| `conversation_history` | object | No | Prior `{role, content}` messages | +| `ranked_chunks` | object | No | Retrieved chunks ranked with metadata | +| `embedding_scores` | object | No | Distance scores for each chunk | + +### Response (200 OK) +```json +{ + "data": { "id": 20, "user_question": "...", "final_answer": "...", "environment": "production" }, + "operationSuccess": true, + "statusCode": 200 +} +``` + +### Response (400 Bad Request) +```json +{ + "data": "[]", + "operationSuccess": false, + "statusCode": 400 +} +``` + +--- + +# LLM Platforms & Models API Endpoints + +## Base URL +``` +/ruuter-private/rag-search +``` + +--- + +## 1. Get LLM Platforms + +### Endpoint +```http +GET /ruuter-private/rag-search/llm/platforms +``` + +Returns all active LLM platforms. + +### Response (200 OK) +```json +[ + { "id": 1, "value": "openai", "label": "OpenAI" }, + { "id": 2, "value": "azure", "label": "Azure AI" }, + { "id": 3, "value": "aws", "label": "AWS Bedrock" } +] +``` + +--- + +## 2. Get LLM Models by Platform + +### Endpoint +```http +GET /ruuter-private/rag-search/llm/models +``` + +### Query Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `platform_key` | string | Yes | Platform key to filter models (e.g. `"openai"`) | + +### Example Request +```http +GET /ruuter-private/rag-search/llm/models?platform_key=openai +``` + +### Response (200 OK) +```json +[ + { "id": 1, "value": "gpt-4o", "label": "GPT-4o", "platform_id": 1, "platform_key": "openai", "platform_name": "OpenAI" }, + { "id": 2, "value": "gpt-4o-mini", "label": "GPT-4o-mini", "platform_id": 1, "platform_key": "openai", "platform_name": "OpenAI" } +] +``` + +--- + +## 3. Get All LLM Models + +### Endpoint +```http +GET /ruuter-private/rag-search/llm/models-list +``` + +Returns all LLM models with no platform filter. + +### Response (200 OK) +```json +[ + { "id": 1, "platform_id": 1, "value": "gpt-4o", "label": "GPT-4o" }, + { "id": 2, "platform_id": 1, "value": "gpt-4o-mini", "label": "GPT-4o-mini" } +] +``` + +--- + +## 4. Get Embedding Platforms + +### Endpoint +```http +GET /ruuter-private/rag-search/embedding/platforms +``` + +Returns all active embedding platforms. + +### Response (200 OK) +```json +[ + { "id": 1, "value": "openai", "label": "OpenAI" }, + { "id": 2, "value": "azure", "label": "Azure AI" } +] +``` + +--- + +## 5. Get Embedding Models by Platform + +### Endpoint +```http +GET /ruuter-private/rag-search/embedding/models +``` + +### Query Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `embedding_platform_key` | string | Yes | Platform key to filter models | + +### Example Request +```http +GET /ruuter-private/rag-search/embedding/models?embedding_platform_key=openai +``` + +### Response (200 OK) +```json +[ + { "id": 1, "value": "text-embedding-3-small", "label": "text-embedding-3-small", "platform_id": 1, "platform_key": "openai", "platform_name": "OpenAI" }, + { "id": 2, "value": "text-embedding-ada-002", "label": "text-embedding-ada-002", "platform_id": 1, "platform_key": "openai", "platform_name": "OpenAI" } +] +``` + +--- + +# Prompt Configuration API Endpoints + +## Base URL +``` +/ruuter-private/rag-search/prompt-configuration +``` + +--- + +## 1. Get Prompt Configuration + +### Endpoint +```http +GET /ruuter-private/rag-search/prompt-configuration/get +``` + +Returns the active custom prompt configuration. Returns an empty array if none is configured. + +### Response (200 OK) +```json +[ + { + "id": 1, + "prompt": "You are a helpful assistant for government services...", + "created_at": "2025-09-02T10:15:30.000Z", + "updated_at": "2025-09-02T12:30:00.000Z" + } +] +``` + +--- + +## 2. Save Prompt Configuration + +### Endpoint +```http +POST /ruuter-private/rag-search/prompt-configuration/save +``` + +Upserts the prompt configuration (inserts if none exists, updates otherwise). Also triggers an LLM cache refresh. + +### Request Body +```json +{ + "prompt": "You are a helpful assistant for government services. Answer questions accurately and concisely." +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `prompt` | string | Yes | Prompt text to save | + +### Response (200 OK) +Returns the saved prompt configuration object. + +```json +{ + "id": 1, + "prompt": "You are a helpful assistant for government services. Answer questions accurately and concisely.", + "updated_at": "2025-09-25T12:00:00.000Z" +} +``` + +--- + +# Vault Secrets API Endpoints + +## Base URL +``` +/ruuter-private/rag-search/vault/secret +``` + +--- + +## 1. Create Vault Secret + +### Endpoint +```http +POST /ruuter-private/rag-search/vault/secret/create +``` + +Stores LLM connection credentials in Vault via CronManager. Supported platforms: `"aws"`, `"azure"`. + +### Request Body (AWS) +```json +{ + "vaultUuid": "550e8400-e29b-41d4-a716-446655440000", + "llmPlatform": "aws", + "llmModel": ["claude-3-sonnet"], + "secretKey": "aws-secret-key", + "accessKey": "aws-access-key", + "embeddingModel": "amazon.titan-embed-text-v1", + "embeddingPlatform": "aws", + "embeddingAccessKey": "embed-access-key", + "embeddingSecretKey": "embed-secret-key", + "deploymentEnvironment": "Production" +} +``` + +### Request Body (Azure) +```json +{ + "vaultUuid": "550e8400-e29b-41d4-a716-446655440000", + "llmPlatform": "azure", + "llmModel": ["gpt-4o"], + "deploymentName": "my-deployment", + "targetUrl": "https://my-endpoint.azure.com", + "apiKey": "azure-api-key", + "embeddingModel": "text-embedding-ada-002", + "embeddingPlatform": "azure", + "embeddingDeploymentName": "embed-deployment", + "embeddingTargetUri": "https://embed-endpoint.azure.com", + "embeddingAzureApiKey": "embed-azure-api-key", + "deploymentEnvironment": "Production" +} +``` + +### Request Parameters +| Parameter | Type | Platform | Description | +|-----------|------|----------|-------------| +| `vaultUuid` | string | Both | Stable UUID for the vault path | +| `llmPlatform` | string | Both | `"aws"` or `"azure"` | +| `llmModel` | array | Both | LLM model identifier(s) | +| `deploymentEnvironment` | string | Both | Deployment environment | +| `embeddingModel` | string | Both | Embedding model identifier | +| `embeddingPlatform` | string | Both | Embedding platform | +| `secretKey` | string | AWS | AWS secret key | +| `accessKey` | string | AWS | AWS access key | +| `embeddingAccessKey` | string | AWS | Embedding AWS access key | +| `embeddingSecretKey` | string | AWS | Embedding AWS secret key | +| `deploymentName` | string | Azure | Azure deployment name | +| `targetUrl` | string | Azure | Azure endpoint URL | +| `apiKey` | string | Azure | Azure API key | +| `embeddingDeploymentName` | string | Azure | Embedding Azure deployment name | +| `embeddingTargetUri` | string | Azure | Embedding Azure endpoint URI | +| `embeddingAzureApiKey` | string | Azure | Embedding Azure API key | + +### Response (200 OK) — AWS +```json +"Executed cron manager successfully to store aws secrets" +``` + +### Response (200 OK) — Azure +```json +"Executed cron manager successfully to store azure secrets" +``` + +### Response (400 Bad Request) +```json +{ + "message": "Platform not supported", + "operationSuccessful": false, + "statusCode": 400 +} +``` + +--- + +## 2. Delete Vault Secret + +### Endpoint +```http +POST /ruuter-private/rag-search/vault/secret/delete +``` + +Removes LLM connection credentials from Vault via CronManager. + +### Request Body +```json +{ + "vaultUuid": "550e8400-e29b-41d4-a716-446655440000", + "llmPlatform": "azure", + "llmModel": "gpt-4o", + "embeddingModel": "text-embedding-ada-002", + "embeddingPlatform": "azure", + "deploymentEnvironment": "Production" +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `vaultUuid` | string | Yes | Vault UUID of the connection | +| `llmPlatform` | string | Yes | LLM platform | +| `llmModel` | string | Yes | LLM model identifier | +| `embeddingModel` | string | Yes | Embedding model identifier | +| `embeddingPlatform` | string | Yes | Embedding platform | +| `deploymentEnvironment` | string | Yes | Deployment environment | + +### Response (200 OK) +```json +"Executed cron manager successfully to delete secrets from vault" +``` + +### Response (404 Not Found) +```json +{ + "message": "Connection not found with the provided vaultUuid", + "operationSuccessful": false, + "statusCode": 404 +} +``` + +--- + +# Data Sync & Services API Endpoints + +## Base URL +``` +/ruuter-public/rag-search +``` + +--- + +## 1. Get Services (for Intent Detection) + +### Endpoint +```http +GET /ruuter-public/rag-search/services/get-services +``` + +Returns all active services if the count is ≤ 10. If count > 10, signals the caller to use semantic search instead. + +### Response (200 OK) — ≤ 10 services +```json +{ + "use_semantic_search": false, + "service_count": 5, + "services": [ + { "id": "svc-1", "name": "Pension Application", "description": "..." } + ] +} +``` + +### Response (200 OK) — > 10 services +```json +{ + "use_semantic_search": true, + "service_count": 23, + "message": "Service count exceeds threshold - use semantic search" +} +``` + +--- + +## 2. Resync Data from KB + +### Endpoint +```http +POST /ruuter-public/rag-search/data/update +``` + +Fetches the latest agency data from CKB, compares the data hash, and if changed triggers vector re-indexing via CronManager. + +### Request Body +None required. + +### Response (200 OK) — sync initiated +```json +{ + "message": "Data synchronization initiated successfully", + "operationSuccessful": true +} +``` + +### Response (200 OK) — already up to date +```json +{ + "success": true, + "message": "No sync required - data is up to date" +} +``` + +### Response (400 Bad Request) +```json +{ + "message": "CKB service returned an error - data synchronization aborted", + "operationSuccessful": false, + "error": "CKB_ERROR" +} +``` + +### Response (404 Not Found) +```json +{ + "success": false, + "message": "Data synchronization failed - CKB agency data not found" +} +``` + +--- + +## 3. Trigger API Tool Endpoint Indexing + +### Endpoint +```http +POST /ruuter-public/rag-search/api-tools/index +``` + +Queues an API tool endpoint for vector indexing in Qdrant via CronManager (async). + +### Request Body +```json +{ + "endpointId": "ep-001", + "serviceId": "svc-1", + "name": "Get Pension Status", + "description": "Retrieve the current pension application status for a citizen", + "method": "GET", + "url": "https://api.example.com/pension/status", + "visibility": "public", + "params": [ + { "name": "nationalId", "type": "string", "required": true } + ] +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `endpointId` | string | Yes | Unique endpoint identifier | +| `name` | string | Yes | Endpoint name | +| `description` | string | Yes | Endpoint description | +| `url` | string | Yes | API URL | +| `serviceId` | string | No | Parent service ID | +| `method` | string | No | HTTP method (default: `"GET"`) | +| `visibility` | string | No | `"public"` or `"private"` (default: `"public"`) | +| `type` | string | No | Endpoint type (default: `"custom_endpoint"`) | +| `params` | array | No | List of parameters | + +### Response (200 OK) +```json +{ + "success": true, + "endpoint_id": "ep-001", + "message": "API Tool indexing job queued successfully. Processing asynchronously." +} +``` + +### Response (400 Bad Request) +```json +{ + "success": false, + "error": "MISSING_REQUIRED_FIELDS", + "message": "endpointId, name, description, and url are required" +} +``` + +### Response (500 Internal Server Error) +```json +{ + "success": false, + "error": "INDEXING_QUEUE_FAILED", + "message": "Failed to queue indexing job. CronManager may be unavailable." +} +``` + +--- + +## 4. Enrich and Index Service + +### Endpoint +```http +POST /ruuter-public/rag-search/services/enrich +``` + +Queues a service for enrichment and Qdrant indexing via CronManager (async). + +### Request Body +```json +{ + "service_id": "svc-001", + "name": "Pension Application", + "description": "Submit a new pension application for eligible citizens", + "examples": ["How do I apply for pension?", "Pension eligibility requirements"], + "entities": ["nationalId", "dateOfBirth"], + "ruuter_type": "POST", + "current_state": "active", + "is_common": false +} +``` + +### Request Parameters +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `service_id` | string | Yes | Unique service identifier | +| `name` | string | Yes | Service name | +| `description` | string | Yes | Service description | +| `examples` | array | No | Example user queries | +| `entities` | array | No | Expected entity names | +| `ruuter_type` | string | No | HTTP method (default: `"GET"`) | +| `current_state` | string | No | `"active"`, `"inactive"`, or `"draft"` (default: `"draft"`) | +| `is_common` | boolean | No | Whether this is a common service (default: `false`) | + +### Response (200 OK) +```json +{ + "success": true, + "service_id": "svc-001", + "message": "Service enrichment job queued successfully. Processing asynchronously." +} +``` + +### Response (400 Bad Request) +```json +{ + "success": false, + "error": "MISSING_REQUIRED_FIELDS", + "message": "service_id, name, and description are required" +} +``` + +### Response (500 Internal Server Error) +```json +{ + "success": false, + "error": "ENRICHMENT_QUEUE_FAILED", + "message": "Failed to queue enrichment job. CronManager may be unavailable." +} +``` + +--- \ No newline at end of file diff --git a/docs/API_TOOL_CALLING.md b/docs/API_TOOL_CALLING.md index bdc5c655..e7467cfd 100644 --- a/docs/API_TOOL_CALLING.md +++ b/docs/API_TOOL_CALLING.md @@ -12,11 +12,17 @@ loop collects all required parameters from the user before the API call is made. | Component | What it does | Status | |---|---|---| -| **Indexing pipeline** | Takes an endpoint definition → enriches it with LLM context → stores hybrid vectors in Qdrant | ✅ Complete | -| **Tool classifier** | At query time, routes to the best matching endpoint via hybrid search + LLM disambiguation | ✅ Complete | -| **Agentic loop** | Multi-turn parameter collection with session persistence, language-aware clarifying questions, param correction, continuation prompt, and intent-switch detection | ✅ Complete | -| **API caller** | Execute collected params against the real API endpoint, with circuit-breaker protection and localized error handling | ✅ Complete | -| **Response formatter** | Convert raw API JSON into a natural-language answer via DSPy, streamed token-by-token to the GUI | ✅ Complete | +| **Indexing pipeline** | Takes an endpoint definition → enriches it with LLM context → stores hybrid vectors in Qdrant | Complete | +| **Tool classifier** | At query time, routes to the best matching endpoint via hybrid search + LLM disambiguation | Complete | +| **Multi-intent detection** | Score-band gate triggers `IntentDecomposer` (DSPy) to decompose a multi-intent query into focused sub-queries; each sub-query is matched in parallel via `asyncio.gather` | Phase 1 & 2 Complete | +| **Agentic loop** | Multi-turn parameter collection with session persistence, language-aware clarifying questions, param correction, continuation prompt, and intent-switch detection | Complete | +| **API caller** | Execute collected params against the real API endpoint, with circuit-breaker protection and localized error handling | Complete | +| **Response formatter** | Convert raw API JSON into a natural-language answer via DSPy, streamed token-by-token to the GUI | Complete | +| **Multi-endpoint loop** | Merges param schemas for all parallel endpoints; collects params across turns with a single deduplicated clarifying question per turn; distributes values back per endpoint | Phase 3 Complete | +| **Parallel API caller** | Fires all completed endpoint calls concurrently via `asyncio.gather` with batch timeout and partial-failure handling | Phase 4 Complete | +| **Multi-response formatter** | DSPy module that synthesises N API results into a single coherent natural-language answer; supports streaming and blocking execution | Phase 5 Complete | +| **Full wiring** | `APIToolWorkflowExecutor` routes parallel sessions through `MultiEndpointAgenticLoop` → `MultiAPICaller` → `MultiResponseFormatterModule` with output guardrails | Phase 6 Complete | +| **ATC Response Cache** | Two-tier Redis cache (L1 exact-match + L2 follow-up context) that eliminates redundant API calls and enables intelligent follow-up handling without re-running the agentic loop | Complete | --- @@ -36,17 +42,31 @@ api_tool_collection (Qdrant) APISemanticSearcher (src/tool_classifier/api_semantic_searcher.py) ↑ called by ToolClassifier._try_api_tool_classification() + │ + ├─ score ≥ HIGH_CONFIDENCE → single path (unchanged) + │ + └─ score in ambiguous band → IntentDecomposer (DSPy) + │ + ├─ mode=single → top candidate, existing path + └─ mode=parallel → asyncio.gather(search per sub-query) + ↓ ClassificationResult(execution_mode=parallel, matched_endpoints=[...]) ↓ ClassificationResult(workflow=API_TOOL_CALLING) APIToolWorkflowExecutor (src/tool_classifier/workflows/api_tool_workflow.py) - ↓ multi-turn param collection -AgenticLoop (src/tool_classifier/agentic_loop.py) - ↓ session state + │ + ├─ execution_mode=single → AgenticLoop + └─ execution_mode=parallel → MultiEndpointAgenticLoop + ↓ merged schema; one clarifying question per turn + ↓ distributes extracted values back per endpoint APIToolSessionStore (Redis, keyed by chat_id, 30-min TTL) - ↓ all params collected -APICaller (src/tool_classifier/api_caller.py) - ↓ raw JSON response -APIResponseFormatterModule (src/tool_classifier/api_response_formatter.py) - ↓ SSE token stream + ↓ all endpoints completed + ├─ single → APICaller + └─ parallel → MultiAPICaller → asyncio.gather per endpoint + ↓ batch timeout: MULTI_API_BATCH_TIMEOUT (30 s) + ↓ raw JSON responses (partial-failure safe) + ├─ single → APIResponseFormatterModule + └─ parallel → MultiResponseFormatterModule (DSPy) + ↓ buffer-first guardrails validation + ↓ SSE token stream User (GUI) ``` @@ -189,21 +209,21 @@ A `PointStruct` is built: ```python PointStruct( - id = endpoint_id, # UUID used directly as Qdrant point ID - vector = { - "dense": [v1, v2, ..., v3072], + id=endpoint_id, # UUID used directly as Qdrant point ID + vector={ + "dense": [v1, v2, ..., v3072], "sparse": {"indices": [...], "values": [...]}, }, - payload = { # Stored metadata — no extra DB lookup needed - "endpoint_id": "...", - "name": "get_national_holidays", - "description": "...", - "url": "https://openholidaysapi.org/PublicHolidays", - "method": "GET", - "params": [...], + payload={ # Stored metadata — no extra DB lookup needed + "endpoint_id": "...", + "name": "get_national_holidays", + "description": "...", + "url": "https://openholidaysapi.org/PublicHolidays", + "method": "GET", + "params": [...], "enriched_context": "...", - "service_id": "...", - } + "service_id": "...", + }, ) ``` @@ -286,9 +306,9 @@ Instantiated once in `ToolClassifier.__init__()` and reuses the shared Qdrant `h ```python APISemanticSearcher( - embedding_service=orchestration_service, # generates dense embeddings - qdrant_client=self._qdrant_client, # shared connection pool - disambiguator=None, # optional: inject for testing + embedding_service=orchestration_service, # generates dense embeddings + qdrant_client=self._qdrant_client, # shared connection pool + disambiguator=None, # optional: inject for testing ) ``` @@ -767,4 +787,690 @@ uv run --no-project --with requests python tests/api_tool_eval/integration_test_ | 6 | Electricity prices | 2 | `datetime` params across two turns | | 7 | Session isolation | 2 | Two different chat IDs — no param leak between sessions | | 8 | AWAITING_CONTINUATION → yes | 4+ | User says “yes” at continuation prompt → loop resumes → API call on completion | -| 9 | MAX_TURNS_REACHED | 5+ | User never provides params → falls back to RAG | \ No newline at end of file +| 9 | MAX_TURNS_REACHED | 5+ | User never provides params → falls back to RAG | + +--- + +## Part 7 — Multi-Intent Handling + +> **Status:** All phases complete. Phase 1 (intent detection + parallel search), Phase 2 (session model extension), Phase 3 (multi-endpoint agentic loop), Phase 4 (parallel API caller), Phase 5 (multi-response formatter), and Phase 6 (full wiring in workflow executor) are production-ready. + +### Overview + +A multi-intent query like *"What are the public holidays in Estonia and what is the current electricity price?"* produces a **diluted embedding** — the dense vector sits between two endpoints rather than close to either one. The cosine score lands in the ambiguous band (≥ `API_TOOL_MIN_THRESHOLD`, < `API_TOOL_HIGH_CONFIDENCE_THRESHOLD`) rather than producing a clean high-confidence hit. + +Phase 1 adds a **score-band gate** that intercepts these ambiguous results and passes the raw query to `IntentDecomposer`. If two or more distinct intents are detected, sub-queries are searched in parallel and the results are stored in the session for eventual multi-endpoint execution. + +--- + +### Phase 1: Score-Band Gate + Intent Decomposer + +#### Gate Logic (`classifier.py` — `_try_api_tool_classification`) + +``` +cosine ≥ HIGH_CONFIDENCE_THRESHOLD → single path, existing code unchanged +cosine in [MIN_THRESHOLD, HIGH_CONFIDENCE) → ambiguous band → IntentDecomposer +cosine < MIN_THRESHOLD → no match → RAG fallback +``` + +The gate only fires when: +- `FeatureFlags.MULTI_INTENT_ENABLED = true` (env: `MULTI_INTENT_ENABLED`, default `true`) +- The top result was **not already LLM-validated** by the disambiguator (`llm_validated=False`) + +The `llm_validated` flag on `APIToolSearchResult` prevents double LLM calls: when the disambiguator has already selected a winner it sets `llm_validated=True`, so the IntentDecomposer gate is skipped for that result. + +#### Disambiguator Edge Case Fix + +When `APISemanticSearcher` runs LLM disambiguation and the disambiguator **rejects all candidates** (returns `winner_id=None`): + +- **Old behaviour:** return `[]` → classified as RAG/CONTEXT even when the query was multi-intent +- **New behaviour:** if there were multiple medium-confidence candidates, return the top cosine result *without* `llm_validated=True` so the IntentDecomposer gate can run + +This is the key fix that allows multi-intent queries to reach `IntentDecomposer` instead of falling through to RAG. + +#### `IntentDecomposerModule` (`src/tool_classifier/intent_decomposer.py`) + +DSPy module — receives the **raw user query** + +| Output | Description | +|---|---| +| `mode` | `"single"` or `"parallel"` | +| `sub_queries` | List of focused sub-queries when `mode=parallel`; empty for single | + +**Conservative by design:** returns `"single"` on any failure or ambiguity — never forces a parallel path. + +Run asynchronously via `asyncio.to_thread(self, user_query)` (calls `__call__`, not `forward`, to avoid a DSPy warning). + +Sub-query count is capped at `MULTI_API_MAX_ENDPOINTS = 3`. + +#### Parallel Sub-Query Search + +When `mode=parallel`, each sub-query is independently searched against `api_tool_collection` using `asyncio.gather`: + +```python +results = await asyncio.gather( + *[self._api_searcher.search(q, ...) for q in sub_queries] +) +``` + +Each search generates its own focused embedding — no dilution from the combined query. + +Results are **deduplicated by endpoint name** (a single endpoint matched by two sub-queries counts once). If fewer than 2 distinct endpoints are found after dedup, the parallel result is discarded and the classifier falls back to the single path. + +#### `ExecutionMode` Enum (`src/tool_classifier/enums.py`) + +```python +class ExecutionMode(str, Enum): + SINGLE = "single" + PARALLEL = "parallel" +``` + +Inherits from `str` so values serialize cleanly to JSON. Always compare against the enum member (e.g. `== ExecutionMode.PARALLEL`), not a string literal. + +#### `ClassificationResult` metadata for parallel mode + +| Key | Type | Description | +|---|---|---| +| `execution_mode` | `ExecutionMode` | `SINGLE` or `PARALLEL` | +| `matched_endpoint` | dict | Set for single mode | +| `matched_endpoints` | list[dict] | Set for parallel mode — all deduplicated endpoints | + +--- + +### Phase 2: Session Model Extension + +Defined in [src/models/session_models.py](../src/models/session_models.py). + +#### `EndpointSessionState` (new model) + +Tracks per-endpoint collection state within a parallel session: + +| Field | Type | Default | Description | +|---|---|---|---| +| `endpoint` | dict | required | Full endpoint payload | +| `collected_params` | dict | `{}` | Params gathered for this endpoint so far | +| `completed` | bool | `False` | Flipped to `True` when all required params for this endpoint have been collected; API execution happens later, gated on `AgenticLoopStatus.COMPLETED` at the loop level | + +#### Extended `APIToolSession` fields + +Three new fields added — all optional with safe defaults so **existing Redis sessions deserialize without error**: + +| Field | Type | Default | Description | +|---|---|---|---| +| `execution_mode` | str | `"single"` | `"single"` or `"parallel"` — drives which loop handles this session | +| `parallel_endpoints` | list[EndpointSessionState] | `[]` | One entry per matched endpoint; empty in single mode | +| `active_endpoint_index` | int | `0` | Reserved for future sequential endpoint-tracking logic; not currently read or written by the parallel workflow | + +The existing single-mode fields (`selected_endpoint`, `collected_params`, `turn_count`, etc.) are **completely unchanged**. Single-mode sessions have `execution_mode="single"` and `parallel_endpoints=[]`. + +#### Session creation in parallel mode (`api_tool_workflow.py`) + +When `context["execution_mode"] == ExecutionMode.PARALLEL`, the workflow captures all matched endpoints from context and populates `parallel_endpoints`: + +```python +APIToolSession( + ... + execution_mode="parallel", + parallel_endpoints=[ + EndpointSessionState(endpoint=e) for e in all_matched + ], +) +``` + + + +--- + +### Phase 3: Multi-Endpoint Agentic Loop + +Defined in [src/tool_classifier/multi_agentic_loop.py](../src/tool_classifier/multi_agentic_loop.py). + +`MultiEndpointAgenticLoop` replaces `AgenticLoop` for parallel sessions. It is stateless between HTTP requests — all mutable state is held in the `EndpointSessionState` list passed in on each call. + +**Key behaviours:** + +- **Merged schema:** All endpoint `params` schemas are merged and deduplicated by param name. A param shared across two endpoints is asked once and applied to both. +- **One question per turn:** The loop generates a single clarifying question covering the next highest-priority missing param across all endpoints. +- **Per-endpoint distribution:** After extraction, `_distribute_params()` copies each extracted value to every endpoint whose schema includes that param name. +- **Completion tracking:** An endpoint is marked `completed=True` in its `EndpointSessionState` once all its required params are present. The loop returns `COMPLETED` only when every endpoint is completed. +- **Turn limit:** Determined by `_compute_turn_limits(num_endpoints)`: + - Multi-intent (`num_endpoints > 1`): fixed `MULTI_INTENT_MAX_TURNS` (6). + - Single-intent (`num_endpoints == 1`): `min(3 × num_endpoints, MULTI_API_MAX_TURNS)` — scales with endpoint count, capped at 9. +- **Continuation threshold:** Also from `_compute_turn_limits`: + - Multi-intent: fixed `MULTI_INTENT_CONTINUATION_TURN` (4). + - Single-intent: `num_endpoints + 1`. + +**`stream_run_turn` signature:** + +```python +await multi_loop.stream_run_turn( + chat_id=chat_id, + user_message=request.message, + conversation_history=conversation_history, + endpoint_states=session.parallel_endpoints, # list[EndpointSessionState] + turn_count=session.turn_count, + awaiting_continuation=session.awaiting_continuation, + session_language=effective_session_language, +) +``` + +Returns `(AgenticLoopResult, list[str])` — the result and pre-tokenised question tokens. + +--- + +### Phase 4: Parallel API Caller + +Defined in [src/tool_classifier/multi_api_caller.py](../src/tool_classifier/multi_api_caller.py). + +`MultiAPICaller` wraps `APICaller` and fires all endpoint calls concurrently via `asyncio.gather`. + +**Key design decisions:** + +- Reuses the shared `APICaller` instance so per-URL circuit breaker state is preserved across single and batch invocations. +- Expects a `"call_params"` key on each endpoint dict (distinct from the `"params"` schema list) to prevent the schema descriptor from being forwarded to the HTTP call. +- A `MULTI_API_BATCH_TIMEOUT` (30 s) caps total wall-clock time. Pending tasks are cancelled on timeout and replaced with failure results — the caller always receives a fully-populated `MultiAPICallResult`. +- Results are returned in the same order as the input endpoint list. +- Partial failure is safe — a failed endpoint produces `APICallResult(success=False, error=)` without affecting other endpoints. + +**Usage:** + +```python +call_payloads = [ + {**state.endpoint, "call_params": state.collected_params} + for state in parallel_endpoints +] +multi_result = await MultiAPICaller(api_caller).call_all( + call_payloads, language=detected_language +) +``` + +--- + +### Phase 5: Multi-Response Formatter + +Defined in [src/tool_classifier/multi_response_formatter.py](../src/tool_classifier/multi_response_formatter.py). + +`MultiResponseFormatterModule` is a DSPy module that synthesises N API results into one unified natural-language answer. + +**DSPy Signature:** `MultiResponseFormatterSignature` + +| Input field | Description | +|---|---| +| `user_query` | The user's original first-turn question | +| `api_results_block` | Formatted block of all API results (name, description, data) | +| `num_results` | Number of results being synthesised | +| `response_language` | `"English"`, `"Estonian"`, or `"Russian"` | +| `custom_instructions` | Optional operator prompt overrides | + +| Output field | Description | +|---|---| +| `unified_answer` | Single coherent natural-language answer covering all endpoints | + +**Rules enforced by signature:** +- Always write in `response_language` regardless of API data language. +- Address every result — do not silently omit any endpoint. +- Gracefully acknowledge failed or empty results without dwelling on them. +- No raw JSON, no markdown headers, no follow-up invitation sentences. + +**Max input size:** 100 KB across all results combined (`_MAX_TOTAL_RESPONSE_BYTES`). + +**Methods:** `forward(user_query, api_results, detected_language)` (blocking) and `stream_forward_multi(user_query, api_results, detected_language)` (async token iterator). + +--- + +### Phase 6: Full Wiring in `APIToolWorkflowExecutor` + +Defined in [src/tool_classifier/workflows/api_tool_workflow.py](../src/tool_classifier/workflows/api_tool_workflow.py). + +`_LoopStep` now has four possible `kind` values: + +| `kind` | Meaning | +|---|---| +| `"api_call"` | Single endpoint; all params collected; call API and format | +| `"multi_api_call"` | Parallel endpoints; call all APIs concurrently and merge results | +| `"question"` | Agentic loop needs more input; return question to user | +| `"fallback"` | Nothing to do; caller falls back to RAG | + +**Parallel fast-path:** When all matched endpoints have no required params, `_compute_loop_step` skips session creation and returns a `"multi_api_call"` step immediately. + +**Streaming path (`_stream_multi_api_and_format`):** +1. Builds `call_params`-keyed payloads for each `EndpointSessionState`. +2. Calls `MultiAPICaller.call_all()` — all HTTP calls fire concurrently. +3. Collects tokens from `MultiResponseFormatterModule.stream_forward_multi()` into a buffer. +4. Runs output guardrails on the full buffered response before yielding any token to the client. +5. Yields tokens one-by-one via `format_sse`, then yields `format_sse(chat_id, "END")`. + +**Blocking path (`_execute_multi_api_and_format`):** +Same API call steps, then `asyncio.to_thread(formatter.forward, ...)` for the synthesis step. + +--- + +### Constants and Feature Flags + +Defined in [src/tool_classifier/constants.py](../src/tool_classifier/constants.py) and [src/llm_orchestrator_config/feature_flags.py](../src/llm_orchestrator_config/feature_flags.py): + +| Name | Value | Description | +|---|---|---| +| `MULTI_INTENT_ENABLED` | `true` (env override) | Feature flag — set `MULTI_INTENT_ENABLED=false` to disable the IntentDecomposer gate globally | +| `MULTI_API_MAX_ENDPOINTS` | `3` | Hard cap on parallel sub-queries per request | +| `MULTI_API_MAX_TURNS` | `9` | Absolute cap on turns for any parallel session (`min(3×N, 9)` per session) | +| `MULTI_API_BATCH_TIMEOUT` | `30` | Seconds before the parallel HTTP batch is cancelled and partial results returned | +| `API_TOOL_HIGH_CONFIDENCE_THRESHOLD` | `0.60` | Cosine score above which single-path is taken immediately | +| `API_TOOL_MIN_THRESHOLD` | `0.40` | Minimum score for any match (below → RAG) | +| `API_TOOL_INTENT_SWITCH_THRESHOLD` | `0.50` | Minimum cosine for the new match to trigger intent-switch detection | + +--- + +### Multi-Intent End-to-End Flow + +``` +Turn 1 — User: "Can you find an address for me and also calculate my vehicle tax?" + │ + ▼ +APISemanticSearcher.search() + → top result: search_address, cosine=0.54 (ambiguous band) + → disambiguator rejects both candidates (multi-intent dilution) + → returns top candidate WITHOUT llm_validated=True + │ + ▼ +_try_api_tool_classification() — gate fires + → MULTI_INTENT_ENABLED=true AND not llm_validated + → IntentDecomposer (DSPy, asyncio.to_thread) + → mode=parallel + → sub_queries=["address lookup and location search", "vehicle tax calculation"] + │ + ▼ +asyncio.gather( + search("address lookup and location search") → search_address, cosine=0.82 + search("vehicle tax calculation") → get_vehicle_tax_info, cosine=0.79 +) + → 2 distinct endpoints after dedup → parallel path confirmed + │ + ▼ +ClassificationResult( + workflow=API_TOOL_CALLING, + metadata={ + execution_mode: ExecutionMode.PARALLEL, + matched_endpoints: [search_address, get_vehicle_tax_info] + } +) + │ + ▼ +APIToolWorkflowExecutor._compute_loop_step() + → no existing session → create new: + APIToolSession( + execution_mode="parallel", + selected_endpoint=search_address, # first endpoint + original_query="Can you find an address...", + parallel_endpoints=[ + EndpointSessionState(endpoint=search_address, collected_params={}), + EndpointSessionState(endpoint=get_vehicle_tax_info, collected_params={}), + ], + turn_count=0, + max_turns=6, # min(3×2, 9) + ) + → MultiEndpointAgenticLoop.stream_run_turn(endpoint_states=[...], turn_count=0) + → merged schema: {address, regNr, calculationYear} (deduped across both endpoints) + → nothing extracted from turn-1 message (intent query, not param values) + → missing: [address, regNr, calculationYear] + → NEEDS_INPUT → clarifying question + │ + Session saved to Redis, _LoopStep(kind="question") + ▼ +Bot: "To help you, I need a few details: the address you'd like to look up, + your vehicle registration number (regNr), and the calculation year for the vehicle tax." + +─────────────────────────────────────────────────────────────── + +Turn 2 — User: "123ABC" + │ + ▼ +ToolClassifier.classify() + → Active session found → intent-switch check + → "123ABC" cosine < API_TOOL_INTENT_SWITCH_THRESHOLD → no switch + → ClassificationResult(reason=active_session_resume) + │ + ▼ +MultiEndpointAgenticLoop.stream_run_turn(turn_count=1) + → ParamExtractionModule extracts regNr="123ABC" + → _distribute_params: regNr → get_vehicle_tax_info.collected_params + → still missing: [address, calculationYear] + → NEEDS_INPUT → clarifying question covers ALL remaining missing params + ▼ +Bot: "Got it! I still need two more things: the address you'd like to look up, + and the calculation year for the vehicle tax." + +─────────────────────────────────────────────────────────────── + +Turn 3 — User: "Viru tn 4, Tallinn and year 2026" + │ + ▼ +MultiEndpointAgenticLoop.stream_run_turn(turn_count=2) + → ParamExtractionModule extracts address="Viru tn 4, Tallinn", calculationYear="2026" + → _distribute_params: + address → search_address.collected_params → completed=True + calculationYear → get_vehicle_tax_info.collected_params + → get_vehicle_tax_info now has [regNr, calculationYear] → completed=True + → ALL endpoints completed → AgenticLoopStatus.COMPLETED + │ + Session DELETED from Redis + _LoopStep(kind="multi_api_call", parallel_endpoints=[...]) + ▼ +APIToolWorkflowExecutor._stream_multi_api_and_format() + │ + ├─ Build call_payloads: + │ [{...search_address, call_params: {address: "Viru tn 4, Tallinn"}}, + │ {...get_vehicle_tax_info, call_params: {regNr: "123ABC", calculationYear: "2026"}}] + │ + ├─ MultiAPICaller.call_all(payloads, language="en") + │ asyncio.gather( + │ GET /address-search?address=Viru+tn+4%2C+Tallinn → 200 OK, address JSON + │ GET /vehicle-tax?regNr=123ABC&year=2026 → 200 OK, tax JSON + │ ) # batch_timeout=30 s; 2/2 succeeded + │ + ├─ MultiResponseFormatterModule.stream_forward_multi( + │ user_query=session.original_query, # full first-turn message + │ api_results=[("search_address", ..., address_data), + │ ("get_vehicle_tax_info", ..., tax_data)], + │ detected_language="en" + │ ) + │ → DSPy streams unified answer tokens + │ → buffer-first: collect all tokens + │ + ├─ Output guardrails on full buffered response → passed + │ + └─ yield format_sse(chat_id, token) per token → yield format_sse(chat_id, "END") + ▼ +Bot: "Here's what I found: The address Viru tn 4 is located in Tallinn city centre + (full address: Viru tn 4, 10111 Tallinn). For vehicle 123ABC, the estimated + vehicle tax for 2026 is €127.40." ← streamed token-by-token +``` + +--- + + + + +## Part 8 — ATC Response Cache + +### Overview + +The ATC Response Cache is a two-tier Redis cache that sits inside `_compute_loop_step()` in +[src/tool_classifier/workflows/api_tool_workflow.py](../src/tool_classifier/workflows/api_tool_workflow.py). +It is checked on every **new request** (no active session) before the agentic loop is created. + +Goal: avoid redundant API calls and agentic loop turns when the user is repeating or +following up on a query that was already answered in the same conversation. + +Gated by `FeatureFlags.ATC_RESPONSE_CACHE_ENABLED` (`ATC_RESPONSE_CACHE_ENABLED` env var, default `true`). +Setting it to `false` disables all cache reads and writes without touching any other ATC logic. + +--- + +### Cache Architecture — Two Tiers + +#### Tier 1 — L1 Exact Response Cache + +``` +Key: atc:cache:{chat_id}:{api_name}:{param_hash} +Value: raw API response JSON (dict or list) +TTL: per-endpoint cache_ttl_seconds OR ATC_CACHE_DEFAULT_TTL_SECONDS (30 min) +``` + +Answers the question: *Has this exact conversation called this exact endpoint with these exact params before?* + +`param_hash` is a 16-character hex digest of the **normalised, sorted** param dict: +- String values are stripped of whitespace +- Purely numeric strings (`"2026"`) are cast to `int` before hashing +- All-alpha strings (enum-like, e.g. `"GET"`, `"EE"`) are lowercased +- Keys are sorted so order does not matter + +This means `{year: "2026", country: "EE"}` and `{country: "ee", year: 2026}` produce +the **same hash** and hit the same cache entry. + +#### Tier 2 — L2 Last Call Context + +``` +Key: atc:last:{chat_id} +Value: JSON list[LastCallContext] +TTL: ATC_LAST_CALL_TTL_SECONDS (30 min, sliding — reset on every write) +``` + +Answers the question: *What was the last API call made in this conversation?* + +Stores a full `LastCallContext` per succeeded endpoint. Single-intent calls write a +one-element list; multi-intent parallel calls write one entry per succeeded endpoint. +The follow-up detector searches this list by `api_name` to find the relevant prior call. + +--- + +### Data Model: `LastCallContext` + +Defined in [src/models/session_models.py](../src/models/session_models.py). + +| Field | Type | Description | +|---|---|---| +| `api_name` | str | Endpoint name (snake_case) that was called | +| `endpoint` | dict | Full endpoint payload from Qdrant (params schema, URL, method, etc.) | +| `collected_params` | dict | Parameter values that were passed to the API call | +| `raw_response` | Any | Parsed API JSON (dict or list) as returned by `APICaller` | +| `original_query` | str | User's first-turn query that triggered this API call | +| `timestamp` | float | Unix timestamp of the call (for staleness reference) | + +--- + + +### Cache Write Points + +L1 and L2 are written **after** every successful API call, as a background +`asyncio.create_task` so they never delay the user-facing response: + +**Single-intent (`_execute_api_and_format`)** +After `api_result.success == True` and before the formatter: +``` +set_l1(chat_id, endpoint["name"], collected_params, response_data, ttl) +set_l2(chat_id, [LastCallContext(...)]) +``` + +**Multi-intent (`_execute_multi_api_and_format` / `_stream_multi_api_and_format`)** +After `multi_result` is received, one write per succeeded endpoint: +``` +for each (endpoint_state, result) where result.success and endpoint.cacheable: + set_l1(chat_id, endpoint["name"], endpoint_state.collected_params, result.response_data, ttl) + append LastCallContext to contexts_list +set_l2(chat_id, contexts_list) ← one write for all endpoints +``` + +--- + +### Cache Read Logic in `_compute_loop_step` + +The cache block runs only when: +- No active Redis session exists (fresh request, not mid-loop) +- Not a parallel multi-intent query (`not all_matched`) +- `endpoint.cacheable == True` +- `ATC_RESPONSE_CACHE_ENABLED == True` + +``` +New request → no session → endpoint resolved + │ + ▼ + ── L1 check ────────────────────────────────────────────────────────── + get_l1(chat_id, endpoint["name"], pre_extracted_params) + │ + ├─ HIT → _LoopStep(kind="cached_response", cache_source="L1") + │ formatter receives cached raw response — no API call, no loop + │ + └─ MISS → continue to L2 + │ + ── L2 check ────────────────────────────────────────────────────────── + get_l2(chat_id) → find entry where api_name == endpoint["name"] + │ + ├─ No match → fall through to normal agentic loop + │ + └─ Match found → FollowUpDetectorModule (DSPy via asyncio.to_thread) + │ + │ Inputs: user_query, previous_query, previous_params, params_schema + │ + ├─ "response_question" + │ → _LoopStep(kind="cached_response", cache_source="L2", + │ cached_raw_response=matching.raw_response) + │ no API call; formatter answers from the previous response + │ + ├─ "param_update" + │ merged = {**matching.collected_params, **updated_params} + │ missing = _missing_required_params(schema, merged) + │ + │ missing == [] + │ ├─ hashes equal (params unchanged) + │ │ → try L1 with matching.collected_params + │ │ hit → cached_response (L1) + │ │ miss → cached_response (L2 raw_response) + │ └─ hashes differ (genuinely new params) + │ → _LoopStep(kind="api_call", collected_params=merged) + │ API called directly — entire agentic loop skipped + │ + │ missing != [] + │ → context["seeded_params"] = merged + │ fall through to agentic loop — only asks for gaps + │ + └─ "new_intent" + → ignore L2; fall through to normal agentic loop + + On any FollowUpDetectorModule exception → fall through to normal loop (fail-open) +``` + +--- + +### Component: `FollowUpDetectorModule` + +Defined in [src/tool_classifier/follow_up_detector.py](../src/tool_classifier/follow_up_detector.py). + +DSPy `Predict` module that classifies the relationship between the new user query +and the previous API call. Run via `asyncio.to_thread` to avoid blocking the event loop. + +**Inputs:** + +| Input | Description | +|---|---| +| `user_query` | The new user message | +| `previous_query` | The user's original question that triggered the last API call | +| `previous_params` | JSON of param values from the last call | +| `params_schema` | JSON of the endpoint's param schema | + +**Output — three possible values for `follow_up_type`:** + +| Value | Meaning | Action | +|---|---|---| +| `response_question` | User is asking about the data already returned | Pass L2 `raw_response` to formatter; no API call | +| `param_update` | User wants the same endpoint with different/additional params | Merge new params into previous; go to API directly if complete, else seed the loop | +| `new_intent` | Completely unrelated query | Ignore L2; run normal agentic loop from scratch | + +**Security:** `updated_params` from the LLM is validated against the endpoint's param +schema — keys not in the schema are silently dropped to prevent injection. + +**Fail-open:** any exception returns `{follow_up_type: "new_intent", updated_params: {}}` +so the user is never blocked. + +--- + +### Param Seeding + +When the L2 `param_update` path finds that merged params are still incomplete, +it sets `context["seeded_params"] = merged` before falling through to the agentic loop. + +`AgenticLoop.run_turn()` and `stream_run_turn()` accept an optional `seeded_params` argument. +On turn 0, the seeds are merged into `collected_params` **before** any extraction runs: + +```python +if turn_count == 0 and seeded_params: + collected_params = {**seeded_params, **collected_params} +``` + +The merge order means existing `collected_params` win — seeds cannot overwrite values +that were already explicitly provided. The seeds are also stored directly in the new +Redis session (`APIToolSession.collected_params = seeded_params`) so they survive +across HTTP requests. + +**Effect:** the agentic loop starts with inherited values already populated and only +generates a question for the genuinely missing params. + +--- + +### L2 Invalidation on Intent Switch + +When an intent switch is detected in `ToolClassifier.classify()` (user mid-session for +endpoint A sends a message that strongly matches endpoint B), both the session and the +L2 key are cleaned up: + +```python +await session_store.delete(request.chatId) # existing behaviour +if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED: + await ATCCacheStore().invalidate_l2(request.chatId) +``` + +`invalidate_l2` deletes only the `atc:last:{chat_id}` key. L1 keys are **not** deleted — +they are param-hash-scoped and expire on their own TTL. Deleting L1 would provide no +safety benefit and would waste valid cached data. + +--- + +### Cache Constants and Feature Flag + +Defined in [src/tool_classifier/constants.py](../src/tool_classifier/constants.py) and +[src/llm_orchestrator_config/feature_flags.py](../src/llm_orchestrator_config/feature_flags.py): + +| Name | Value | Description | +|---|---|---| +| `ATC_CACHE_KEY_PREFIX` | `atc:cache` | Redis key prefix for L1 entries | +| `ATC_LAST_CALL_KEY_PREFIX` | `atc:last` | Redis key prefix for L2 entries | +| `ATC_CACHE_DEFAULT_TTL_SECONDS` | `1800` | Default L1 TTL (30 min); overridable per endpoint via `cache_ttl_seconds` | +| `ATC_LAST_CALL_TTL_SECONDS` | `1800` | L2 TTL (30 min, sliding) | +| `ATC_RESPONSE_CACHE_ENABLED` | `true` (env) | Master kill-switch — disables all reads and writes when `false` | + +--- + +### Cache Component Reference + +| Class / File | Responsibility | +|---|---| +| `ATCCacheStore` ([src/utils/atc_cache_store.py](../src/utils/atc_cache_store.py)) | All Redis operations for L1 and L2; param normalisation and hashing | +| `FollowUpDetectorModule` ([src/tool_classifier/follow_up_detector.py](../src/tool_classifier/follow_up_detector.py)) | DSPy classifier for follow-up type detection | +| `LastCallContext` ([src/models/session_models.py](../src/models/session_models.py)) | Pydantic model stored in L2 | +| `_compute_loop_step` ([src/tool_classifier/workflows/api_tool_workflow.py](../src/tool_classifier/workflows/api_tool_workflow.py)) | Where L1 + L2 are read and routing decisions are made | +| `_execute_api_and_format` / `_stream_api_and_format` | Where L1 + L2 are written after single-intent calls | +| `_execute_multi_api_and_format` / `_stream_multi_api_and_format` | Where L1 + L2 are written after parallel calls | +| `ToolClassifier.classify` ([src/tool_classifier/classifier.py](../src/tool_classifier/classifier.py)) | L2 invalidation on intent switch | + +--- + +### End-to-End Cache Example + +``` +Turn 1 — "What are public holidays in Estonia in 2026?" + Agentic loop collects {countryIsoCode:"EE", validFrom:"2026-01-01", validTo:"2026-12-31"} + API called → 12 holidays returned + L1 written: atc:cache:{id}:get_national_holidays:{hash({EE,2026-01-01,2026-12-31})} + L2 written: atc:last:{id} = [LastCallContext{api_name="get_national_holidays", ...}] + +Turn 2 — "Same for Latvia?" + No session → L1 miss (country changed) → L2 hit + FollowUpDetectorModule → param_update, updated_params={countryIsoCode:"LV"} + merged = {countryIsoCode:"LV", validFrom:"2026-01-01", validTo:"2026-12-31"} + no missing params + hashes differ → api_call step + API called directly — zero agentic loop turns + L1 + L2 updated with new result + +Turn 3 — "Which of those is a bank holiday?" + No session → L1 miss → L2 hit + FollowUpDetectorModule → response_question + Formatter receives Latvia raw_response from L2 → answers from cached data + No API call, no loop + +Turn 4 — "What is the weather in Tallinn?" + Classifier: get_weather matched (different endpoint) + Intent switch → session_store.delete + invalidate_l2 + Fresh session for get_weather starts with empty L1 and L2 +``` + +--- \ No newline at end of file diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md new file mode 100644 index 00000000..43a9a532 --- /dev/null +++ b/docs/ARCHITECTURE.md @@ -0,0 +1,98 @@ +# Architecture + +This page describes how the **LLM Module** fits together, using the +[C4 model](https://c4model.com/) to move from a high-level system view down to the internal +components of the LLM Orchestration Service. Each level links out to the detailed flow documents +that explain the behaviour in depth. + +> The LLM Module is a **multi-workflow orchestrator** for the Bürokratt / Estonian Government AI +> assistant. A tool classifier inspects every user query and routes it to the most appropriate +> workflow — **Service**, **Context**, **RAG**, **API-Tool**, or **Out-of-Domain (OOD)**. +--- + +## Level 1 — System Context + +The context diagram shows the LLM Module as a single system, the people who use it, and the external +systems it depends on (LLM providers, the Central Knowledge Base, observability tooling). + +![LLM Module — C4 System Context Diagram](./images/LLM%20Module%20Context%20Diagram%20(Current).png) + +**Key relationships** + +- **End users / chatbot** send natural-language queries and receive grounded, cited answers. +- **Administrators** configure LLM connections, prompts, budgets, and view analytics. +- **LLM providers** (Azure OpenAI, AWS Bedrock, OpenAI, Anthropic, Google Cloud, self-hosted) supply + chat and embedding models, selected per connection. +- **Central Knowledge Base (CKB)** provides the source content that is indexed for retrieval. +- **Observability** (Langfuse, Grafana/Loki) captures traces, costs, and logs. + +--- + +## Level 2 — Containers + +The container diagram zooms into the deployable units of the system and the data stores they rely on. + +![LLM Module — C4 Container Diagram](./images/LLM%20Module%20App%20Diagram%20(Current).png) + +**Containers & data stores** + +| Container / Store | Role | +| --- | --- | +| **GUI** | Admin web interface for connections, prompts, budgets, and analytics. | +| **Ruuter (public/private)** | API gateway that routes and authorises requests to backend services. | +| **LLM Orchestration Service** | FastAPI service (port `8100`) — the core that runs the workflows. | +| **Notification Server** | Node service pushing real-time updates (e.g. cost alerts, streaming relay). | +| **Qdrant** | Vector database for knowledge-base and API-tool embeddings. | +| **Redis** | Conversation history, session state, and rate-limit counters. | +| **PostgreSQL + ClickHouse** | Relational + columnar stores backing Langfuse analytics. | +| **MinIO (S3)** | Object storage for datasets and documents. | +| **HashiCorp Vault** | Encrypted storage for LLM provider credentials. | + +For how requests traverse the gateway and the orchestration service, see +[API_REFERENCE.md](./API_REFERENCE.md). + +--- + +## Level 3 — Components (LLM Orchestration Service) + +The component diagram opens up the LLM Orchestration Service to show its internal building blocks: the +tool classifier, the per-workflow executors, the contextual retriever, response generation, and +guardrails. + +![LLM Orchestration Service — C4 Component Diagram](./images/LLM%20Orchestration%20Service%20Component%20Diagram%20(Current).png) + +### Request lifecycle (high level) + +1. **Validation & safety** — input is sanitised and checked against guardrails. +2. **Tool classification** — hybrid dense + sparse (BM25) search over indexed examples routes the + query to a workflow. See [HYBRID_SEARCH_CLASSIFICATION.md](./HYBRID_SEARCH_CLASSIFICATION.md). +3. **Workflow execution** — one of: + - **Service** — maps the query to a backend service/intent. + See [TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md](./TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md). + - **Context** — greeting/conversation handling with Redis-backed history. + See [CONTEXT_WORKFLOW_GREETING_DETECTION.md](./CONTEXT_WORKFLOW_GREETING_DETECTION.md). + - **RAG** — contextual retrieval (hybrid search + RRF fusion) then grounded generation. + See [CONTEXTUAL_RETRIEVAL_FLOW.md](./CONTEXTUAL_RETRIEVAL_FLOW.md). + - **API-Tool** — agentic multi-endpoint API calling. + See [API_TOOL_CALLING.md](./API_TOOL_CALLING.md). + - **OOD** — graceful fallback for out-of-domain queries. +4. **Generation & guardrails** — the response is generated with citations and re-checked before return. +5. **Observability** — the full trace and cost are recorded in Langfuse; logs go to Loki. + +### Component → documentation map + +| Concern | Detailed doc | +| --- | --- | +| Tool classifier overview | [TOOL_CLASSIFIER.md](./TOOL_CLASSIFIER.md) | +| Hybrid search & intent enrichment | [HYBRID_SEARCH_CLASSIFICATION.md](./HYBRID_SEARCH_CLASSIFICATION.md) | +| Contextual retrieval (RAG) | [CONTEXTUAL_RETRIEVAL_FLOW.md](./CONTEXTUAL_RETRIEVAL_FLOW.md) | +| API tool calling | [API_TOOL_CALLING.md](./API_TOOL_CALLING.md) | +| Service workflow (UI trace) | [TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md](./TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md) | +| Conversation history & sessions | [REDIS_SESSION_STORE.md](./REDIS_SESSION_STORE.md), [CONTEXT_WORKFLOW_GREETING_DETECTION.md](./CONTEXT_WORKFLOW_GREETING_DETECTION.md) | +| LLM credentials & Vault | [LLM_CONFIG_VAULT_INTEGRATION.md](./LLM_CONFIG_VAULT_INTEGRATION.md), [VAULT_SETUP_AND_USAGE.md](./VAULT_SETUP_AND_USAGE.md), [VAULT_SECURITY_ARCHITECTURE.md](./VAULT_SECURITY_ARCHITECTURE.md) | +| Connection swapping | [CONNECTION_SWAP_FLOW.md](./CONNECTION_SWAP_FLOW.md) | +| Prompt configuration | [CUSTOM_PROMPT_CONFIGURATION.md](./CUSTOM_PROMPT_CONFIGURATION.md) | + +--- + +For the full catalogue of documentation, see the [Documentation Index](./README.md). \ No newline at end of file diff --git a/docs/CONNECTION_SWAP_FLOW.md b/docs/CONNECTION_SWAP_FLOW.md new file mode 100644 index 00000000..bf071534 --- /dev/null +++ b/docs/CONNECTION_SWAP_FLOW.md @@ -0,0 +1,305 @@ +# LLM Connection Swapping Flow + +## Overview + +LLM connections use a **UUID-based Vault path** design where environment swaps are **pure database operations** with zero Vault I/O. Each connection stores its credentials in Vault once at a fixed path (`secret/llm/connections/{platform}/{vault_uuid}`), and switching which connection is "production" vs "testing" only updates the `environment` column in PostgreSQL. + +## Architecture Diagram + +```mermaid +flowchart TD + subgraph GUI["GUI (React)"] + A[User clicks Swap/Edit] + end + + subgraph Ruuter["Ruuter Private DSL"] + B[POST /llm-connections/edit] + C[Demote old production → testing] + D[Clear LLM cache] + E[Promote connection → production] + end + + subgraph Resql["Resql (PostgreSQL)"] + F[(rag_search.llm_connections)] + end + + subgraph LLM["LLM Orchestration Service"] + G[POST /cache/clear] + H[ConnectionIdFetcher] + I[POST /orchestrate/stream] + J[_initialize_llm_manager] + K[ConfigurationLoader] + L[SecretResolver] + end + + subgraph Vault["HashiCorp Vault"] + M[secret/llm/connections/platform/uuid] + N[secret/embeddings/connections/platform/uuid] + end + + A --> B + B --> C + C --> F + C --> D + D --> G + G --> H + B --> E + E --> F + + I --> J + J --> H + H --> F + J --> K + K --> L + L --> M + L --> N +``` + +## Database Schema + +```sql +-- Table: rag_search.llm_connections +CREATE TABLE rag_search.llm_connections ( + id SERIAL PRIMARY KEY, + vault_uuid UUID NOT NULL DEFAULT gen_random_uuid() UNIQUE, + connection_name TEXT, + llm_platform TEXT, -- e.g., "azure_openai", "aws_bedrock" + llm_model TEXT, + embedding_platform TEXT, + embedding_model TEXT, + environment TEXT, -- "production" or "testing" + connection_status TEXT, -- "active" or "inactive" + -- ... budget fields, timestamps, etc. +); +``` + +**Key point:** `vault_uuid` is immutable and assigned at row creation. The `environment` column is the only thing that changes during a swap. + +## Vault Path Structure + +``` +secret/ +├── llm/ +│ └── connections/ +│ ├── azure_openai/ +│ │ ├── ← credentials for connection 1 +│ │ └── ← credentials for connection 2 +│ └── aws_bedrock/ +│ └── +└── embeddings/ + └── connections/ + ├── azure_openai/ + │ └── + └── aws_bedrock/ + └── +``` + +Environment is **NOT** part of the Vault path. This is the core design decision — it means swapping environments never touches Vault. + +## Connection Swap Flow (Step by Step) + +### 1. User Initiates Swap (GUI → Ruuter) + +The user edits a "testing" connection and changes its environment to "production" via the GUI. This triggers `POST /llm-connections/edit`. + +### 2. Ruuter Orchestrates the Swap (`edit.yml`) + +```yaml +# Step 1: Check if promoting testing → production +check_deployment_environment: + condition: environment == "production" && existing.environment == "testing" + next: get_existing_production_connection + +# Step 2: Find current production connection +get_existing_production_connection: + call: Resql → get-production-connection + next: update_existing_production_to_testing + +# Step 3: Demote current production to testing (DB only) +update_production_connection: + call: Resql → update-llm-connection-environment + body: { connection_id: , environment: "testing" } + next: clear_llm_cache + +# Step 4: Invalidate LLM service cache +clear_llm_cache: + call: POST http://llm-orchestration-service:8100/cache/clear + next: update_llm_connection + +# Step 5: Update the edited connection to production (DB only) +update_llm_connection: + call: Resql → update-llm-connection + body: { connection_id: , environment: "production", ... } +``` + +**Total Vault operations: 0** — only PostgreSQL updates and a cache invalidation. + +### 3. Cache Invalidation (`POST /cache/clear`) + +The LLM orchestration service caches `vault_uuid` and `connection_id` in memory (via `ConnectionIdFetcher`). Without invalidation, it would keep using the old production connection's vault_uuid. + +```python +# src/llm_orchestration_service_api.py +@app.post("/cache/clear") +async def clear_connection_cache(): + fetcher = get_connection_id_fetcher() + fetcher.clear_cache() # Clears all cached connection_id + vault_uuid entries + return {"status": "ok"} +``` + +### 4. Next Request Resolves New Connection + +On the next `/orchestrate/stream` request: + +``` +Request arrives → _initialize_llm_manager(environment="production", connection_id=None) + ↓ +ConnectionIdFetcher.fetch_vault_uuid_sync("production") + ↓ (cache is empty after clear) +POST Resql → get-production-connection → returns new production row with vaultUuid + ↓ +vault_uuid cached in memory for subsequent requests + ↓ +ConfigurationLoader(environment="production", connection_id=) + ↓ +SecretResolver._build_vault_path() → "llm/connections/{provider}/{vault_uuid}" + ↓ +VaultAgentClient.get_secret() → fetches credentials from Vault + ↓ +LLM initialized with new credentials +``` + +## Connection Resolution Chain + +```mermaid +sequenceDiagram + participant Client + participant API as FastAPI (/orchestrate/stream) + participant OrcSvc as OrchestrationService + participant Fetcher as ConnectionIdFetcher + participant Resql as Resql (PostgreSQL) + participant Loader as ConfigurationLoader + participant Resolver as SecretResolver + participant Vault as HashiCorp Vault + + Client->>API: POST /orchestrate/stream + API->>OrcSvc: orchestrate(request) + OrcSvc->>OrcSvc: _initialize_llm_manager("production", None) + + Note over OrcSvc,Fetcher: Auto-resolve vault_uuid for production + OrcSvc->>Fetcher: fetch_vault_uuid_sync("production") + alt Cache hit + Fetcher-->>OrcSvc: cached vault_uuid + else Cache miss + Fetcher->>Resql: POST /get-production-connection + Resql-->>Fetcher: [{id, vaultUuid, llmPlatform, ...}] + Note over Fetcher: Resql converts snake_case → camelCase + Fetcher-->>OrcSvc: vault_uuid (cached for next time) + end + + OrcSvc->>Loader: load(environment, connection_id=vault_uuid) + Loader->>Resolver: get_secret_for_model(provider, env, "", vault_uuid) + Resolver->>Resolver: _build_vault_path → "llm/connections/{provider}/{vault_uuid}" + Resolver->>Vault: GET secret/llm/connections/{provider}/{vault_uuid} + Vault-->>Resolver: {api_key, endpoint, model, ...} + Resolver-->>Loader: AzureOpenAISecret / AWSBedrockSecret + Loader-->>OrcSvc: config with resolved secrets + OrcSvc-->>API: LLM response (streamed) +``` + +## Where Cache Clear is Triggered + +| Operation | File | Triggers `/cache/clear`? | +|-----------|------|--------------------------| +| Add new production connection (demotes existing) | `add.yml` | Yes | +| Edit connection to promote testing → production | `edit.yml` | Yes | +| Delete a connection | `delete.yml` | No (not needed — deleted connections aren't cached) | +| Edit without environment change | `edit.yml` | No (no swap happening) | + +## Caching Behavior + +### ConnectionIdFetcher Cache (In-Memory) + +```python +# Cache keys: +# "production_connection_id" → int (DB row ID) +# "production_vault_uuid" → str (UUID for Vault path) +# "testing_connection_id" → int +# "testing_vault_uuid" → str + +_connection_cache: Dict[str, int | str] = {} +``` + +- **Populated on:** First request after service start or cache clear +- **Cleared by:** `POST /cache/clear` (called by Ruuter after swap) +- **Thread-safe:** Uses `threading.Lock` + +### SecretResolver Cache (TTL-Based) + +```python +# Cache key: vault_path string → CachedSecret (data + expires_at) +# TTL: 5 minutes (configurable) +_cache: Dict[str, CachedSecret] = {} +``` + +- **Populated on:** Successful Vault read +- **Evicted after:** 5-minute TTL +- **Background refresh:** Expired entries are refreshed asynchronously +- **Not cleared by `/cache/clear`** — addressed by TTL expiry + +### Why Two Caches? + +1. **ConnectionIdFetcher cache** — Avoids repeated DB calls to resolve which connection is currently "production". Cleared instantly on swap. +2. **SecretResolver cache** — Avoids repeated Vault reads for the same credentials. Not cleared on swap because the vault_uuid changes, so the new path is a cache miss anyway. + +## Testing vs Production Connection Resolution + +| Aspect | Production | Testing | +|--------|-----------|---------| +| `connection_id` in request | Optional (auto-resolved from DB) | Required (must be provided) | +| DB lookup | `get-production-connection` SQL | Not needed (UUID in request) | +| Vault path | `llm/connections/{provider}/{auto_resolved_uuid}` | `llm/connections/{provider}/{provided_uuid}` | +| Caching | Cached in ConnectionIdFetcher | Not cached (provided per-request) | + +## Adding a New Connection (Vault Write) + +Vault is only written to during **connection creation**, not during swaps: + +```mermaid +sequenceDiagram + participant GUI + participant Ruuter + participant Resql as Resql (DB) + participant CronMgr as CronManager + participant Vault + + GUI->>Ruuter: POST /llm-connections/add + Ruuter->>Resql: INSERT into llm_connections (generates vault_uuid) + Resql-->>Ruuter: {id, vaultUuid, ...} + + Note over Ruuter: If production: demote old, clear cache + + Ruuter->>CronMgr: POST /vault/secret/create {vaultUuid, platform, credentials} + CronMgr->>CronMgr: build_vault_path → "secret/llm/connections/{platform}/{vaultUuid}" + CronMgr->>Vault: Write credentials to path + Vault-->>CronMgr: OK + CronMgr-->>Ruuter: Success +``` + +## Key Design Decisions + +1. **No environment in Vault paths** — Swapping is instantaneous (DB UPDATE only), no secret migration needed. +2. **UUID generated by PostgreSQL** — `gen_random_uuid()` DEFAULT ensures it's assigned atomically at INSERT. +3. **Cache invalidation via HTTP** — Ruuter calls `/cache/clear` after any swap to ensure the orchestration service picks up the new production connection on the next request. +4. **Resql auto-converts snake_case to camelCase** — `vault_uuid` in PostgreSQL becomes `vaultUuid` in all JSON responses. All DSL/code must use `vaultUuid`. +5. **Graceful degradation** — If Vault is unavailable, `SecretResolver` falls back to last-known-good cached credentials. + +## Troubleshooting + +| Symptom | Likely Cause | Fix | +|---------|-------------|-----| +| Old connection still used after swap | Cache not cleared | Verify `clear_llm_cache` step fires in Ruuter DSL; manually call `POST /cache/clear` | +| `vault_uuid` is null in Ruuter | Using `vault_uuid` instead of `vaultUuid` | Resql converts to camelCase — always use `vaultUuid` in DSL | +| "No production connection found" | DB has no row with `environment = 'production'` | Create a production connection via GUI | +| Vault secret not found | Vault path mismatch | Verify CronManager `store_secrets_in_vault.sh` used same `vaultUuid` as DB | +| Stale credentials after 5+ minutes | SecretResolver TTL expired but Vault is down | Check Vault connectivity; fallback cache serves last-known-good | diff --git a/docs/CONTEXTUAL_RETRIEVAL_FLOW.md b/docs/CONTEXTUAL_RETRIEVAL_FLOW.md index c59c342c..9c6c0d5b 100644 --- a/docs/CONTEXTUAL_RETRIEVAL_FLOW.md +++ b/docs/CONTEXTUAL_RETRIEVAL_FLOW.md @@ -82,12 +82,12 @@ For each of the 6 refined queries, the system performs parallel semantic and BM2 - **<0.3**: Likely irrelevant **0.4 is the optimal balance** because: -- ✅ Captures semantically related content beyond exact matches -- ✅ Includes contextual information (e.g., implementation details, legal context) -- ✅ Maintains quality while maximizing diversity -- ✅ Industry standard for production RAG systems -- ❌ Lower values (0.3) introduce too much noise -- ❌ Higher values (0.5+) miss valuable context +- Captures semantically related content beyond exact matches +- Includes contextual information (e.g., implementation details, legal context) +- Maintains quality while maximizing diversity +- Industry standard for production RAG systems +- Lower values (0.3) introduce too much noise +- Higher values (0.5+) miss valuable context **Performance Impact:** - Threshold 0.5: ~17 results, 4 unique chunks (too narrow) @@ -172,10 +172,10 @@ The k-parameter determines how quickly scores decay with rank position: | k=90 | 0.0110 | 0.0100 | Very narrow | Too democratic | **k=35 Advantages:** -- ✅ **65-70% higher top-rank scores** vs k=60 (0.0541 vs 0.0328) -- ✅ **Clear score separation** between highly relevant and marginal chunks -- ✅ **Balanced approach** - respects both top results and broader context -- ✅ **Better signal for response generator** - easier to identify best chunks +- **65-70% higher top-rank scores** vs k=60 (0.0541 vs 0.0328) +- **Clear score separation** between highly relevant and marginal chunks +- **Balanced approach** - respects both top results and broader context +- **Better signal for response generator** - easier to identify best chunks **Score Differentiation Example:** ``` @@ -302,12 +302,12 @@ For each of the top 10 chunks: | Metric | Value | Target | Status | |--------|-------|--------|--------| -| Semantic Results per Query | 27.3 | >5 | ✅ Excellent | -| Unique Semantic Chunks | 42 | >10 | ✅ Excellent | -| Fusion Coverage | 100% | >80% | ✅ Perfect | -| Both-sources Validation | 12/12 | >50% | ✅ Perfect | -| Score Differentiation | High | Clear gaps | ✅ Excellent | -| Retrieval Speed | 1.6s | <3s | ✅ Excellent | +| Semantic Results per Query | 27.3 | >5 | Excellent | +| Unique Semantic Chunks | 42 | >10 | Excellent | +| Fusion Coverage | 100% | >80% | Perfect | +| Both-sources Validation | 12/12 | >50% | Perfect | +| Score Differentiation | High | Clear gaps | Excellent | +| Retrieval Speed | 1.6s | <3s | Excellent | --- @@ -573,10 +573,10 @@ When evaluating the quality of the contextual retrieval system and response gene ### Alert Thresholds -- ⚠️ Semantic yield drops below 5 results/query -- ⚠️ Fusion coverage drops below 80% -- ⚠️ Retrieval time exceeds 3 seconds -- ⚠️ BM25 index build fails or incomplete +- Semantic yield drops below 5 results/query +- Fusion coverage drops below 80% +- Retrieval time exceeds 3 seconds +- BM25 index build fails or incomplete --- diff --git a/docs/CONTEXT_WORKFLOW_GREETING_DETECTION.md b/docs/CONTEXT_WORKFLOW_GREETING_DETECTION.md index 8a67e841..66d88603 100644 --- a/docs/CONTEXT_WORKFLOW_GREETING_DETECTION.md +++ b/docs/CONTEXT_WORKFLOW_GREETING_DETECTION.md @@ -2,12 +2,14 @@ ## Overview -The **Context Workflow (Layer 2)** intercepts user queries that can be answered without searching the knowledge base. It handles two categories: +The **Context Workflow (Layer 3)** intercepts user queries that can be answered without searching the knowledge base. It handles two categories: 1. **Greetings** — Detects and responds to social exchanges (hello, goodbye, thanks) in multiple languages 2. **Conversation history references** — Answers follow-up questions that refer to information already discussed in the session -When the context workflow can answer, a response is returned immediately, bypassing the RAG pipeline entirely. When it cannot answer, the query falls through to the RAG workflow (Layer 3). +Conversation history is now sourced from a **Redis-backed store** (canonical source) rather than the GUI-provided `request.conversationHistory`. The store retains the most recent 10 rounds per session and maintains an incremental summary of evicted older rounds, enabling context detection to cover the full conversation lifetime. + +When the context workflow can answer, a response is returned immediately, bypassing the RAG pipeline entirely. When it cannot answer, the query falls through to the RAG workflow (Layer 4). --- @@ -18,23 +20,31 @@ When the context workflow can answer, a response is returned immediately, bypass ``` User Query ↓ -Layer 1: SERVICE → External API calls +Layer 1: SERVICE → External API calls + ↓ (cannot handle) +Layer 2: API_TOOL_CALLING → Agentic API tool execution ↓ (cannot handle) -Layer 2: CONTEXT → Greetings + conversation history ←── This document +Layer 3: CONTEXT → Greetings + conversation history ←── This document ↓ (cannot handle) -Layer 3: RAG → Knowledge base retrieval +Layer 4: RAG → Knowledge base retrieval ↓ (cannot handle) -Layer 4: OOD → Out-of-domain fallback +Layer 5: OOD → Out-of-domain fallback ``` +> **Note**: The classifier also checks for an **active API tool session** (`chatId` present in `session_store`) before any layer evaluation. If found, the request short-circuits directly to the `API_TOOL_CALLING` workflow to continue parameter collection — the context workflow is never reached in this path. + ### Key Components | Component | File | Responsibility | |-----------|------|----------------| -| `ContextAnalyzer` | `src/tool_classifier/context_analyzer.py` | LLM-based greeting detection and context analysis | -| `ContextWorkflowExecutor` | `src/tool_classifier/workflows/context_workflow.py` | Orchestrates the workflow, handles streaming/non-streaming | +| `ContextAnalyzer` | `src/tool_classifier/context_analyzer.py` | LLM-based greeting detection, context analysis, summary generation | +| `ContextWorkflowExecutor` | `src/tool_classifier/workflows/context_workflow.py` | Orchestrates the workflow, history fetching, streaming/non-streaming | | `ToolClassifier` | `src/tool_classifier/classifier.py` | Invokes `ContextAnalyzer` during classification and routes to `ContextWorkflowExecutor` | -| `greeting_constants.py` | `src/tool_classifier/greeting_constants.py` | Fallback greeting responses for Estonian and English | +| `FeatureFlags.CONTEXT_WORKFLOW_ENABLED` | `src/llm_orchestrator_config/feature_flags.py` | Guards the context workflow; if `False`, the layer is skipped in the fallback chain and the request proceeds directly to RAG | +| `ConversationHistoryStore` | `src/utils/conversation_history_store.py` | Redis CRUD store for per-session rounds and incremental summary | +| `conversation_summary_generator` | `src/utils/conversation_summary_generator.py` | Factory for the incremental summarizer callable injected into the store | +| `redis_client` | `src/utils/redis_client.py` | Singleton async Redis client (db=1, TLS-capable) | +| `greeting_constants.py` | `src/tool_classifier/greeting_constants.py` | Static greeting response templates for Estonian and English | --- @@ -44,47 +54,105 @@ Layer 4: OOD → Out-of-domain fallback User Query + Conversation History ↓ ToolClassifier.classify() + ├─ Pre-check: active API tool session for chatId? + │ └─ Yes → short-circuit to API_TOOL_CALLING (context skipped) + │ ├─ Layer 1 (SERVICE): Embedding-based intent routing - │ └─ If no service tool matches → route to CONTEXT workflow + │ └─ If no service tool matches → try Layer 2 + │ + ├─ Layer 2 (API_TOOL_CALLING): Semantic search in api_tool_collection + │ └─ If no API tool matches (or feature disabled) → route to CONTEXT workflow │ └─ ClassificationResult(workflow=CONTEXT) ToolClassifier.route_to_workflow() ├─ Non-streaming → ContextWorkflowExecutor.execute_async() - │ ├─ Phase 1: _detect() → context_analyzer.detect_context() [classification only] - │ ├─ If greeting → return greeting OrchestrationResponse + │ ├─ _build_history() → ConversationHistoryStore.get_context() [Redis, with fallback] + │ ├─ Phase 1: _detect() → context_analyzer.detect_context_with_summary_fallback() + │ │ ├─ Step 1: detect_context() on last 10 turns + │ │ ├─ Step 2 (if needed): use pre_computed_summary from Redis OR generate summary + │ │ └─ Step 3 (if needed): _analyze_from_summary() on summary + │ ├─ If greeting → return greeting OrchestrationResponse (static template) │ ├─ If can_answer → _generate_response_async() → context_analyzer.generate_context_response() │ └─ Otherwise → return None (RAG fallback) │ └─ Streaming → ContextWorkflowExecutor.execute_streaming() - ├─ Phase 1: _detect() → context_analyzer.detect_context() [classification only] - ├─ If greeting → _stream_greeting() async generator + ├─ _build_history() → ConversationHistoryStore.get_context() [Redis, with fallback] + ├─ Phase 1: _detect() → context_analyzer.detect_context_with_summary_fallback() + ├─ If greeting → _stream_greeting() async generator (static template) ├─ If can_answer → _create_history_stream() → context_analyzer.stream_context_response() └─ Otherwise → return None (RAG fallback) ``` --- +## Redis-Backed Conversation History + +### Overview + +`ConversationHistoryStore` is a Redis-backed CRUD store (db=1) that holds per-session conversation data. It is the **canonical source of truth** for conversation history, replacing the GUI-provided `request.conversationHistory` when available. + +### Key Layout + +| Redis Key | Content | TTL | +|-----------|---------|-----| +| `conv:{chat_id}` | JSON list of up to 10 `ConversationRound` objects | 30 minutes (sliding) | +| `conv:summary:{chat_id}` | Plain-text incremental summary of evicted rounds | 30 minutes (sliding) | + +Both keys share a sliding TTL: every write resets the expiry on both keys to keep them in sync. + +### History Capping and Eviction + +The store caps history at **10 rounds** (`_MAX_ROUNDS`). When appending a new round causes the count to exceed 10, the oldest rounds are trimmed. Trimmed (evicted) rounds are passed to an optional `summarizer` callable as a fire-and-forget background `asyncio.Task`, which merges them into the running summary using `IncrementalSummarySignature`. + +### `_build_history()` — History Resolution in the Workflow + +`ContextWorkflowExecutor._build_history()` resolves the history and pre-computed summary to pass to Phase 1: + +1. If `ConversationHistoryStore` is wired in, call `get_context(chat_id)` to retrieve rounds and the optional Redis summary. +2. If rounds are present, flatten them into `{"authorRole", "message", "timestamp"}` dicts and return `(history, summary)`. +3. If the store is absent, raises, or returns no rounds → fall back to `request.conversationHistory` with `summary=None`. + +The returned `pre_computed_summary` is forwarded to `detect_context_with_summary_fallback()` to skip an expensive LLM summarisation step when Redis already has one. + +### Optimistic Locking + +`save_round()` uses Redis `WATCH`/`MULTI`/`EXEC` (optimistic locking) to detect concurrent writes and retries up to 3 times on conflict. + +--- + ## Phase 1: Detection (Classify Only) -### LLM Task +### Three-Step Detection Flow -Every query is checked against the **most recent 10 conversation turns** using a single LLM call (`detect_context()`). This phase **does not generate an answer** — it only classifies the query and extracts a relevant context snippet for Phase 2. +Every query is processed by `detect_context_with_summary_fallback()`, which implements a three-step detection pipeline: -The `ContextDetectionSignature` DSPy signature instructs the LLM to: +**Step 1 — Recent turns check (`detect_context`)** -1. Detect if the query is a greeting in any supported language -2. Check if the query references something discussed in the last 10 turns -3. If the query can be answered from history, extract the relevant snippet -4. Do **not** generate the final answer here — detection only +Runs `ContextDetectionSignature` via `dspy.ChainOfThought` against the **most recent 10 conversation turns**. This phase **does not generate an answer** — it only classifies the query and extracts a relevant context snippet for Phase 2. + +**Step 2 — Summary path (triggered when Step 1 cannot answer)** + +Triggered when `can_answer_from_context=False` AND one of the following is true: +- Total history exceeds 10 turns (older turns exist), OR +- Redis has a `pre_computed_summary` (covers evicted rounds beyond the current active window) + +Two sub-paths: +- **Redis path**: Pre-computed summary is available → used directly (no LLM call, zero cost). +- **On-demand path**: No pre-computed summary → older turns (beyond last 10) are summarised via `ConversationSummarySignature`. + +**Step 3 — Summary analysis (`_analyze_from_summary`)** + +Runs `SummaryAnalysisSignature` against the summary string to determine if the query can be answered from it. If so, the summary-derived answer is returned as `context_snippet` (with `answered_from_summary=True`) for Phase 2 generation. ### LLM Output Format -The LLM returns a JSON object parsed into `ContextDetectionResult`: +`ContextDetectionSignature` returns a JSON object parsed into `ContextDetectionResult`: ```json { "is_greeting": false, + "greeting_type": "hello", "can_answer_from_context": true, "reasoning": "User is asking about tax rate discussed earlier", "context_snippet": "Bot confirmed the flat rate is 20%, applying equally to all income brackets." @@ -94,18 +162,18 @@ The LLM returns a JSON object parsed into `ContextDetectionResult`: | Field | Type | Description | |-------|------|-------------| | `is_greeting` | `bool` | Whether the query is a greeting | +| `greeting_type` | `str` | One of `hello`, `goodbye`, `thanks`, `casual` (relevant when `is_greeting=True`) | | `can_answer_from_context` | `bool` | Whether the query can be answered from conversation history | | `reasoning` | `str` | Brief explanation of the detection decision | -| `context_snippet` | `str \| null` | Relevant excerpt from history for use in Phase 2, or `null` | - -> **Internal field**: `answered_from_summary` (bool, default `False`) is reserved for future summary-based detection paths. +| `context_snippet` | `str \| null` | Relevant excerpt from history or summary for Phase 2, or `null` | +| `answered_from_summary` | `bool` | `True` when the answer was derived from the summary path (internal, default `False`) | ### Decision After Phase 1 ``` -is_greeting=True → Phase 2: return greeting response (no LLM call) +is_greeting=True → Phase 2: return greeting response (static template, no LLM) can_answer_from_context=True AND snippet set → Phase 2: generate answer from snippet -Otherwise → Fall back to RAG +Otherwise (all steps exhausted) → Fall back to RAG ``` --- @@ -118,9 +186,12 @@ Calls `generate_context_response(query, context_snippet)` which uses `ContextRes ### Streaming (`_create_history_stream` → `stream_context_response`) -Calls `stream_context_response(query, context_snippet)` which uses DSPy native streaming (`dspy.streamify`) with `ContextResponseGenerationSignature`. Tokens are yielded in real time and passed through NeMo Guardrails before being SSE-formatted. +Calls `stream_context_response(query, context_snippet)` which uses DSPy native streaming (`dspy.streamify`) with `ContextResponseGenerationSignature`. A fresh `StreamListener` is created per call to avoid stale state. Tokens are yielded in real time and passed through NeMo Guardrails before being SSE-formatted. ---- +**Fallback chain inside `stream_context_response`:** +1. DSPy `streamify` → yield `StreamResponse` tokens as they arrive. +2. If no stream tokens received but the final `Prediction` has an answer, yield it in word-group chunks. +3. If that is also empty, call `generate_context_response()` directly and yield its result in word-group chunks. --- @@ -144,28 +215,28 @@ Calls `stream_context_response(query, context_snippet)` which uses DSPy native s ### Greeting Response Generation -Greeting detection is handled in **Phase 1 (`detect_context`)**, where the LLM classifies whether the query is a greeting and, if so, identifies the language and greeting type. This phase does **not** generate the final natural-language reply. -In **Phase 2**, `ContextWorkflowExecutor` calls `get_greeting_response(...)`, which returns a response based on predefined static templates in `greeting_constants.py`, ensuring the reply is in the detected language. If greeting detection fails or the greeting type is unsupported, the query falls through to the next workflow layer instead of attempting LLM-based greeting generation. +Greeting detection is handled in **Phase 1 (`detect_context`)**, where the LLM classifies whether the query is a greeting, identifies the `greeting_type`, and sets `is_greeting=True`. A message is only treated as a greeting if it contains **nothing beyond the greeting itself** — a greeting combined with a question is routed to RAG instead. + +In **Phase 2**, `ContextWorkflowExecutor` calls `get_greeting_response(greeting_type=..., language=...)`, which returns a static template from `greeting_constants.py`. The language is determined by `detect_language()` on the user query. No LLM call is made for greeting responses. + **Greeting response templates (`greeting_constants.py`):** ```python GREETINGS_ET = { - "hello": "Tere! Kuidas ma saan sind aidata?", + "hello": "Tere! Kuidas ma saan sind aidata?", "goodbye": "Nägemist! Head päeva!", - "thanks": "Palun! Kui on veel küsimusi, küsi julgelt.", - "casual": "Tere! Mida ma saan sinu jaoks teha?", + "thanks": "Palun! Kui on veel küsimusi, küsi julgelt.", + "casual": "Tere! Mida ma saan sinu jaoks teha?", } GREETINGS_EN = { - "hello": "Hello! How can I help you?", + "hello": "Hello! How can I help you?", "goodbye": "Goodbye! Have a great day!", - "thanks": "You're welcome! Feel free to ask if you have more questions.", - "casual": "Hey! What can I do for you?", + "thanks": "You're welcome! Feel free to ask if you have more questions.", + "casual": "Hey! What can I do for you?", } ``` -The fallback greeting type is determined by keyword matching in `_detect_greeting_type()` — checking for `thank/tänan/aitäh`, `bye/goodbye/nägemist/tšau`, before defaulting to `hello`. - --- ## Streaming Support @@ -174,7 +245,7 @@ The context workflow supports both response modes: ### Non-Streaming (`execute_async`) -Returns a complete `OrchestrationResponse` object with the answer as a single string. Output guardrails are applied before the response is returned. +Returns a complete `OrchestrationResponse` object with the answer as a single string. Output guardrails are applied before the response is returned. If a `pre_computed_analysis_result` is present in the classifier context, Phase 1 is skipped entirely (reuses the already-computed detection). ### Streaming (`execute_streaming`) @@ -184,6 +255,10 @@ Returns an `AsyncIterator[str]` that yields SSE (Server-Sent Events) chunks. **History responses** use DSPy native streaming (`dspy.streamify`) with `ContextResponseGenerationSignature`. Tokens are emitted in real time as they arrive from the LLM, then passed through NeMo Guardrails (`stream_with_guardrails`) before being SSE-formatted. If a guardrail violation is detected in a chunk, streaming stops and the violation message is sent instead. +### Conversation History Persistence (Streaming) + +After streaming completes, `llm_orchestration_service.py` saves a `ConversationRound` to `ConversationHistoryStore` — but **only for non-RAG workflows** (SERVICE, API_TOOL_CALLING, CONTEXT). RAG has its own internal save hook inside `_stream_rag_pipeline` and does not go through this path. The accumulated content is filtered to exclude SSE control tokens (`END`) and predefined excluded messages before saving. + **SSE Format:** ``` data: {"chatId": "abc123", "payload": {"content": "Tere! Kuidas ma"}, "timestamp": "...", "sentTo": []} @@ -199,12 +274,13 @@ data: {"chatId": "abc123", "payload": {"content": "END"}, "timestamp": "...", "s LLM token usage and cost is tracked via `get_lm_usage_since()` and stored in `costs_metric` within the workflow executor. Costs are logged via `orchestration_service.log_costs()` at the end of each execution path. -Two cost keys are tracked separately: +Two cost keys are tracked separately. When the summary fallback path is taken, its LLM calls are **merged into** `context_detection`: ```python costs_metric = { "context_detection": { - # Phase 1: detect_context() — single LLM call + # Phase 1: detect_context() + optional summary generation + summary analysis + # All summary-path costs are merged here via _merge_cost_dicts() "total_cost": 0.0012, "total_tokens": 180, "total_prompt_tokens": 150, @@ -222,9 +298,7 @@ costs_metric = { } ``` -Greeting responses skip Phase 2, so only `"context_detection"` cost is populated. - ---- +Greeting responses skip Phase 2, so only `"context_detection"` cost is populated. When the Redis pre-computed summary is used, the summary-generation cost is zero (no LLM call). --- @@ -232,10 +306,15 @@ Greeting responses skip Phase 2, so only `"context_detection"` cost is populated | Failure Point | Behaviour | |---------------|-----------| +| Redis unavailable (`ConversationHistoryStore`) | Logged as warning → falls back to `request.conversationHistory` | +| Redis fetch raises exception | Logged as warning → falls back to `request.conversationHistory` | | Phase 1 LLM call raises exception | `can_answer_from_context=False` → falls back to RAG | | Phase 1 returns invalid JSON | Logged as warning, all flags default to `False` → falls back to RAG | +| Summary generation (on-demand) fails | Logged as error → summary path skipped → falls back to RAG | +| Summary analysis returns no answer | Logged as info → falls back to RAG | | Phase 2 LLM call raises exception | Logged as error, `_generate_response_async` returns `None` → falls back to RAG | | Phase 2 returns empty answer | Logged as warning → falls back to RAG | +| All Phase 2 streaming fallbacks exhausted | Logged as error → empty response | | Output guardrails fail | Logged as warning, response returned without guardrail check | | Guardrail violation in streaming | `OUTPUT_GUARDRAIL_VIOLATION_MESSAGE` sent, stream terminated | | `orchestration_service` unavailable | History streaming skipped → `None` returned → RAG fallback | @@ -252,13 +331,21 @@ Key log entries emitted during a request: |-------|---------|------| | `INFO` | `CONTEXT WORKFLOW (NON-STREAMING) \| Query: '...'` | `execute_async()` entry | | `INFO` | `CONTEXT WORKFLOW (STREAMING) \| Query: '...'` | `execute_streaming()` entry | +| `DEBUG` | `[chatId] Using Redis history: N rounds, summary=present\|absent` | Redis history fetched successfully | +| `WARNING` | `[chatId] Redis history fetch failed, falling back to request history: ...` | Redis read error | | `INFO` | `CONTEXT DETECTOR: Phase 1 \| Query: '...' \| History: N turns` | `detect_context()` entry | | `INFO` | `DETECTION RESULT \| Greeting: ... \| Can Answer: ... \| Has snippet: ...` | Phase 1 LLM response parsed | | `INFO` | `Detection cost \| Total: $... \| Tokens: N` | After Phase 1 cost tracked | +| `INFO` | `Pre-computed summary available \| Skipping LLM summary generation, using Redis summary directly` | Redis summary reused | +| `INFO` | `History has N turns (> 10) \| Cannot answer from recent 10 \| Attempting summary-based detection` | On-demand summary path triggered | +| `INFO` | `DETECTION: Can answer from summary \| Reasoning: ...` | Summary path answered query | +| `INFO` | `Cannot answer from summary either \| Falling back to RAG` | Summary path failed | | `INFO` | `Detection: greeting=... can_answer=...` | After `_detect()` returns in executor | | `INFO` | `CONTEXT GENERATOR: Phase 2 non-streaming \| Query: '...'` | `generate_context_response()` entry | | `INFO` | `CONTEXT GENERATOR: Phase 2 streaming \| Query: '...'` | `stream_context_response()` entry | | `INFO` | `Context response streaming complete (final Prediction received)` | DSPy streaming finished | +| `WARNING` | `Stream tokens not received — yielding answer from final Prediction in chunks.` | Streaming fallback 1 | +| `WARNING` | `No answer from streamify — falling back to generate_context_response.` | Streaming fallback 2 | | `WARNING` | `[chatId] Phase 2 empty answer — fallback to RAG` | Phase 2 returned no content | | `WARNING` | `[chatId] Guardrails violation in context streaming` | Violation detected mid-stream | | `WARNING` | `[chatId] Cannot answer from context — falling back to RAG` | Neither phase could answer | @@ -267,32 +354,76 @@ Key log entries emitted during a request: ## Data Models +### `ConversationRound` (Redis storage unit) + +```python +class ConversationRound(BaseModel): + user_message: str # The user's message text + bot_message: str # The bot's response text + timestamp: float # Unix timestamp of the round +``` + +### `ConversationHistoryState` (Redis fetch result) + +```python +class ConversationHistoryState(BaseModel): + chat_id: str # Unique conversation identifier + rounds: list[ConversationRound] # Ordered rounds (newest last), capped at 10 + summary: Optional[str] # Incremental summary of evicted older rounds +``` + ### `ContextDetectionResult` (Phase 1 output) ```python class ContextDetectionResult(BaseModel): - is_greeting: bool # True if query is a greeting - can_answer_from_context: bool # True if query can be answered from last 10 turns - reasoning: str # LLM's brief explanation - answered_from_summary: bool # Reserved; always False in current workflow + is_greeting: bool # True if query is a greeting + greeting_type: str # "hello" | "goodbye" | "thanks" | "casual" + can_answer_from_context: ( + bool # True if query can be answered from history or summary + ) + reasoning: str # LLM's brief explanation + answered_from_summary: bool # True when answer derived from summary path context_snippet: Optional[str] # Relevant excerpt for Phase 2 generation, or None ``` -### `ContextDetectionSignature` (DSPy — Phase 1) +### `ContextDetectionSignature` (DSPy — Phase 1, recent turns) | Field | Type | Description | |-------|------|-------------| | `conversation_history` | Input | Last 10 turns formatted as JSON | | `user_query` | Input | Current user query | -| `detection_result` | Output | JSON with `is_greeting`, `can_answer_from_context`, `reasoning`, `context_snippet` | +| `detection_result` | Output | JSON with `is_greeting`, `greeting_type`, `can_answer_from_context`, `reasoning`, `context_snippet` | > Detection only — **no answer generated here**. +### `ConversationSummarySignature` (DSPy — on-demand summary generation) + +| Field | Type | Description | +|-------|------|-------------| +| `conversation_history` | Input | JSON of older turns to summarize | +| `summary` | Output | Concise summary preserving key facts, names, numbers, dates | + +### `IncrementalSummarySignature` (DSPy — background eviction summary) + +| Field | Type | Description | +|-------|------|-------------| +| `existing_summary` | Input | Current summary (may be empty for first eviction) | +| `new_rounds` | Input | JSON array of just-evicted rounds | +| `updated_summary` | Output | Merged summary incorporating new rounds | + +### `SummaryAnalysisSignature` (DSPy — Phase 1, summary path) + +| Field | Type | Description | +|-------|------|-------------| +| `conversation_summary` | Input | Summary of earlier conversation | +| `user_query` | Input | Current user query | +| `analysis_result` | Output | JSON with `can_answer_from_context`, `answer`, `reasoning` | + ### `ContextResponseGenerationSignature` (DSPy — Phase 2) | Field | Type | Description | |-------|------|-------------| -| `context_snippet` | Input | Relevant excerpt from Phase 1 | +| `context_snippet` | Input | Relevant excerpt from Phase 1 (or summary-derived answer) | | `user_query` | Input | Current user query | | `answer` | Output | Natural language response in the same language as the query | @@ -302,11 +433,15 @@ class ContextDetectionResult(BaseModel): | Scenario | Phase 1 LLM Calls | Phase 2 LLM Calls | Outcome | |----------|--------------------|--------------------|---------| -| Greeting detected | 1 (`detect_context`) | 0 (static response) | Context responds (greeting) | +| Greeting detected | 1 (`detect_context`) | 0 (static template) | Context responds (greeting) | | Follow-up answerable from last 10 turns | 1 (`detect_context`) | 1 (`generate_context_response` or `stream_context_response`) | Context responds | -| Cannot answer from last 10 turns | 1 (`detect_context`) | 0 | Falls back to RAG | +| Cannot answer from 10 turns; Redis summary answers | 1 + 1 (`detect_context` + `_analyze_from_summary`) | 1 | Context responds (summary path) | +| Cannot answer from 10 turns; Redis summary reused (no new LLM call) | 1 + 1 (`detect_context` + `_analyze_from_summary`; 0 for summary gen) | 1 | Context responds (Redis summary path) | +| Cannot answer from 10 turns; on-demand summary answers | 1 + 1 + 1 (`detect_context` + `_generate_conversation_summary` + `_analyze_from_summary`) | 1 | Context responds (on-demand summary path) | +| Cannot answer from any path | 1–3 (all detection steps) | 0 | Falls back to RAG | | Phase 1 LLM error / JSON parse failure | — | 0 | Falls back to RAG | -| Phase 2 LLM error or empty answer | 1 | — | Falls back to RAG | +| Phase 2 LLM error or empty answer | 1–3 | — | Falls back to RAG | +| Redis unavailable | 0 (fallback to request history) | varies | Proceeds normally with request history | --- @@ -314,10 +449,14 @@ class ContextDetectionResult(BaseModel): | File | Purpose | |------|---------| -| `src/tool_classifier/context_analyzer.py` | Core LLM analysis logic (all three steps) | -| `src/tool_classifier/workflows/context_workflow.py` | Workflow executor (streaming + non-streaming) | +| `src/tool_classifier/context_analyzer.py` | Core LLM analysis logic (detection, summary generation, response generation) | +| `src/tool_classifier/workflows/context_workflow.py` | Workflow executor (history fetching, streaming + non-streaming) | | `src/tool_classifier/classifier.py` | Classification layer that invokes context analysis | -| `src/tool_classifier/greeting_constants.py` | Static fallback greeting responses (ET/EN) | +| `src/tool_classifier/greeting_constants.py` | Static greeting response templates (ET/EN) | +| `src/utils/conversation_history_store.py` | Redis CRUD store for rounds and incremental summary | +| `src/utils/conversation_summary_generator.py` | Factory for the incremental summarizer callable | +| `src/utils/redis_client.py` | Singleton async Redis client (db=1, TLS-capable) | +| `src/models/conversation_history_models.py` | Pydantic models: `ConversationRound`, `ConversationHistoryState` | | `tests/test_context_analyzer.py` | Unit tests for `ContextAnalyzer` | | `tests/test_context_workflow.py` | Unit tests for `ContextWorkflowExecutor` | | `tests/test_context_workflow_integration.py` | Integration tests for the full classify → route → execute chain | \ No newline at end of file diff --git a/docs/CUSTOM_PROMPT_CONFIGURATION.md b/docs/CUSTOM_PROMPT_CONFIGURATION.md index 8a7f94ef..4647a26c 100644 --- a/docs/CUSTOM_PROMPT_CONFIGURATION.md +++ b/docs/CUSTOM_PROMPT_CONFIGURATION.md @@ -2,14 +2,58 @@ ## Overview -The custom prompt configuration system allows admins to configure prompts via UI that automatically apply to all response generation operations. Changes are cached with a 5-minute TTL and can be immediately refreshed when updated. +The custom prompt configuration system allows admins to configure a single organisation-level prompt via the UI that automatically applies to user-facing answer generation. Changes are cached with a 5-minute TTL and can be immediately refreshed when updated. + +The same configured prompt is consumed by **two** workflows — the **RAG workflow** and the **API Tool Calling workflow** — through one shared `PromptConfigurationLoader`. The **Context workflow** does **not** apply custom prompts (greetings use static templates and history answers use their own signature). See [Where Custom Prompts Are Applied (by Workflow)](#where-custom-prompts-are-applied-by-workflow). + +--- + +## Where Custom Prompts Are Applied (by Workflow) + +All consumers read the same prompt from the shared `PromptConfigurationLoader` +(`src/utils/prompt_config_loader.py`, 5-minute TTL cache). They differ in **how** they inject it. + +| Workflow | Custom prompt applied? | Where / how | +|---|---|---| +| **RAG** | Yes | `ResponseGeneratorAgent` — the prompt is wrapped as `[SYSTEM INSTRUCTIONS]…[USER QUESTION]` and appended to the question for both streaming and non-streaming generation. | +| **API Tool Calling** | Yes | `APIToolWorkflowExecutor._get_custom_instructions()` loads the raw prompt and passes it into parameter extraction and response formatting (see below). | +| **Context** | No | `context_workflow.py` / `context_analyzer.py` do not load or apply the custom prompt. Greetings return static templates; history answers use `ContextResponseGenerationSignature` without injection. | +| **Service** | No | The response is pre-formed text from Ruuter/DMapper — there is no LLM generation step to steer. | +| **OOD** | No | Fixed localized out-of-scope message. | + +### RAG workflow + +- Source: [`src/llm_orchestration_service.py`](../src/llm_orchestration_service.py) → `_get_custom_instructions_for_response_generation()` builds the prefix + `"[SYSTEM INSTRUCTIONS]\n{prompt}\n\n[USER QUESTION]\n"` and passes it as + `ResponseGeneratorAgent(custom_instructions_prefix=…)`. +- Application: [`src/response_generator/response_generate.py`](../src/response_generator/response_generate.py) applies it as + `augmented_question = f"{question}\n\n{custom_instructions_prefix}"` in both `forward()` and + `stream_response()`. +- **Note:** despite the name `custom_instructions_prefix`, the string is **appended after** the + question (the wrapper text itself carries the `[USER QUESTION]` marker). It is **not** applied to + `PromptRefinerAgent`, which only optimises the query for retrieval. + +### API Tool Calling workflow + +- Source: [`src/tool_classifier/workflows/api_tool_workflow.py`](../src/tool_classifier/workflows/api_tool_workflow.py) → `_get_custom_instructions()` reads the **same** `prompt_config_loader` + (via `asyncio.to_thread`, fail-open to `""`). Unlike the RAG path, it passes the **raw** prompt + (no `[SYSTEM INSTRUCTIONS]` wrapper). +- It is injected as a dedicated DSPy `custom_instructions` input field into: + - `ParamExtractionModule` ([`param_extractor.py`](../src/tool_classifier/param_extractor.py)) — steers how parameters are extracted from the user. + - `APIResponseFormatterModule` ([`api_response_formatter.py`](../src/tool_classifier/api_response_formatter.py)) — single-endpoint natural-language answer. + - `MultiResponseFormatterModule` ([`multi_response_formatter.py`](../src/tool_classifier/multi_response_formatter.py)) — multi-endpoint synthesis. +- In the formatter/extractor signatures, a non-empty `custom_instructions` is followed with + **HIGHEST PRIORITY**, overriding defaults such as language policy, tone, and formatting. +- Additionally, the workflow derives the response language from the prompt via + `_language_from_custom_instructions()` and merges any Redis conversation summary into the same + `custom_instructions` string before extraction. --- ## Architecture Components ### 1. **Database Layer** -- **Table**: `public.prompt_configuration` +- **Table**: `rag_search.prompt_configuration` - **Columns**: `id` (BIGINT), `prompt` (TEXT) - Stores the custom prompt text configured by admins @@ -36,7 +80,7 @@ The custom prompt configuration system allows admins to configure prompts via UI - **ResponseGeneratorAgent** (`src/response_generator/response_generate.py`) - Accepts `custom_instructions_prefix` parameter - - Prepends custom instructions to user questions + - Appends custom instructions after the user question - Applied in both streaming and non-streaming modes ### 4. **API Endpoints** @@ -120,8 +164,8 @@ The custom prompt configuration system allows admins to configure prompts via UI ▼ ┌─────────────────────────────────────────────────────────────────┐ │ ResponseGeneratorAgent.forward() or stream_response() │ -│ - Prepends custom_instructions_prefix to user question │ -│ - Modified question = "{prefix}{user_question}" │ +│ - Appends custom_instructions_prefix after user question │ +│ - Modified question = "{user_question}{prefix}" │ └────────────────┬────────────────────────────────────────────────┘ │ ▼ @@ -169,8 +213,7 @@ The custom prompt configuration system allows admins to configure prompts via UI ### **User Request Processing** 1. **Request Received** (Any of 3 endpoints) - - `/orchestrate` - Standard response - - `/orchestrate/test` - Test response + - `/orchestrate/test` - Test Sresponse - `/orchestrate/stream` - Streaming response 2. **Service Components Initialization** @@ -242,7 +285,7 @@ RAG_SEARCH_PROMPT_REFRESH=http://llm-orchestration-service:8100/prompt-config/re ### **1. Insert Test Prompt** ```sql -INSERT INTO public.prompt_configuration (id, prompt) +INSERT INTO rag_search.prompt_configuration (id, prompt) VALUES (1, 'Always respond in Estonian language. Be professional and concise.') ON CONFLICT (id) DO UPDATE SET prompt = EXCLUDED.prompt; ``` @@ -260,7 +303,7 @@ curl -X POST http://localhost:8100/orchestrate/test \ ### **3. Update Prompt** ```sql -UPDATE public.prompt_configuration +UPDATE rag_search.prompt_configuration SET prompt = 'Provide concise answers using bullet points. Be helpful and clear.' WHERE id = 1; ``` @@ -290,14 +333,14 @@ curl -X POST http://localhost:8100/prompt-config/refresh ## Key Features -✅ **TTL Caching** - 5-minute cache reduces database calls -✅ **Immediate Updates** - Admin changes trigger instant refresh -✅ **Graceful Degradation** - If refresh fails, TTL cache continues working -✅ **Thread-Safe** - Multiple concurrent requests handled safely -✅ **Retry Logic** - 3 attempts with exponential backoff for HTTP failures -✅ **Instruction Prepending** - Preserves DSPy optimization compatibility -✅ **Applied Consistently** - Works across all 3 orchestration endpoints -✅ **Applied to ResponseGenerator Only** - Not applied to PromptRefinerAgent + **TTL Caching** - 5-minute cache reduces database calls + **Immediate Updates** - Admin changes trigger instant refresh + **Graceful Degradation** - If refresh fails, TTL cache continues working + **Thread-Safe** - Multiple concurrent requests handled safely + **Retry Logic** - 3 attempts with exponential backoff for HTTP failures + **Instruction Appending** - Custom instructions appended to the question without modifying the DSPy signature + **Applied Consistently** - Works across all 3 orchestration endpoints + **Applied to RAG & API Tool Calling** - In RAG, only the ResponseGenerator (not the PromptRefiner); in API Tool Calling, the param extractor and response formatters. Not applied to the Context workflow. --- @@ -322,10 +365,10 @@ Context: [retrieved documentation chunks...] ``` **Expected Response:** -- In Estonian language ✅ -- Professional tone ✅ -- Concise format ✅ -- Citations included ✅ +- In Estonian language +- Professional tone +- Concise format +- Citations included --- @@ -348,7 +391,7 @@ Context: [retrieved documentation chunks...] ### **Prompt Not Applied** - Check logs for: "Custom prompt configuration loaded at startup" -- Verify database has prompt: `SELECT * FROM public.prompt_configuration;` +- Verify database has prompt: `SELECT * FROM rag_search.prompt_configuration;` - Test refresh endpoint: `curl -X POST http://localhost:8100/prompt-config/refresh` ### **Cache Not Refreshing** @@ -365,7 +408,9 @@ Context: [retrieved documentation chunks...] ## Notes -- Custom prompts apply **only to ResponseGeneratorAgent** (not PromptRefinerAgent) -- PromptRefiner focuses on query optimization for retrieval -- ResponseGenerator needs language policy and interaction style for user-facing content -- This design preserves DSPy optimization compatibility by using instruction prepending instead of signature modification +- In the **RAG** path, custom prompts apply **only to `ResponseGeneratorAgent`** (not `PromptRefinerAgent`). + The wrapped instructions are appended to the question rather than modifying the DSPy signature. +- The **API Tool Calling** workflow consumes the **same** configured prompt (via the shared loader) and + injects the raw text as a dedicated `custom_instructions` DSPy input field on the parameter extractor + and the response formatters, where it is followed with highest priority. +- The **Context** and **Service** workflows do not apply custom prompts (see the per-workflow section). diff --git a/docs/HYBRID_SEARCH_CLASSIFICATION.md b/docs/HYBRID_SEARCH_CLASSIFICATION.md index 1de3f7f5..a4521e24 100644 --- a/docs/HYBRID_SEARCH_CLASSIFICATION.md +++ b/docs/HYBRID_SEARCH_CLASSIFICATION.md @@ -115,9 +115,7 @@ tokens = re.findall(r"\w+", text.lower()) # ["mis", "suhe", "on", "euro", ...] ```python # Collection: "intent_collections" -vectors_config = { - "dense": VectorParams(size=3072, distance=Distance.COSINE) -} +vectors_config = {"dense": VectorParams(size=3072, distance=Distance.COSINE)} sparse_vectors_config = { "sparse": SparseVectorParams(index=SparseIndexParams(on_disk=False)) } @@ -197,12 +195,12 @@ Queries Qdrant using only the dense vector to get **actual cosine similarity sco ```python # classifier.py → _dense_search() -POST /collections/intent_collections/points/query +POST / collections / intent_collections / points / query { "query": [0.023, -0.041, ...], # 3072-dim dense vector "using": "dense", - "limit": 6, # DENSE_SEARCH_TOP_K * 2 (3 * 2 = 6, allows dedup) - "with_payload": true + "limit": 6, # DENSE_SEARCH_TOP_K * 2 (3 * 2 = 6, allows dedup) + "with_payload": true, } ``` @@ -219,15 +217,15 @@ Sparse prefetch is only included if the query produces a non-empty sparse vector ```python # classifier.py → _hybrid_search() # First checks collection exists and has data (points_count > 0) -POST /collections/intent_collections/points/query +POST / collections / intent_collections / points / query { "prefetch": [ {"query": dense_vector, "using": "dense", "limit": 10}, - {"query": {"indices": [...], "values": [...]}, "using": "sparse", "limit": 10} + {"query": {"indices": [...], "values": [...]}, "using": "sparse", "limit": 10}, ], "query": {"fusion": "rrf"}, "limit": 5, - "with_payload": true + "with_payload": true, } ``` diff --git a/docs/LLM_CONFIG_VAULT_INTEGRATION.md b/docs/LLM_CONFIG_VAULT_INTEGRATION.md deleted file mode 100644 index 563054ee..00000000 --- a/docs/LLM_CONFIG_VAULT_INTEGRATION.md +++ /dev/null @@ -1,546 +0,0 @@ -# LLM Config Module - HashiCorp Vault Integration - -## Overview - -The LLM Config Module integrates with HashiCorp Vault to securely store and manage API keys, endpoints, and other sensitive configuration data for various LLM providers (AWS Bedrock, Azure OpenAI, etc.). This integration replaces the traditional `.env` file approach with a more secure, centralized secret management system. - -## Architecture - -### Components - -1. **VaultSecretResolver** - Core component that interfaces with Vault -2. **ConfigurationLoader** - Loads configuration and resolves secrets from Vault -3. **LLMManager** - Main entry point that initializes with Vault-backed configuration -4. **Connection Management** - Dynamic discovery of provider connections from Vault - -### Key Features - -- **Environment-Aware**: Automatically discovers and uses appropriate secrets based on environment (production/development/test) -- **User-Independent**: No hardcoded user lists - dynamically discovers available connections -- **Provider Discovery**: Automatically detects which LLM providers are available based on Vault contents -- **Fallback Protection**: Graceful handling when Vault is unavailable (fails securely) - -## Vault Data Structure - -### Secret Storage Schema - -The Vault integration uses the KV v2 secrets engine with the following hierarchical structure: - -``` -secret/ -├── users/ -│ ├── user1/ -│ │ ├── conn_12345abc/ -│ │ │ ├── data/ -│ │ │ │ ├── provider: "aws_bedrock" -│ │ │ │ ├── environment: "production" -│ │ │ │ ├── aws_access_key_id: "AKIA..." -│ │ │ │ ├── aws_secret_access_key: "..." -│ │ │ │ ├── aws_region: "us-east-1" -│ │ │ │ └── model_id: "anthropic.claude-3-sonnet-20240229-v1:0" -│ │ └── conn_67890def/ -│ │ ├── data/ -│ │ │ ├── provider: "azure_openai" -│ │ │ ├── environment: "development" -│ │ │ ├── api_key: "sk-..." -│ │ │ ├── endpoint: "https://myservice.openai.azure.com/" -│ │ │ ├── deployment_name: "gpt-4" -│ │ │ └── api_version: "2024-02-15-preview" -│ └── user2/ -│ └── conn_11111xyz/ -│ └── data/ -│ ├── provider: "aws_bedrock" -│ ├── environment: "production" -│ └── ... -``` - -### Connection Metadata - -Each connection contains: - -- **Provider Type**: `aws_bedrock`, `azure_openai`, etc. -- **Environment**: `production`, `development`, `test` -- **Provider-specific secrets**: API keys, endpoints, regions, model IDs -- **Connection ID**: Unique identifier for the connection - -## Development Container Setup - -### Current Container Configuration - -The project includes a development Vault container configured in `docker-compose.yml`: - -```yaml -vault: - image: hashicorp/vault:latest - container_name: vault - command: ["vault", "server", "-dev", "-dev-listen-address=0.0.0.0:8200", "-dev-root-token-id=myroot"] - cap_add: - - IPC_LOCK - ports: - - "8200:8200" - environment: - - VAULT_ADDR=http://0.0.0.0:8200 - - VAULT_API_ADDR=http://localhost:8200 - - VAULT_DEV_ROOT_TOKEN_ID=myroot - - VAULT_DEV_LISTEN_ADDRESS=0.0.0.0:8200 - volumes: - - vault-data:/vault/data - networks: - - bykstack - restart: unless-stopped - healthcheck: - test: ["CMD", "vault", "status"] - interval: 10s - timeout: 5s - retries: 5 -``` - -### Starting the Development Environment - -1. **Start Vault Container**: - ```bash - docker-compose up vault -d - ``` - -2. **Verify Vault is Running**: - ```bash - curl http://localhost:8200/v1/sys/health - ``` - -3. **Access Vault UI**: - - URL: http://localhost:8200 - - Token: `myroot` - -### Development Configuration - -For development, set these environment variables: - -```bash -export VAULT_ADDR="http://localhost:8200" -export VAULT_TOKEN="myroot" -``` - -## Usage Examples - -### Production Environment - -```python -import os -from llm_config_module import LLMManager - -# Set Vault connection details -os.environ["VAULT_ADDR"] = "https://vault.company.com" -os.environ["VAULT_TOKEN"] = "your-production-token" - -# Initialize LLM Manager - automatically discovers production providers -manager = LLMManager(environment="production") - -# Get available providers (discovered from Vault) -providers = manager.get_available_providers() -print(f"Available providers: {list(providers.keys())}") - -# Use the LLM -llm = manager.get_llm() -response = llm.generate("Hello, world!") -``` - -### Development Environment - -```python -# Development requires a specific connection ID -manager = LLMManager( - environment="development", - connection_id="conn_12345abc" # Specific dev connection -) - -llm = manager.get_llm() -``` - -### Dynamic Provider Discovery - -The system automatically discovers which providers are available: - -```python -manager = LLMManager(environment="production") - -# Only providers with valid Vault secrets will be available -if manager.is_provider_available(LLMProvider.AWS_BEDROCK): - print("AWS Bedrock is configured and available") - -if manager.is_provider_available(LLMProvider.AZURE_OPENAI): - print("Azure OpenAI is configured and available") -``` - -## Configuration Details - -### Vault Configuration (llm_config.yaml) - -```yaml -vault: - enabled: true - url: "${VAULT_ADDR}" - token: "${VAULT_TOKEN}" - mount_point: "secret" - secrets_engine: "kv-v2" - -providers: - aws_bedrock: - enabled: true # Will be dynamically determined from Vault - model_id: "anthropic.claude-3-sonnet-20240229-v1:0" - max_tokens: 1000 - temperature: 0.7 - - azure_openai: - enabled: true # Will be dynamically determined from Vault - max_tokens: 1000 - temperature: 0.7 -``` - -### Environment Variable Resolution - -The configuration supports environment variable substitution: - -- `${VAULT_ADDR}` - Vault server URL -- `${VAULT_TOKEN}` - Vault authentication token - -## Production Considerations - -### Security Best Practices - -#### 1. Authentication & Authorization - -**🔒 Token Management**: -```bash -# Use short-lived tokens in production -vault write auth/userpass/users/llm-service password="secure-password" policies="llm-read-policy" - -# Generate service token -vault write -field=token auth/userpass/login/llm-service password="secure-password" -``` - -**🔒 Policy Configuration**: -```hcl -# llm-read-policy.hcl -path "secret/data/users/*/conn_*" { - capabilities = ["read"] -} - -path "secret/metadata/users/*" { - capabilities = ["list", "read"] -} -``` - -#### 2. Network Security - -**🔒 TLS Configuration**: -```hcl -# vault.hcl (Production) -listener "tcp" { - address = "0.0.0.0:8200" - tls_cert_file = "/etc/ssl/vault/vault.crt" - tls_key_file = "/etc/ssl/vault/vault.key" - tls_min_version = "tls12" -} -``` - -**🔒 Network Isolation**: -- Deploy Vault in private subnets -- Use VPC endpoints for AWS services -- Implement network ACLs and security groups -- Enable Vault audit logging - -#### 3. High Availability Setup - -**🏗️ Raft Storage Backend**: -```hcl -storage "raft" { - path = "/vault/data" - node_id = "vault-1" - - retry_join { - leader_api_addr = "https://vault-1.internal:8200" - } - retry_join { - leader_api_addr = "https://vault-2.internal:8200" - } - retry_join { - leader_api_addr = "https://vault-3.internal:8200" - } -} -``` - -**🏗️ Auto-Unseal** (recommended): -```hcl -seal "awskms" { - region = "us-east-1" - kms_key_id = "alias/vault-unseal-key" -} -``` - -#### 4. Monitoring & Logging - -**📊 Health Checks**: -```yaml -# kubernetes health check -livenessProbe: - httpGet: - path: /v1/sys/health - port: 8200 - scheme: HTTPS - initialDelaySeconds: 60 - timeoutSeconds: 5 -``` - -**📊 Audit Logging**: -```hcl -audit "file" { - file_path = "/vault/logs/audit.log" -} -``` - -#### 5. Backup & Recovery - -**💾 Automated Snapshots**: -```bash -#!/bin/bash -# backup-vault.sh -vault operator raft snapshot save "vault-snapshot-$(date +%Y%m%d-%H%M%S).snap" -aws s3 cp "vault-snapshot-*.snap" s3://vault-backups/ -``` - -### Production Deployment Architecture - -```mermaid -graph TB - subgraph "Load Balancer" - ALB[Application Load Balancer] - end - - subgraph "Vault Cluster" - V1[Vault Node 1
Active] - V2[Vault Node 2
Standby] - V3[Vault Node 3
Standby] - end - - subgraph "Application Tier" - APP1[LLM App 1] - APP2[LLM App 2] - APP3[LLM App 3] - end - - subgraph "External Services" - AWS[AWS Bedrock] - AZURE[Azure OpenAI] - end - - ALB --> V1 - ALB --> V2 - ALB --> V3 - - APP1 --> ALB - APP2 --> ALB - APP3 --> ALB - - APP1 --> AWS - APP2 --> AZURE - APP3 --> AWS -``` - -### Environment-Specific Configurations - -#### Production -```yaml -# Production values -vault: - url: "https://vault.company.com" - token: "${VAULT_SERVICE_TOKEN}" # From secure secret management - -# Use IAM roles where possible -providers: - aws_bedrock: - use_iam_role: true # Preferred over access keys -``` - -#### Staging -```yaml -vault: - url: "https://vault-staging.company.com" - token: "${VAULT_STAGING_TOKEN}" -``` - -#### Development -```yaml -vault: - url: "http://localhost:8200" - token: "myroot" # Development only -``` - -## Migration from .env Files - -### Step-by-Step Migration - -1. **Identify Current Secrets**: - ```bash - # List current .env variables - grep -E "(API_KEY|SECRET|TOKEN)" .env - ``` - -2. **Create Vault Connections**: - ```bash - # Example: Migrate AWS credentials - vault kv put secret/users/production/conn_aws_prod \ - provider="aws_bedrock" \ - environment="production" \ - aws_access_key_id="$AWS_ACCESS_KEY_ID" \ - aws_secret_access_key="$AWS_SECRET_ACCESS_KEY" \ - aws_region="us-east-1" \ - model_id="anthropic.claude-3-sonnet-20240229-v1:0" - ``` - -3. **Update Application Code**: - ```python - # Before (using .env) - manager = LLMManager(config_path="config.yaml", environment="production") - - # After (using Vault) - manager = LLMManager(environment="production") # Auto-discovers from Vault - ``` - -4. **Verify Migration**: - ```python - # Test that providers are discovered correctly - providers = manager.get_available_providers() - assert len(providers) > 0, "No providers discovered from Vault" - ``` - -## Testing - -### Unit Tests - -The integration includes comprehensive test coverage: - -- **Vault Integration Tests**: `test_integration_vault_llm_config.py` -- **Provider-Specific Tests**: `test_aws.py`, `test_azure.py` -- **Helper Functions**: `vault_test_helpers.py` - -### Running Tests - -```bash -# Run all tests -uv run pytest -v - -# Run only Vault integration tests -uv run pytest tests/test_integration_vault_llm_config.py -v - -# Run provider-specific tests -uv run pytest tests/test_aws.py tests/test_azure.py -v -``` - -### Test Helpers - -The `vault_test_helpers.py` provides utilities for test discovery: - -```python -from tests.vault_test_helpers import ( - check_vault_available, - get_available_providers_from_vault, - should_skip_aws_test, - should_skip_azure_test -) - -# Conditionally skip tests based on Vault provider availability -@pytest.mark.skipif(should_skip_aws_test(), reason="AWS not available in Vault") -def test_aws_integration(): - # Test will only run if AWS Bedrock is configured in Vault - pass -``` - -## Troubleshooting - -### Common Issues - -#### 1. Vault Connection Failures -```python -# Check Vault connectivity -try: - from rag_config_manager.vault import VaultClient - vault = VaultClient() - print(f"Vault available: {vault.is_vault_available()}") -except Exception as e: - print(f"Vault error: {e}") -``` - -#### 2. Provider Discovery Issues -```python -# Debug provider discovery -import os -os.environ["VAULT_ADDR"] = "http://localhost:8200" -os.environ["VAULT_TOKEN"] = "myroot" - -manager = LLMManager(environment="production") -providers = manager.get_available_providers() -print(f"Discovered providers: {list(providers.keys())}") -``` - -#### 3. Authentication Errors -- Verify `VAULT_TOKEN` is valid and not expired -- Check token policies have required permissions -- Ensure Vault server is accessible from application network - -#### 4. Secret Path Issues -- Verify secret paths match the expected structure -- Check that secrets exist in the correct mount point -- Ensure proper KV v2 format is used - -### Logging - -Enable debug logging to troubleshoot issues: - -```python -import logging -logging.basicConfig(level=logging.DEBUG) - -# The LLM Config Module uses loguru for logging -from loguru import logger -logger.add("vault_debug.log", level="DEBUG") -``` - -## Best Practices Summary - -### ✅ Do: -- Use production-grade Vault deployment with HA -- Implement proper authentication (avoid root tokens) -- Enable TLS in production -- Use auto-unseal mechanisms -- Implement comprehensive monitoring -- Regular backup and recovery testing -- Use IAM roles where possible instead of static keys -- Rotate secrets regularly - -### ❌ Don't: -- Use development mode Vault in production -- Store root tokens in application code -- Disable TLS in production environments -- Skip audit logging -- Use overly permissive policies -- Store Vault tokens in environment files -- Forget to implement proper secret rotation - -## Support & Maintenance - -### Vault Version Compatibility -- **Minimum**: Vault 1.12+ -- **Recommended**: Vault 1.15+ -- **Tested With**: Vault 1.15.1 - -### Dependencies -- `rag_config_manager` - Vault client interface -- `hvac` - HashiCorp Vault client library -- `pydantic` - Data validation and settings management - -### Monitoring Endpoints -- Health: `GET /v1/sys/health` -- Metrics: `GET /v1/sys/metrics` (Prometheus format) -- Status: `vault status` (CLI command) - -This integration provides a robust, secure, and scalable approach to managing LLM provider secrets using HashiCorp Vault, replacing traditional environment variable-based configuration with enterprise-grade secret management. diff --git a/docs/README.md b/docs/README.md new file mode 100644 index 00000000..32115576 --- /dev/null +++ b/docs/README.md @@ -0,0 +1,64 @@ +# Documentation Index + +Welcome to the documentation for the **LLM Module** — the LLM orchestration component of the +Bürokratt / Estonian Government AI assistant. This index catalogues every document under `docs/`, +grouped by topic. Start with the [project README](../README.md) for the big picture, then dive into +the areas below. + +> New to the system? Read [ARCHITECTURE.md](./ARCHITECTURE.md) first — it walks the C4 diagrams from +> system context down to the internal components. + +--- + +## Architecture & Design + +| Document | What it covers | +| --- | --- | +| [ARCHITECTURE.md](./ARCHITECTURE.md) | C4 model walkthrough (context → containers → components) with links into every detailed flow. The recommended starting point. | + +## Retrieval & Search + +| Document | What it covers | +| --- | --- | +| [CONTEXTUAL_RETRIEVAL_FLOW.md](./CONTEXTUAL_RETRIEVAL_FLOW.md) | The RAG workflow in depth: multi-query expansion, hybrid (semantic + BM25) search, RRF rank fusion, thresholds, and quality testing. | +| [HYBRID_SEARCH_CLASSIFICATION.md](./HYBRID_SEARCH_CLASSIFICATION.md) | Tool-classifier architecture using per-example dense (3072-dim) + sparse (BM25) vectors in Qdrant; offline indexing and query-time classification. | + +## Tool Classification & Workflows + +| Document | What it covers | +| --- | --- | +| [TOOL_CLASSIFIER.md](./TOOL_CLASSIFIER.md) | High-level overview of the classifier: routing model, key functions, and configuration. Start here, then read the per-workflow docs below. | +| [TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md](./TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md) | Service-workflow architecture: high-confidence vs ambiguous routes and service-discovery logic. | +| [CONTEXT_WORKFLOW_GREETING_DETECTION.md](./CONTEXT_WORKFLOW_GREETING_DETECTION.md) | The Context workflow: greeting detection and Redis-backed conversation history with incremental summaries. | +| [API_TOOL_CALLING.md](./API_TOOL_CALLING.md) | The agentic API-Tool workflow end to end: indexing, multi-intent decomposition, multi-endpoint agentic loop, calling, and response formatting. | +| [TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md](./TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md) | Service workflow + streaming via the TestProductionLLM page (three-hop SSE relay). | + +## Configuration & Secrets + +| Document | What it covers | +| --- | --- | +| [LLM_CONFIG_VAULT_INTEGRATION.md](./LLM_CONFIG_VAULT_INTEGRATION.md) | HashiCorp Vault integration for LLM credentials: KV v2 layout, dev setup, production HA, and migration from `.env`. | +| [VAULT_SETUP_AND_USAGE.md](./VAULT_SETUP_AND_USAGE.md) | Operational guide: dual-network topology, vault-agents, bootstrap flow, AppRole auth, and credential reconciliation. | +| [VAULT_SECURITY_ARCHITECTURE.md](./VAULT_SECURITY_ARCHITECTURE.md) | Security model: threat model, network isolation, AppRole authentication, and the per-policy access-control matrix. | +| [CONNECTION_SWAP_FLOW.md](./CONNECTION_SWAP_FLOW.md) | UUID-based Vault path design enabling zero-I/O environment swaps (promote/demote) between LLM connections. | +| [CUSTOM_PROMPT_CONFIGURATION.md](./CUSTOM_PROMPT_CONFIGURATION.md) | Admin-facing prompt management: database → Ruuter DSL → Python loader with TTL cache and invalidation. | + +## Data & Sessions + +| Document | What it covers | +| --- | --- | +| [REDIS_SESSION_STORE.md](./REDIS_SESSION_STORE.md) | Redis session-store usage: CRUD for agentic-loop state, TTL behaviour, and the async API. | + +## API Reference + +| Document | What it covers | +| --- | --- | +| [API_REFERENCE.md](./API_REFERENCE.md) | HTTP API reference: LLM Connections management, Inference Results storage/retrieval, and the chatbot Inquiry endpoint. | + +--- + +## Related resources + +- [Project README](../README.md) — overview, quick start, and component map. +- [CONTRIBUTING.md](../CONTRIBUTING.md) — development environment, tooling, and CI checks. +- Architecture diagrams (source images): [`images/`](./images). diff --git a/docs/REDIS_SESSION_STORE.md b/docs/REDIS_SESSION_STORE.md index 502f24fe..00082425 100644 --- a/docs/REDIS_SESSION_STORE.md +++ b/docs/REDIS_SESSION_STORE.md @@ -65,9 +65,9 @@ if session is None: # No active session → this is a fresh conversation ... else: - print(session.state) # "collecting_params" - print(session.collected_params) # {"city": "Tallinn"} - print(session.turn_count) # 2 + print(session.state) # "collecting_params" + print(session.collected_params) # {"city": "Tallinn"} + print(session.turn_count) # 2 ``` --- @@ -129,12 +129,14 @@ session_store = request.app.state.session_store session = await session_store.get(request.chatId) if session is None: detected_endpoint = ... # endpoint detected from user query - await session_store.save(APIToolSession( - chat_id=request.chatId, - state="collecting_params", - selected_endpoint=detected_endpoint, - turn_count=1, - )) + await session_store.save( + APIToolSession( + chat_id=request.chatId, + state="collecting_params", + selected_endpoint=detected_endpoint, + turn_count=1, + ) + ) return "Which city would you like weather for?" # --- Turn 2+ --- diff --git a/docs/TESTMODEL_SERVICE_WORKFLOW.md b/docs/TESTMODEL_SERVICE_WORKFLOW.md deleted file mode 100644 index e1649a9f..00000000 --- a/docs/TESTMODEL_SERVICE_WORKFLOW.md +++ /dev/null @@ -1,405 +0,0 @@ -# TestModel Page — Service Workflow Documentation - -This document traces the **service workflow** end-to-end through the **TestModel** UI page. It covers two scenarios: - -1. **Natural-language service detection** — the user types a free-text query that the system classifies as a service. -2. **MCQ button-click** — the user clicks a choice button whose payload is a `#service` command, short-circuiting the NLU pipeline. - ---- - -## Architecture Overview - -``` -┌─────────────┐ POST /rag-search/inference/test ┌─────────────────┐ -│ TestModel │ ──────────────────────────────────────────▷ │ Ruuter (proxy) │ -│ (GUI) │ │ /rag-search/ │ -│ index.tsx │ ◁─────────── JSON response ──────────────── │ inference/test │ -└─────────────┘ └───────┬─────────┘ - │ - POST /orchestrate/test - ▼ - ┌───────────────────────┐ - │ llm_orchestration_ │ - │ service_api.py │ - │ test_orchestrate_ │ - │ llm_request() │ - └───────┬───────────────┘ - │ - OrchestrationRequest (mapped with defaults) - ▼ - ┌───────────────────────┐ - │ llm_orchestration_ │ - │ service.py │ - │ process_orchestration_ │ - │ request() │ - └───────────────────────┘ -``` - ---- - -## Key Files - -| Layer | File | Purpose | -|---|---|---| -| **GUI** | `GUI/src/pages/TestModel/index.tsx` | UI page with connection selector, text input, result display, MCQ buttons | -| **GUI Service** | `GUI/src/services/inference.ts` | `viewInferenceResult()` — POST to `/rag-search/inference/test` | -| **API Layer** | `src/llm_orchestration_service_api.py` | `/orchestrate/test` handler — maps `TestOrchestrationRequest` → `OrchestrationRequest` and calls `process_orchestration_request()` | -| **Orchestration** | `src/llm_orchestration_service.py` | `process_orchestration_request()` — the core pipeline (language detection → `#service` prefix check → query validation → guardrails → classifier → service workflow) | -| **Service Workflow** | `src/tool_classifier/workflows/service_workflow.py` | `ServiceWorkflowExecutor` — service discovery, intent detection, entity extraction, endpoint call, direct step execution | -| **Models** | `src/models/request_models.py` | `OrchestrationRequest`, `OrchestrationResponse`, `TestOrchestrationRequest`, `TestOrchestrationResponse`, `ChoiceButton` | -| **Constants** | `src/tool_classifier/constants.py` | `SERVICE_STEP_PREFIXES`, `RUUTER_SERVICE_BASE_URL`, search thresholds | - ---- - -## Flow 1: Natural-Language Service Detection - -### 1.1 Frontend — User Sends a Message - -The user selects an LLM connection from the dropdown and types a query (e.g., *"My keyboard is not working"*). - -**`TestModel/index.tsx` → `handleSend()`** (line 72): -```tsx -inferenceMutation.mutate({ - llmConnectionId: Number(testLLM.connectionId), - message: testLLM.text, -}); -``` - -**`inference.ts` → `viewInferenceResult()`** (line 50): -```ts -const { data } = await apiDev.post(inferenceEndpoints.VIEW_TEST_INFERENCE_RESULT(), { - connectionId: request.llmConnectionId, - message: request.message, -}); -``` - -This POST goes to `/rag-search/inference/test` (via Ruuter proxy), which maps to the backend `/orchestrate/test` endpoint. - -### 1.2 API Layer — Request Mapping - -**`llm_orchestration_service_api.py` → `test_orchestrate_llm_request()`** (line 313): -- Receives a `TestOrchestrationRequest` (only `message`, `environment`, optional `connectionId`). -- Maps to a full `OrchestrationRequest` with defaults: - -```python -full_request = OrchestrationRequest( - chatId="test-session", - message=request.message, - authorId="test-user", - conversationHistory=[], - url="test-context", - environment=request.environment, - connection_id=str(request.connectionId) if request.connectionId is not None else None, -) -``` - -- Calls `orchestration_service.process_orchestration_request(full_request)`. - -### 1.3 Orchestration Pipeline — Service Detection - -**`llm_orchestration_service.py` → `process_orchestration_request()`** (line 279): - -``` -STEP 0: Language detection → detect_language(request.message) → "en" - -STEP 0.1: Check request.message.startswith(SERVICE_STEP_PREFIXES) - → FALSE (natural language) → skip, continue normally - -STEP 0.5: Query validation → validate_query_basic(request.message) → valid - -STEP 1: Component initialization → LLM manager, guardrails adapter - -STEP 2: Input guardrails check → allowed - -STEP 3: ToolClassifier.classify(query, conversation_history, language) - → Classification(workflow=SERVICE, confidence=0.92) - The classifier uses hybrid search (dense + BM25) against - the intent_collections Qdrant collection. - -STEP 4: route_to_workflow(classification, request, is_streaming=False) - → Routes to ServiceWorkflowExecutor.execute_async() -``` - -### 1.4 ServiceWorkflowExecutor — `execute_async()` - -**`service_workflow.py` → `execute_async()`** (line 699): - -The executor uses **classification metadata** from hybrid search to decide how to proceed. There are three paths: - -| Condition | Path | -|---|---| -| `needs_llm_confirmation == False` | High-confidence match — run intent detection on the single top match only | -| `needs_llm_confirmation == True` | Ambiguous — run intent detection on top-N candidates | -| No metadata | Fall back to full discovery flow (`_log_request_details`) | - -#### Service Discovery (full flow) - -1. **`_call_service_discovery(chat_id)`** — calls `GET http://ruuter-public:8086/rag-search/services/get-services` -2. Checks if `service_count > SERVICE_COUNT_THRESHOLD (10)` → triggers **semantic search** via Qdrant -3. Otherwise uses the services list directly - -#### Intent Detection - -**`_process_intent_detection(services, request, chat_id, context, costs_metric)`**: -1. Calls **`_detect_service_intent()`** → uses `IntentDetectionModule` (DSPy LLM) with: - - The user query - - The candidate services list - - Conversation history -2. Returns matched `service_id`, `confidence`, `entities` -3. **`_validate_detected_service()`** — confirms the matched service exists in the active services list - -#### Entity Extraction & Validation - -```python -service_metadata = self._extract_service_metadata(context, chat_id) -# → {service_id, service_name, entities_dict, entity_schema, ruuter_type, is_common} - -validation_result = self._validate_entities(entities_dict, entity_schema, service_name, chat_id) -# → checks missing, extra, empty entities - -entities_array = self._transform_entities_to_array(entities_dict, entity_schema) -# → ordered list of entity values matching service schema -``` - -#### Service Endpoint Call - -```python -endpoint_url = self._construct_service_endpoint(service_name, chat_id, is_common) -# → "http://ruuter:8086/services/services/active/Klaviatuuri_probleemi_lahendamine" - -service_result = await self._call_service_endpoint( - endpoint_url, http_method, entities_array, chat_id, author_id -) -``` - -**`_call_service_endpoint()`** (line 523): -1. Sends POST/GET to the Ruuter endpoint with payload `{chatId, authorId, input: entities_array}` -2. Ruuter executes the DSL → DMapper produces the response -3. Parses the response: - - Unwraps `{"response": ...}` wrapper - - Extracts `data[0].content` → text content - - Extracts `data[0].buttons` → JSON string or list of `{title, payload}` objects -4. Returns `{"content": str, "buttons": List[Dict]}` - -#### Build Response - -```python -service_buttons = service_result["buttons"] -buttons_list = [ChoiceButton(**b) for b in service_buttons if "title" in b and "payload" in b] - -return OrchestrationResponse( - chatId=request.chatId, - llmServiceActive=True, - questionOutOfLLMScope=False, - inputGuardFailed=False, - content=service_content, - buttons=buttons_list if buttons_list else None, -) -``` - -### 1.5 API Layer — Response Conversion - -Back in `test_orchestrate_llm_request()` (line 382): - -```python -test_response = TestOrchestrationResponse( - llmServiceActive=response.llmServiceActive, - questionOutOfLLMScope=response.questionOutOfLLMScope, - inputGuardFailed=response.inputGuardFailed, - content=response.content, - buttons=response.buttons, # ← forwarded - chunks=None, -) -``` - -### 1.6 Frontend — Display Result - -**`TestModel/index.tsx`**: - -```tsx -// onSuccess callback (line 51-54) -setInferenceResult(data?.response); -setPendingButtons(data?.response?.buttons ?? []); -``` - -- Response text is rendered inside `` (line 166-168) -- MCQ buttons are rendered if `pendingButtons.length > 0` (line 173-186): - -```tsx -{pendingButtons.map((btn) => ( - -))} -``` - ---- - -## Flow 2: MCQ Button Click (Direct Step — Short-Circuit) - -### 2.1 Frontend — Button Click - -When the user clicks a button (e.g., *[Windows]*), `handleButtonClick()` is called (line 81): - -```tsx -const handleButtonClick = (payload: string) => { - if (!testLLM.connectionId) return; - setPendingButtons([]); - inferenceMutation.mutate({ - llmConnectionId: Number(testLLM.connectionId), - message: payload, // e.g. "#service, /POST/services/active/Klaviatuuri_probleemi_lahendamine_mcq_1_0" - }); -}; -``` - -The button **payload** becomes the next message. The same API call (`/rag-search/inference/test` → `/orchestrate/test`) is made. - -### 2.2 Input Sanitizer Safety - -The `#service, /POST/...` payload goes through Pydantic's `validate_and_sanitize_message()` on `OrchestrationRequest.message` (line 64-88 of `request_models.py`). The `InputSanitizer.sanitize_message()` strips HTML tags and normalizes whitespace but leaves `#`, `,`, `/` characters intact. The payload passes through unchanged. - -### 2.3 Orchestration — `#service` Prefix Short-Circuit - -**`process_orchestration_request()`** (line 324-340): - -```python -# STEP 0.1: Multi-step service prefix check (bypass NLU pipeline) -if request.message.startswith(SERVICE_STEP_PREFIXES): - logger.info(f"[{request.chatId}] #service prefix detected - direct step execution") - executor = self._get_service_workflow_executor() - direct_response = await executor.execute_direct_step( - request=request, - time_metric=time_metric, - ) - if direct_response is not None: - log_step_timings(time_metric, request.chatId) - return direct_response -``` - -**`SERVICE_STEP_PREFIXES`** = `("#service,", "#common_service,")` from `constants.py`. - -**`_get_service_workflow_executor()`** (line 264): -- Reuses the existing `tool_classifier.service_workflow` if a ToolClassifier has been initialized -- Otherwise creates a lightweight `ServiceWorkflowExecutor(llm_manager=None, orchestration_service=self)` — no LLM is needed for direct steps - -### 2.4 ServiceWorkflowExecutor — `execute_direct_step()` - -**`service_workflow.py` → `execute_direct_step()`** (line 1000): - -1. **Parse**: `_parse_service_prefix(request.message)` - - Input: `"#service, /POST/services/active/Klaviatuuri_probleemi_lahendamine_mcq_1_0"` - - Splits off prefix → remainder: `/POST/services/active/...` - - Extracts HTTP method: `POST` - - Builds URL: `http://ruuter:8086/services/services/active/Klaviatuuri_probleemi_lahendamine_mcq_1_0` - - Returns: `("POST", "http://ruuter:8086/services/services/active/...")` - -2. **Call endpoint**: `_call_service_endpoint(url, "POST", [], chat_id, author_id)` - - `entities_array=[]` — no entities for MCQ steps - - Same parsing as Flow 1 (extracts `content` + `buttons`) - -3. **Build response**: Same `OrchestrationResponse` construction as Flow 1 - -### What Gets Skipped (Short-Circuit) - -| Skipped Step | Why | -|---|---| -| Query validation | Would reject `#service` as gibberish | -| Component initialization | Expensive (LLM manager, Vault, guardrails) | -| Input guardrails | Would block a machine-generated payload | -| ToolClassifier.classify() | LLM call — unnecessary cost | -| Intent detection LLM | Another LLM call — URL is already known | -| Entity extraction | No natural language entities to extract | -| Semantic search (Qdrant) | No need to find a service | - -### 2.5 Response & Loop - -The response follows the same path back through `test_orchestrate_llm_request()` → frontend. - -- If the response has `buttons` → frontend renders the next set of MCQ buttons -- If the response has `buttons=null` → the MCQ flow is complete, only the final text answer is shown - ---- - -## Data Models - -### Request - -```python -class TestOrchestrationRequest: - message: str - environment: Literal["production", "testing", "development"] - connectionId: Optional[int] -``` - -### Response - -```python -class TestOrchestrationResponse: - llmServiceActive: bool - questionOutOfLLMScope: bool - inputGuardFailed: bool - content: str - buttons: Optional[List[ChoiceButton]] # MCQ buttons - chunks: Optional[List[ChunkInfo]] # RAG context chunks - -class ChoiceButton: - title: str # "Windows" - payload: str # "#service, /POST/services/active/..." -``` - ---- - -## Complete MCQ Sequence Diagram - -``` -User TestModel UI API (/orchestrate/test) ServiceWorkflow Ruuter/DMapper - │ │ │ │ │ - │ types "keyboard │ │ │ │ - │ not working" │ │ │ │ - ├───────────────────▶│ POST inference/test │ │ │ - │ ├────────────────────────▶│ process_orchestration_ │ │ - │ │ │ request() │ │ - │ │ │ startswith(#service)?→NO │ │ - │ │ │ classify()→SERVICE │ │ - │ │ ├───────────────────────────▶│ execute_async() │ - │ │ │ │ intent detect→matched │ - │ │ │ │ _call_service_endpoint │ - │ │ │ ├──────────────────────▶│ - │ │ │ │◁ {content,buttons}─────│ - │ │ │◁─ OrchestrationResponse ───│ │ - │ │◁─ TestOrchResponse ─────│ │ │ - │◁─ render text + │ │ │ │ - │ [Windows] [Mac] │ │ │ │ - │ │ │ │ │ - │ clicks [Windows] │ │ │ │ - ├───────────────────▶│ POST inference/test │ │ │ - │ │ msg="#service,/POST/…" │ │ │ - │ ├────────────────────────▶│ startswith(#service)?→YES │ │ - │ │ │ SKIP classifier+guardrails│ │ - │ │ ├───────────────────────────▶│ execute_direct_step() │ - │ │ │ │ _parse_service_prefix │ - │ │ │ │ _call_service_endpoint │ - │ │ │ ├──────────────────────▶│ - │ │ │ │◁ {content,buttons}─────│ - │ │ │◁─ OrchestrationResponse ───│ │ - │ │◁─ TestOrchResponse ─────│ │ │ - │◁─ render next MCQ │ │ │ │ - │ or final answer │ │ │ │ -``` - ---- - -## Error Handling - -| Error Scenario | Behavior | -|---|---| -| Service discovery fails | `execute_async()` returns `None` → falls back to RAG/context pipeline | -| Intent detection fails | Returns `None` → falls back to RAG/context pipeline | -| `_parse_service_prefix()` fails | `execute_direct_step()` returns `None` → falls through to normal pipeline | -| Service endpoint timeout | `_call_service_endpoint()` returns `None` → falls back | -| Service endpoint HTTP error | Logged, returns `None` → falls back | -| Buttons JSON parsing fails | Logs warning, returns empty buttons list `[]` | diff --git a/docs/TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md b/docs/TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md index 3e4ccfc0..05c0211f 100644 --- a/docs/TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md +++ b/docs/TESTPRODUCTIONLLM_SERVICE_WORKFLOW.md @@ -212,6 +212,7 @@ orchestration_service = self.orchestration_service service_content = service_result["content"] service_buttons = service_result["buttons"] + async def service_stream() -> AsyncIterator[str]: yield orchestration_service.format_sse( chat_id, service_content, service_buttons or None @@ -219,6 +220,7 @@ async def service_stream() -> AsyncIterator[str]: yield orchestration_service.format_sse(chat_id, "END") orchestration_service.log_costs(costs_metric) + return service_stream() ``` @@ -229,12 +231,13 @@ return service_stream() **`llm_orchestration_service.py` → `format_sse()`** (line 1195): ```python -def format_sse(self, chat_id: str, content: str, - buttons: Optional[List[Dict[str, Any]]] = None) -> str: +def format_sse( + self, chat_id: str, content: str, buttons: Optional[List[Dict[str, Any]]] = None +) -> str: inner_payload: Dict[str, Any] = {"content": content} if buttons: inner_payload["buttons"] = buttons - + payload = { "chatId": chat_id, "payload": inner_payload, diff --git a/docs/TOOL_CLASSIFIER.md b/docs/TOOL_CLASSIFIER.md new file mode 100644 index 00000000..d3e04dbf --- /dev/null +++ b/docs/TOOL_CLASSIFIER.md @@ -0,0 +1,154 @@ +# Tool Classifier + +## Overview + +The **Tool Classifier** is the entry router of the LLM Module. For every incoming query it inspects the +message (and conversation state) and dispatches it to exactly one **workflow** that knows how to answer. +Retrieval-Augmented Generation (RAG) is just one of those workflows. + +This document gives a high-level understanding of the classifier itself — its routing model, key +functions, and configuration. The **behaviour of each individual workflow is documented separately**; +see [Related documentation](#related-documentation). + +Source: [`src/tool_classifier/`](../src/tool_classifier/). The classifier only runs when +`TOOL_CLASSIFIER_ENABLED=true`; otherwise the service uses the RAG-only pipeline (backward compatible). + +--- + +## Routing model + +Workflows are evaluated as a **layer-wise chain** (Strategy pattern). Each workflow either handles the +query or returns `None` to fall through to the next layer. The order is defined by +`WORKFLOW_LAYER_ORDER` in [`enums.py`](../src/tool_classifier/enums.py): + +``` +User query + │ + ├─ active API-tool session for this chat_id? ──► short-circuit to API_TOOL_CALLING + │ + ▼ +Layer 1: SERVICE → external Bürokratt service calls + ↓ (None) +Layer 2: API_TOOL_CALLING → agentic external API tool calling + ↓ (None) +Layer 3: CONTEXT → greetings + conversation-history answers + ↓ (None) +Layer 4: RAG → knowledge-base retrieval + generation + ↓ (None) +Layer 5: OOD → out-of-domain fallback (always answers) +``` + +If the classifier errors at any point, it falls back to RAG (`FALLBACK_TO_RAG_ON_ERROR = True`). + +--- + +## How classification works + +`classify()` uses a **two-step search** to decide the workflow: + +1. **Dense search** — cosine similarity against the service collection for a relevance check. +2. **Hybrid search** — dense + sparse (BM25) vectors fused with Reciprocal Rank Fusion (RRF) to + identify the best-matching service. + +High-confidence matches route straight to `SERVICE`; if no service matches, an API-tool search may route +to `API_TOOL_CALLING`; otherwise the query falls through to `CONTEXT` / `RAG`. The full scoring scheme +(thresholds, score-gap logic, sparse encoding) is documented in +[HYBRID_SEARCH_CLASSIFICATION.md](./HYBRID_SEARCH_CLASSIFICATION.md). + +--- + +## Key functions + +### `ToolClassifier` ([`classifier.py`](../src/tool_classifier/classifier.py)) + +| Method | Purpose | +| --- | --- | +| `classify(query, conversation_history, language, request=None)` | Runs the two-step search and returns a `ClassificationResult` indicating the target workflow. Also handles the active-session short-circuit and intent-switch detection. | +| `route_to_workflow(classification, request, is_streaming, ...)` | Executes the chosen workflow with layer-wise fallback. Returns an `OrchestrationResponse` (non-streaming) or an SSE `AsyncIterator[str]` (streaming). | +| `aclose()` | Releases the shared Qdrant `httpx` client. | + +Internal helpers (not part of the public surface): `_dense_search()`, `_hybrid_search()`, +`_try_api_tool_classification()`, and `_execute_with_fallback_async/streaming()`. + +### `BaseWorkflow` ([`base_workflow.py`](../src/tool_classifier/base_workflow.py)) + +Every workflow executor inherits this contract: + +| Method | Purpose | +| --- | --- | +| `execute_async(request, context, time_metric=None)` | Non-streaming execution (`/orchestrate`, `/orchestrate/test`). Returns a response, or `None` to fall back to the next layer. | +| `execute_streaming(request, context, time_metric=None)` | Streaming execution (`/orchestrate/stream`). Returns an SSE `AsyncIterator[str]`, or `None` to fall back. | + +The **return-`None` fallback** is the mechanism that powers the layer chain. + +### `ClassificationResult` ([`models.py`](../src/tool_classifier/models.py)) + +| Field | Type | Description | +| --- | --- | --- | +| `workflow` | `WorkflowType` | Which workflow should handle the query. | +| `confidence` | `float` (0.0–1.0) | Confidence in the classification. | +| `metadata` | `dict` | Workflow-specific data passed to the executor (e.g. matched service/endpoint). | +| `reasoning` | `str \| None` | Human-readable explanation of the decision. | + +--- + +## Workflow executors + +Each `WorkflowType` maps to one executor under +[`src/tool_classifier/workflows/`](../src/tool_classifier/workflows/). Detailed behaviour lives in the +linked docs. + +| Workflow | Executor | Detailed documentation | +| --- | --- | --- | +| `SERVICE` | `service_workflow.py` | [TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md](./TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md) | +| `API_TOOL_CALLING` | `api_tool_workflow.py` | [API_TOOL_CALLING.md](./API_TOOL_CALLING.md) | +| `CONTEXT` | `context_workflow.py` | [CONTEXT_WORKFLOW_GREETING_DETECTION.md](./CONTEXT_WORKFLOW_GREETING_DETECTION.md) | +| `RAG` | `rag_workflow.py` | [CONTEXTUAL_RETRIEVAL_FLOW.md](./CONTEXTUAL_RETRIEVAL_FLOW.md) | +| `OOD` | `ood_workflow.py` | — (fixed out-of-domain response) | + +--- + +## Configuration + +### Feature flags ([`feature_flags.py`](../src/llm_orchestrator_config/feature_flags.py)) + +All are environment variables read at startup. + +| Flag | Default | Effect | +| --- | --- | --- | +| `TOOL_CLASSIFIER_ENABLED` | `false` | Master switch. When `false`, the service uses the RAG-only pipeline. | +| `SERVICE_WORKFLOW_ENABLED` | `true` | Enables Layer 1 (Service). | +| `API_TOOL_CALLING_WORKFLOW_ENABLED` | `true` | Enables Layer 2 (API tool calling). | +| `CONTEXT_WORKFLOW_ENABLED` | `true` | Enables Layer 3 (Context). | +| `MULTI_INTENT_ENABLED` | `true` | Enables the parallel multi-intent path (IntentDecomposer) in API tool calling. | +| `ATC_RESPONSE_CACHE_ENABLED` | `true` | Enables the two-tier Redis response cache for API tool calling. | +| `FALLBACK_TO_RAG_ON_ERROR` | `true` (constant) | Routes to RAG if the classifier raises. | + +> RAG and OOD have no flags — RAG is the core fallback and OOD is the final safety net. + +### Classification constants ([`constants.py`](../src/tool_classifier/constants.py)) + +| Constant | Value | Purpose | +| --- | --- | --- | +| `QDRANT_COLLECTION` | `intent_collections` | Qdrant collection for Bürokratt services. | +| `API_TOOL_COLLECTION` | `api_tool_collection` | Qdrant collection for registered API tool endpoints. | +| `DENSE_MIN_THRESHOLD` | `0.5` | Below this cosine → not a service match. | +| `DENSE_HIGH_CONFIDENCE_THRESHOLD` | `0.55` | At/above (with gap) → SERVICE without LLM confirmation. | +| `DENSE_SCORE_GAP_THRESHOLD` | `0.05` | Required lead of the top service over the runner-up. | +| `API_TOOL_MIN_THRESHOLD` | `0.40` | Below this → no API-tool match. | +| `API_TOOL_HIGH_CONFIDENCE_THRESHOLD` | `0.60` | At/above → API-tool single-path immediately. | +| `API_TOOL_INTENT_SWITCH_THRESHOLD` | `0.50` | Min cosine for a new endpoint to interrupt an active session. | + +See [HYBRID_SEARCH_CLASSIFICATION.md](./HYBRID_SEARCH_CLASSIFICATION.md) for how these thresholds combine, +and [API_TOOL_CALLING.md](./API_TOOL_CALLING.md) for the multi-intent / caching constants. + +--- + +## Related documentation + +- [Hybrid Search Classification](./HYBRID_SEARCH_CLASSIFICATION.md) — scoring, sparse encoding, intent enrichment. +- [Service Workflow](./TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md) — Layer 1 in depth. +- [API Tool Calling](./API_TOOL_CALLING.md) — Layer 2 agentic loop, multi-intent, response cache. +- [Context Workflow](./CONTEXT_WORKFLOW_GREETING_DETECTION.md) — Layer 3 greetings & history. +- [Contextual Retrieval Flow](./CONTEXTUAL_RETRIEVAL_FLOW.md) — Layer 4 RAG. +- [Architecture](./ARCHITECTURE.md) · [Documentation Index](./README.md) diff --git a/docs/TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md b/docs/TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md index ac92abb2..46a36863 100644 --- a/docs/TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md +++ b/docs/TOOL_CLASSIFIER_AND_SERVICE_WORKFLOW.md @@ -24,7 +24,9 @@ Layer 4: OOD → Out-of-domain fallback (polite rejection) ```python # Non-streaming mode classification = await classifier.classify(query, history, language) -response = await classifier.route_to_workflow(classification, request, is_streaming=False) +response = await classifier.route_to_workflow( + classification, request, is_streaming=False +) # Streaming mode classification = await classifier.classify(query, history, language) @@ -153,14 +155,13 @@ embedding = orchestration_service.create_embeddings_for_indexer([query]) # 2. Search Qdrant collection search_payload = { "vector": query_embedding, - "limit": 10, # Top 10 services (SEMANTIC_SEARCH_TOP_K) - "score_threshold": 0.2, # Minimum similarity (SEMANTIC_SEARCH_THRESHOLD) - "with_payload": True + "limit": 10, # Top 10 services (SEMANTIC_SEARCH_TOP_K) + "score_threshold": 0.2, # Minimum similarity (SEMANTIC_SEARCH_THRESHOLD) + "with_payload": True, } response = qdrant_client.post( - f"/collections/{QDRANT_COLLECTION}/points/search", - json=search_payload + f"/collections/{QDRANT_COLLECTION}/points/search", json=search_payload ) ``` @@ -182,12 +183,12 @@ Uses **DSPy + LLM** to intelligently match user query to a specific service and ```python class ServiceIntentDetector(dspy.Signature): # Inputs - user_query: str # "How much is 100 EUR in USD?" - available_services: str # JSON of service definitions - conversation_context: str # Recent 3 conversation turns - + user_query: str # "How much is 100 EUR in USD?" + available_services: str # JSON of service definitions + conversation_context: str # Recent 3 conversation turns + # Output - intent_result: str # JSON: {matched_service_id, confidence, entities, reasoning} + intent_result: str # JSON: {matched_service_id, confidence, entities, reasoning} ``` ### LLM Call Flow @@ -200,7 +201,7 @@ services_formatted = [ "name": "Currency Conversion", "description": "Convert EUR to other currencies", "required_entities": ["target_currency"], - "examples": ["How much is EUR in USD?", "Convert EUR to JPY"] # Top 3 examples + "examples": ["How much is EUR in USD?", "Convert EUR to JPY"], # Top 3 examples } ] @@ -216,7 +217,7 @@ with self.llm_manager.use_task_local(): intent_result = intent_module.forward( user_query="How much is 100 EUR in USD?", services=services_formatted, - conversation_history=conversation_history + conversation_history=conversation_history, ) ``` @@ -360,12 +361,12 @@ validation_errors = ["Entity 'target_currency' has empty value"] ```python { - "is_valid": True, # Always true (lenient validation) - "missing_entities": ["amount"], # Will send empty strings - "extra_entities": ["random_field"], # Will be ignored - "validation_errors": [ # Warnings only - "Entity 'amount' has empty value" - ] + "is_valid": True, # Always true (lenient validation) + "missing_entities": ["amount"], # Will send empty strings + "extra_entities": ["random_field"], # Will be ignored + "validation_errors": [ # Warnings only + "Entity 'amount' has empty value" + ], } ``` @@ -393,11 +394,7 @@ Ruuter services expect parameters in specific order: entities_schema = ["target_currency", "source_currency", "amount"] # LLM extraction (unordered dict) -entities_dict = { - "amount": "100", - "target_currency": "USD", - "source_currency": "EUR" -} +entities_dict = {"amount": "100", "target_currency": "USD", "source_currency": "EUR"} # Transform to ordered array entities_array = ["USD", "EUR", "100"] @@ -409,9 +406,7 @@ entities_array = ["USD", "EUR", "100"] ```python def _transform_entities_to_array( - self, - entities_dict: Dict[str, str], - entity_order: List[str] + self, entities_dict: Dict[str, str], entity_order: List[str] ) -> List[str]: """Transform entity dict to ordered array.""" if not entity_order: @@ -461,7 +456,7 @@ def _construct_service_endpoint(self, service_name: str, chat_id: str) -> str: payload = { "chatId": chat_id, "authorId": author_id, - "input": entities_array, # ["USD", "EUR", "100"] + "input": entities_array, # ["USD", "EUR", "100"] } ``` @@ -540,10 +535,10 @@ entities_dict = {"target_currency": "THB"} #### 4. Entity Validation ```python validation_result = { - "is_valid": True, - "missing_entities": [], - "extra_entities": [], - "validation_errors": [] + "is_valid": True, + "missing_entities": [], + "extra_entities": [], + "validation_errors": [], } ``` @@ -563,7 +558,7 @@ response = await _call_service_endpoint( http_method="POST", entities_array=["THB"], chat_id="...", - author_id="..." + author_id="...", ) # Returns content string from Ruuter response ``` @@ -632,25 +627,25 @@ RUUTER_SERVICE_BASE_URL = "http://ruuter-public:8086/services" RAG_SEARCH_RUUTER_PUBLIC = "http://ruuter-public:8086/rag-search" # Service call timeouts -SERVICE_CALL_TIMEOUT = 10 # seconds for external service calls -SERVICE_DISCOVERY_TIMEOUT = 10.0 # seconds for service discovery +SERVICE_CALL_TIMEOUT = 10 # seconds for external service calls +SERVICE_DISCOVERY_TIMEOUT = 10.0 # seconds for service discovery # Service selection thresholds -SERVICE_COUNT_THRESHOLD = 10 # Switch to semantic search if exceeded -MAX_SERVICES_FOR_LLM_CONTEXT = 50 # Max services to pass to LLM +SERVICE_COUNT_THRESHOLD = 10 # Switch to semantic search if exceeded +MAX_SERVICES_FOR_LLM_CONTEXT = 50 # Max services to pass to LLM # Semantic search QDRANT_COLLECTION = "intent_collections" -SEMANTIC_SEARCH_TOP_K = 10 # Top 10 relevant services -SEMANTIC_SEARCH_THRESHOLD = 0.2 # Minimum similarity score -QDRANT_TIMEOUT = 10.0 # seconds +SEMANTIC_SEARCH_TOP_K = 10 # Top 10 relevant services +SEMANTIC_SEARCH_THRESHOLD = 0.2 # Minimum similarity score +QDRANT_TIMEOUT = 10.0 # seconds # Hybrid search classification (see HYBRID_SEARCH_CLASSIFICATION.md) -DENSE_MIN_THRESHOLD = 0.38 # Minimum cosine to consider service match +DENSE_MIN_THRESHOLD = 0.38 # Minimum cosine to consider service match DENSE_HIGH_CONFIDENCE_THRESHOLD = 0.40 # Cosine for high-confidence path -DENSE_SCORE_GAP_THRESHOLD = 0.05 # Required gap between top two services -DENSE_SEARCH_TOP_K = 3 # Unique services from dense search -HYBRID_SEARCH_TOP_K = 5 # Results from hybrid RRF search +DENSE_SCORE_GAP_THRESHOLD = 0.05 # Required gap between top two services +DENSE_SEARCH_TOP_K = 3 # Unique services from dense search +HYBRID_SEARCH_TOP_K = 5 # Results from hybrid RRF search ``` --- diff --git a/docs/TOOL_CLASSIFIER_EXTENSION_SPEC.md b/docs/TOOL_CLASSIFIER_EXTENSION_SPEC.md deleted file mode 100644 index 38d81898..00000000 --- a/docs/TOOL_CLASSIFIER_EXTENSION_SPEC.md +++ /dev/null @@ -1,1940 +0,0 @@ -# Tool Classifier Extension - System Specification - -**Version**: 1.0 -**Date**: February 13, 2026 -**Status**: Design Specification - ---- - -## 1. Overview - -This document specifies the extension of the existing RAG Module with a **Tool Classifier** that implements layer-wise workflow routing. The classifier determines whether a user query should be handled by: - -1. **Service Workflow** - External service/API calls -2. **Context Workflow** - Conversation history-based responses -3. **RAG Workflow** - Knowledge base retrieval (existing) -4. **OOD Response** - Out of domain fallback - -### 1.1 Current State - -**Existing Flow:** -``` -User Query → Input Guardrails → Prompt Refiner → Contextual Retrieval → Response Generator → Output Guardrails -``` - -**Entry Points:** -- `POST /orchestrate` - Non-streaming orchestration -- `POST /orchestrate/test` - Testing environment with simplified input -- `POST /orchestrate/stream` - Server-sent events streaming - -### 1.2 Proposed Extension - -**New Flow:** -``` -User Query → Input Guardrails → Tool Classifier → [Service | Context | RAG | OOD] - ↓ - Layer 1: Service Check - ↓ (no match) - Layer 2: Context Check - ↓ (no match) - Layer 3: RAG Retrieval - ↓ (no chunks) - Layer 4: OOD Response -``` - ---- - -## 2. Architecture Changes - -### 2.1 Component Integration - -The Tool Classifier will be integrated into the existing `LLMOrchestrationService` with minimal disruption: - -```python -# Location: src/llm_orchestration_service.py - -def process_orchestration_request(self, request: OrchestrationRequest): - """ - Modified orchestration pipeline with tool classifier. - - Pipeline: - 1. Language Detection (existing) - 2. Query Validation (existing) - 3. Input Guardrails (existing, relocated) - 4. Tool Classifier (NEW) - 5. Workflow Routing (NEW) - """ - - # Existing: Step 0, 0.5 - detected_language = detect_language(request.message) - validation_result = validate_query_basic(request.message) - - # Existing: Component initialization - components = self._initialize_service_components(request) - - # Existing: Step 1 - Input Guardrails (RELOCATED before classifier) - if components["guardrails_adapter"]: - input_blocked = self.handle_input_guardrails(...) - if input_blocked: - return input_blocked - - # NEW: Step 2 - Tool Classifier - classifier_result = self.tool_classifier.classify( - query=request.message, - conversation_history=request.conversationHistory, - language=detected_language - ) - - # NEW: Step 3 - Workflow Routing - if classifier_result.workflow == WorkflowType.SERVICE: - return self._execute_service_workflow(request, classifier_result) - elif classifier_result.workflow == WorkflowType.CONTEXT: - return self._execute_context_workflow(request, classifier_result) - elif classifier_result.workflow == WorkflowType.RAG: - return self._execute_rag_workflow(request, classifier_result) - else: - return self._create_out_of_scope_response(request, detected_language) -``` - -### 2.2 New Components - -| Component | Location | Purpose | -|-----------|----------|---------| -| `ToolClassifier` | `src/tool_classifier/classifier.py` | Main classifier logic | -| `ServiceWorkflowExecutor` | `src/tool_classifier/service_workflow.py` | Service discovery and triggering | -| `ContextWorkflowExecutor` | `src/tool_classifier/context_workflow.py` | LLM-based conversation history analysis | -| `IntentEntityExtractor` | `src/tool_classifier/intent_extractor.py` | LLM-based intent/entity detection | -| `ServiceDiscoveryManager` | `src/tool_classifier/service_discovery.py` | Qdrant semantic search for services | -| `IntentCollectionSync` | `src/tool_classifier/intent_sync_service.py` | Database → Qdrant synchronization | -| `ContextAnalyzer` | `src/tool_classifier/context_analyzer.py` | LLM-based context availability checker | - -### 2.3 LLM Config Module Integration - -The existing LLM Config Module (`src/llm_config_module/`) is reused by the tool classifier for all LLM-based operations. No modifications to the core module are required. - -**Current LLM Config Module Capabilities:** -- **Multi-Provider Support**: Azure OpenAI, AWS Bedrock, OpenAI, Anthropic -- **Vault Integration**: Secure credential management via HashiCorp Vault -- **Connection Management**: Dynamic LLM connection selection based on `connection_id` from requests -- **Usage Tracking**: Token counting and cost calculation across providers - -**Tool Classifier LLM Usage:** - -| Workflow | LLM Operation | Config Usage | Temperature | -|----------|---------------|--------------|-------------| -| **Service (Layer 1)** | Intent & entity extraction | `llm_manager.call_llm_async()` | 0.0 (deterministic) | -| **Context (Layer 2)** | Context availability check | `llm_manager.call_llm_async()` | 0.0 (deterministic) | -| **RAG (Layer 3)** | Response generation | Existing integration | 0.7 (default) | -| **OOD (Layer 4)** | No LLM call | N/A | N/A | - -**Integration Pattern:** - -```python -# Tool classifier workflows use the same LLMManager instance -class ToolClassifier: - def __init__(self, llm_manager: LLMManager, ...): - self.llm_manager = llm_manager # Reuse existing instance - - async def detect_intent(self, query: str, services: List[Service]): - """Use LLM Config Module for intent detection.""" - response = await self.llm_manager.call_llm_async( - prompt=INTENT_DETECTION_PROMPT.format(...), - temperature=0.0, # Deterministic for classification - max_tokens=200 - ) - return parse_intent(response) -``` - -**Configuration Reuse:** -- Same connection selection logic (`connection_id` from `OrchestrationRequest`) -- Same Vault credential retrieval -- Same cost tracking pattern (`get_lm_usage_since()`) -- Same error handling and retry logic -- Same provider-specific implementations - -**No Changes Required**: The LLM Config Module is provider-agnostic and supports all tool classifier LLM calls out of the box. - ---- - -## 3. Layer 1: Service Workflow - -### 3.1 Workflow Logic - -When a user query is received, the system determines if it's a service-related request through the following steps: - -``` -1. Service Count Check → 2. Service Discovery → 3. Intent Detection → 4. Service Validation → 5. Entity Transformation → 6. Service Triggering -``` - -### 3.2 Step-by-Step Implementation - -#### Step 1: Service Count Check - -**Purpose**: Optimize performance based on service catalog size - -```python -# Query: SELECT COUNT(*) FROM services WHERE current_state = 'active' AND deleted = FALSE - -if service_count <= 50: - # Use all services for LLM context - services = get_all_active_services() -else: - # Use semantic search for top 20 most relevant - services = semantic_search_services(user_query, top_k=20) -``` - -**Database Query:** -```sql -SELECT COUNT(*) FROM public.services -WHERE current_state = 'active' AND deleted = FALSE; -``` - -#### Step 2: Semantic Search (When Service Count > 50) - -**Tool**: Qdrant vector database -**Collection**: `intent_collection` -**Vector Dimension**: 3072 (text-embedding-3-large) - -**Search Configuration:** -```python -search_params = { - "collection_name": "intent_collection", - "query_vector": embed_query(user_query), - "limit": 20, - "score_threshold": 0.5, # Higher threshold for service matching -} -``` - -**Output Format:** -```json -[ - { - "service_id": "exchange-rate-001", - "service_name": "ExchangeRateService", - "description": "Provides currency exchange rates", - "entities": ["fromCurrency", "toCurrency"], - "score": 0.87 - }, - ... -] -``` - -#### Step 3: LLM Intent Detection - -**Action**: Call LLM with user query and service context to extract: -- `intent`: Service name to trigger -- `entities`: Key-value pairs of extracted parameters - -**Prompt Template:** -```python -INTENT_DETECTION_PROMPT = """ -You are an intent classifier for government services. Analyze the user query and determine which service should handle the request. - -Available Services: -{service_list} - -User Query: "{user_query}" - -Task: -1. If the query matches a service, extract: - - intent: The exact service name to trigger - - entities: Key-value pairs of required parameters - -2. If NO service matches, respond with: {{"intent": null, "entities": null}} - -Response Format (JSON only, no explanation): -{{"intent": "ServiceName", "entities": {{"param1": "value1", "param2": "value2"}}}} -""" -``` - -**Expected LLM Response:** -```json -{ - "choices": [ - { - "message": { - "content": "{\"intent\": \"ExchangeRateService\", \"entities\": {\"fromCurrency\": \"EUR\", \"toCurrency\": \"USD\"}}" - } - } - ] -} -``` - -**Parsing Logic:** -```python -# Parse LLM response -content = response["choices"][0]["message"]["content"] -parsed = json.loads(content) - -if parsed["intent"] is None: - # No service match - move to Layer 2 (Context Workflow) - return WorkflowType.CONTEXT -``` - -#### Step 4: Service Validation - -**Action**: Validate the detected service against the database - -**Validation Query:** -```sql -SELECT service_id, name, ruuter_type, endpoints, structure, entities -FROM public.services -WHERE service_id = %(detected_service_id)s - AND current_state = 'active' - AND deleted = FALSE; -``` - -**Validation Checks:** -- Service exists in database -- `current_state = 'active'` -- `deleted = FALSE` - -**Failure Handling:** -```python -if not service_exists or not service_active: - logger.warning(f"Service validation failed: {detected_service_id}") - # Fallback to Layer 2 (Context Workflow) - return WorkflowType.CONTEXT -``` - -#### Step 5: Entity Transformation - -**Purpose**: Convert LLM entity object to array format for service payload - -**Input (from LLM):** -```json -{ - "fromCurrency": "EUR", - "toCurrency": "USD" -} -``` - -**Output (for service call):** -```json -["EUR", "USD"] -``` - -**Transformation Logic:** -```python -def transform_entities(entities: Optional[Dict[str, str]], - entity_order: List[str]) -> List[str]: - """ - Transform entity dictionary to ordered array. - - Args: - entities: LLM-extracted entity key-value pairs - entity_order: Expected entity order from service schema - - Returns: - Ordered list of entity values - """ - if not entities or entities is None: - return [] - - # Maintain order defined in service schema - return [entities.get(key, "") for key in entity_order] -``` - -**Example:** -```python -# Service schema defines: entities = ["fromCurrency", "toCurrency"] -transform_entities( - {"fromCurrency": "EUR", "toCurrency": "USD"}, - ["fromCurrency", "toCurrency"] -) -# Output: ["EUR", "USD"] -``` - -#### Step 6: Service Triggering - -**Purpose**: Call the external service endpoint with formatted payload - -**URL Construction:** -```python -# From database field 'endpoints' -base_url = "http://ruuter:8086" # From environment or service config -service_endpoint = f"{base_url}/services/active{service_name}" - -# Example: http://ruuter:8086/services/activeExchangeRateService -``` - -**HTTP Method:** -```python -# Retrieved from database field 'ruuter_type' -method = service.ruuter_type # 'GET' or 'POST' (ENUM) -``` - -**Payload Format:** -```json -{ - "input": ["EUR", "USD"], - "authorId": "user-67890", - "chatId": "chat-12345" -} -``` - -**Implementation:** -```python -async def trigger_service( - service: ServiceRecord, - entities: List[str], - request: OrchestrationRequest -) -> Dict[str, Any]: - """ - Trigger external service via Ruuter. - - Args: - service: Validated service record from database - entities: Transformed entity array - request: Original orchestration request - - Returns: - Service response or error - """ - url = f"{RUUTER_BASE_URL}/services/active{service.name}" - payload = { - "input": entities, - "authorId": request.authorId, - "chatId": request.chatId - } - - try: - if service.ruuter_type == "GET": - response = await http_client.get(url, params=payload, timeout=10) - else: # POST - response = await http_client.post(url, json=payload, timeout=10) - - response.raise_for_status() - return response.json() - - except httpx.TimeoutException: - logger.error(f"Service timeout: {service.service_id}") - raise ServiceTimeoutError() - except httpx.HTTPStatusError as e: - logger.error(f"Service error: {e.response.status_code}") - raise ServiceExecutionError() -``` - -**Response Handling:** - -**Non-Streaming:** -```python -service_response = await trigger_service(service, entities, request) -formatted_content = format_service_response(service_response) - -# Apply output guardrails -if guardrails_adapter: - output_check = await guardrails_adapter.check_output_async(formatted_content) - costs_metric["output_guardrails"] = output_check.usage - - if not output_check.allowed: - logger.warning(f"Service response blocked by guardrails: {output_check.reason}") - return create_guardrail_violation_response(request) - -# Return validated service response -return OrchestrationResponse( - chatId=request.chatId, - llmServiceActive=True, - questionOutOfLLMScope=False, - inputGuardFailed=False, - content=formatted_content -) -``` - -**Streaming:** -```python -service_response = await trigger_service(service, entities, request) -formatted_content = format_service_response(service_response) - -# Apply output guardrails validation -if guardrails_adapter: - output_check = await guardrails_adapter.check_output_async(formatted_content) - costs_metric["output_guardrails"] = output_check.usage - - if not output_check.allowed: - logger.warning(f"Service response blocked by guardrails") - yield format_sse(request.chatId, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE) - yield format_sse(request.chatId, "END") - return - -# Stream validated response token-by-token -for token in split_into_tokens(formatted_content, chunk_size=5): - yield format_sse(request.chatId, token) - await asyncio.sleep(0.01) # Maintain streaming UX - -yield format_sse(request.chatId, "END") -``` - -### 3.3 Failure Scenarios - -| Scenario | Action | -|----------|--------| -| No intent detected | Move to Layer 2 (Context Workflow) | -| Service validation failed | Move to Layer 2 (Context Workflow) | -| Service call timeout | Return `SERVICE_TIMEOUT_ERROR` message | -| Service returns error | Return `SERVICE_EXECUTION_ERROR` message | -| Entity extraction incomplete | Attempt service call with partial entities, or fallback to Layer 2 | -| Output guardrails blocked | Return `OUTPUT_GUARDRAIL_VIOLATION_MESSAGE` or fallback to Layer 2 | - -### 3.4 Output Guardrails for Service Responses - -**Why Service Responses Need Guardrails:** -- External services may return PII (personal identifiable information) -- Service errors could expose sensitive system details -- Third-party API responses are untrusted content -- Ensures consistent safety across all workflows - -**Integration Pattern:** - -Both non-streaming and streaming modes validate service responses before sending to users: - -```python -# Get service response -service_response = await trigger_service(...) - -# Apply output guardrails (validation-first) -if guardrails_adapter: - output_check = await guardrails_adapter.check_output_async(service_response) - if not output_check.allowed: - # Blocked - return error or fallback - return create_guardrail_violation_response(request) - -# Validated - return/stream to user -return/stream service_response -``` - ---- - -## 4. Layer 2: Context Workflow - -### 4.1 Workflow Logic - -If Layer 1 fails (no service match), use LLM to determine if the query is a greeting or can be answered from conversation history. - -**Trigger Conditions:** -- No service intent detected in Layer 1 -- Query is a greeting (hello, hi, good morning, etc.) **OR** -- Conversation history exists (at least 1 previous turn) and query references it - -### 4.2 Greeting Detection - -Greetings and conversational pleasantries are handled by the Context Workflow to provide natural, friendly responses without triggering service discovery or RAG retrieval. - -**Greeting Patterns (Multilingual):** - -```python -# Estonian greetings -ESTONIAN_GREETINGS = [ - "tere", "tervist", "tere hommikust", "tere päevast", "tere õhtust", - "hei", "hommikust", "õhtust", "päevast", "nägemist", - "tsau", "moi", "moikka" -] - -# English greetings -ENGLISH_GREETINGS = [ - "hello", "hi", "hey", "good morning", "good afternoon", "good evening", - "greetings", "howdy", "morning", "afternoon", "evening" -] - -# Farewell patterns -FAREWELL_PATTERNS = [ - "goodbye", "bye", "see you", "talk to you later", "ttyl", - "nägemist", "head aega", "kuni", "tsau" -] -``` - -**LLM-Based Greeting Detection:** - -Instead of rigid pattern matching, the LLM analyzes whether the query is a greeting or conversational message: - -```python -async def detect_greeting( - query: str, - llm_manager: LLMManager, - language: str -) -> GreetingResult: - """ - Use LLM to detect if query is a greeting/conversational message. - - Args: - query: User's message - llm_manager: LLM manager instance - language: Detected language (et/en) - - Returns: - GreetingResult with is_greeting flag and optional response - """ - prompt = GREETING_DETECTION_PROMPT.format( - user_query=query, - language=language - ) - - response = await llm_manager.call_llm_async( - prompt=prompt, - temperature=0.0, - max_tokens=150 - ) - - content = response["choices"][0]["message"]["content"] - result = json.loads(content) - - return GreetingResult( - is_greeting=result["is_greeting"], - greeting_type=result.get("greeting_type"), # 'hello', 'goodbye', 'thanks', etc. - suggested_response=result.get("suggested_response") - ) -``` - -**Greeting Detection Prompt:** - -```python -GREETING_DETECTION_PROMPT = """ -You are a greeting classifier. Determine if the user's message is a greeting, farewell, or conversational pleasantry. - -User Message: "{user_query}" -Language: {language} - -Task: -1. Identify if this is a greeting/conversational message (hello, hi, goodbye, thanks, etc.) -2. If YES: Classify the type and suggest an appropriate response -3. If NO: Indicate it's not a greeting - -Response Format (JSON only): -{{ - "is_greeting": true/false, - "greeting_type": "hello" | "goodbye" | "thanks" | "casual" | null, - "suggested_response": "friendly response in same language" | null -}} - -Examples of greetings: -- "Tere!" → {"is_greeting": true, "greeting_type": "hello"} -- "Good morning" → {"is_greeting": true, "greeting_type": "hello"} -- "Thanks for your help" → {"is_greeting": true, "greeting_type": "thanks"} -- "What are digital signatures?" → {"is_greeting": false} -""" -``` - -**Response Generation:** - -```python -if greeting_result.is_greeting: - # Use LLM-suggested response or fallback to predefined messages - response = greeting_result.suggested_response or get_default_greeting_response( - greeting_type=greeting_result.greeting_type, - language=language - ) - - return OrchestrationResponse( - chatId=request.chatId, - llmServiceActive=True, - questionOutOfLLMScope=False, - inputGuardFailed=False, - content=response - ) -``` - -### 4.3 LLM-Based Context Analysis - -Instead of using regex patterns, we use the LLM to intelligently determine if the query references conversation history and can be answered from it. - -**Conversation Window:** -```python -# Consider last 10 conversation turns (5 user + 5 bot pairs) -CONTEXT_WINDOW_SIZE = 10 - -def get_recent_history(history: List[ConversationItem]) -> List[ConversationItem]: - """Get recent conversation history for context analysis.""" - return history[-CONTEXT_WINDOW_SIZE:] if history else [] -``` - -**LLM Context Check Prompt:** -```python -CONTEXT_CHECK_PROMPT = """ -You are a conversation context analyzer. Analyze if the user's current query can be answered using ONLY the conversation history provided. - -Conversation History: -{conversation_history} - -Current User Query: "{user_query}" - -Task: -1. First check if this is a greeting/conversational message (hi, hello, thanks, goodbye, etc.) -2. If it's a greeting: Provide an appropriate friendly response -3. If NOT a greeting: Determine if the query references or can be answered from the conversation history above -4. If YES: Extract and provide the answer from the conversation history -5. If NO: Indicate that it cannot be answered from conversation history - -Response Format (JSON only, no explanation): -{{ - "is_greeting": true/false, - "can_answer_from_context": true/false, - "answer": "extracted answer from history OR greeting response" OR null, - "reasoning": "brief explanation of why it can/cannot be answered" -}} - -Examples of GREETINGS (handle with friendly response): -- "Tere!" → {"is_greeting": true, "answer": "Tere! Kuidas saan teid aidata?"} -- "Hello" → {"is_greeting": true, "answer": "Hello! How can I help you?"} -- "Thanks!" → {"is_greeting": true, "answer": "You're welcome!"} -- "Good morning" → {"is_greeting": true, "answer": "Good morning! What can I do for you?"} - -Examples of queries that CAN be answered from context: -- "What did you say earlier about that?" -- "Can you repeat that?" -- "What was the rate you mentioned?" -- "Tell me more about what you just said" - -Examples of queries that CANNOT be answered from context: -- Completely new topics -- Requests for real-time data -- Questions requiring external knowledge -""" -``` - -**Implementation:** -```python -async def check_context_availability( - query: str, - conversation_history: List[ConversationItem], - llm_manager: LLMManager -) -> ContextCheckResult: - """ - Use LLM to check if query can be answered from conversation history. - - Args: - query: Current user query - conversation_history: Recent conversation turns - llm_manager: LLM manager for making calls - - Returns: - ContextCheckResult with can_answer flag and optional answer - """ - # Get recent history - recent_history = get_recent_history(conversation_history) - - if not recent_history: - # No conversation history available - return ContextCheckResult( - can_answer_from_context=False, - answer=None, - reasoning="No conversation history available" - ) - - # Format conversation history for prompt - history_text = format_conversation_history(recent_history) - - # Call LLM with structured output request - prompt = CONTEXT_CHECK_PROMPT.format( - conversation_history=history_text, - user_query=query - ) - - try: - response = await llm_manager.call_llm_async( - prompt=prompt, - temperature=0.0, # Deterministic for classification - max_tokens=300 - ) - - # Parse structured JSON response - content = response["choices"][0]["message"]["content"] - result = json.loads(content) - - return ContextCheckResult( - is_greeting=result.get("is_greeting", False), - can_answer_from_context=result["can_answer_from_context"], - answer=result.get("answer"), - reasoning=result.get("reasoning", "") - ) - - except (json.JSONDecodeError, KeyError) as e: - logger.error(f"Failed to parse LLM context check response: {e}") - # Fallback: assume cannot answer from context - return ContextCheckResult( - can_answer_from_context=False, - answer=None, - reasoning="Failed to parse LLM response" - ) - -def format_conversation_history(history: List[ConversationItem]) -> str: - """Format conversation history for LLM prompt.""" - formatted = [] - for i, item in enumerate(history, 1): - role = "User" if item.authorRole == "user" else "Assistant" - formatted.append(f"{i}. {role}: {item.message}") - return "\n".join(formatted) -``` - -**Response Models:** -```python -from pydantic import BaseModel - -class ContextCheckResult(BaseModel): - """Result from LLM context availability check.""" - is_greeting: bool = False - can_answer_from_context: bool - answer: Optional[str] = None - reasoning: str = "" - -class GreetingResult(BaseModel): - """Result from greeting detection.""" - is_greeting: bool - greeting_type: Optional[str] = None # 'hello', 'goodbye', 'thanks', 'casual' - suggested_response: Optional[str] = None -``` - -### 4.3 Workflow Execution - -**Non-Streaming Response:** -```python -async def execute_context_workflow( - request: OrchestrationRequest, - llm_manager: LLMManager, - guardrails_adapter: Optional[NeMoRailsAdapter], - costs_metric: Dict -) -> Optional[OrchestrationResponse]: - """ - Execute context-based response workflow with output guardrails. - - Returns: - OrchestrationResponse with context-based answer or None to fallback to next layer - """ - # Check if query can be answered from conversation history - context_result = await check_context_availability( - query=request.message, - conversation_history=request.conversationHistory, - llm_manager=llm_manager - ) - - # Track costs - costs_metric["context_check"] = get_lm_usage_since(history_before) - - if (context_result.is_greeting or context_result.can_answer_from_context) and context_result.answer: - logger.info( - f"[{request.chatId}] Query answered from context " - f"(greeting: {context_result.is_greeting})" - ) - - # Apply output guardrails validation - if guardrails_adapter: - output_check = await guardrails_adapter.check_output_async( - context_result.answer - ) - costs_metric["output_guardrails"] = output_check.usage - - if not output_check.allowed: - logger.warning( - f"[{request.chatId}] Context response blocked by guardrails: " - f"{output_check.reason}" - ) - return create_guardrail_violation_response(request) - - # Return validated context-based response - return OrchestrationResponse( - chatId=request.chatId, - llmServiceActive=True, - questionOutOfLLMScope=False, - inputGuardFailed=False, - content=context_result.answer - ) - - else: - logger.info( - f"[{request.chatId}] Cannot answer from context: {context_result.reasoning}" - ) - # Fallback to Layer 3 (RAG Workflow) - return None # Signal to move to next layer -``` - -**Streaming Response:** -```python -async def execute_context_workflow_streaming( - request: OrchestrationRequest, - llm_manager: LLMManager, - guardrails_adapter: Optional[NeMoRailsAdapter], - costs_metric: Dict -) -> Optional[AsyncIterator[str]]: - """ - Execute context workflow with streaming support and output guardrails. - - Yields: - SSE-formatted strings with validated context-based response - - Returns: - None if cannot answer from context (signals fallback to next layer) - """ - # Check context availability (non-streaming, fast) - context_result = await check_context_availability( - query=request.message, - conversation_history=request.conversationHistory, - llm_manager=llm_manager - ) - - # Track costs - costs_metric["context_check"] = get_lm_usage_since(history_before) - - if (context_result.is_greeting or context_result.can_answer_from_context) and context_result.answer: - logger.info( - f"[{request.chatId}] Validating and streaming context-based response " - f"(greeting: {context_result.is_greeting})" - ) - - # Apply output guardrails validation BEFORE streaming - if guardrails_adapter: - output_check = await guardrails_adapter.check_output_async( - context_result.answer - ) - costs_metric["output_guardrails"] = output_check.usage - - if not output_check.allowed: - logger.warning( - f"[{request.chatId}] Context response blocked by guardrails (streaming)" - ) - yield format_sse(request.chatId, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE) - yield format_sse(request.chatId, "END") - return - - # Response validated - stream token by token for consistent UX - for token in split_into_tokens(context_result.answer, chunk_size=5): - yield format_sse(request.chatId, token) - await asyncio.sleep(0.01) # Maintain streaming pace - - # Signal completion - yield format_sse(request.chatId, "END") - - else: - logger.info(f"[{request.chatId}] No context match, falling back to RAG") - # Return None to signal fallback to next layer - # Caller will handle RAG workflow - return None - -def split_into_tokens(text: str, chunk_size: int = 5) -> List[str]: - """Split text into token-like chunks for streaming simulation.""" - words = text.split() - tokens = [] - for i in range(0, len(words), chunk_size): - chunk = " ".join(words[i:i + chunk_size]) - tokens.append(chunk + " " if i + chunk_size < len(words) else chunk) - return tokens -``` - -### 4.4 Advantages of LLM-Based Approach - - **No Regex Pattern Maintenance**: LLM understands semantic context references naturally - **Handles Edge Cases**: Can detect implicit references that regex would miss - **Multilingual Support**: Works across Estonian, English, and other languages - **Structured Output**: Consistent JSON format for easy parsing - **Reasoning Transparency**: Includes explanation of decision - **Streaming Compatible**: Fast context check + token-by-token answer delivery - **Greeting Detection**: Automatically handles greetings, farewells, and conversational pleasantries - **Natural Responses**: LLM generates contextually appropriate greeting responses - -### 4.7 Fallback Strategy - -**Fallback to Layer 3 (RAG):** -- If `is_greeting = false` AND `can_answer_from_context = false` -- If LLM response parsing fails -- If conversation history is empty (and not a greeting) -- If output guardrails block the response (fallback to RAG for alternative answer) - -**Error Handling:** -```python -try: - result = await execute_context_workflow( - request, llm_manager, guardrails_adapter, costs_metric - ) - if result: - return result # Context-based answer (validated) - else: - # Move to Layer 3 (RAG) - return await execute_rag_workflow(request, components, costs_metric) -except Exception as e: - logger.error(f"Context workflow failed: {e}") - # Fallback to RAG workflow - return await execute_rag_workflow(request, components, costs_metric) -``` - -**Guardrail Violation Fallback:** -```python -# Option 1: Return error message (current approach) -if not output_check.allowed: - return create_guardrail_violation_response(request) - -# Option 2: Fallback to RAG (alternative approach) -if not output_check.allowed: - logger.warning("Context response blocked, trying RAG workflow") - return await execute_rag_workflow(request, components, costs_metric) -``` - ---- - -## 5. Layer 3: RAG Workflow - -### 5.1 Integration with Existing System - -**Trigger**: When both Layer 1 (Service) and Layer 2 (Context) fail to match - -**Implementation:** -```python -# Reuse existing RAG pipeline -return self._execute_orchestration_pipeline( - request, components, costs_metric, time_metric -) -``` - -**Existing Flow (No Changes Required):** -1. Prompt Refinement -2. Contextual Retrieval (Qdrant + BM25) -3. Rank Fusion (RRF) -4. Response Generation -5. Output Guardrails (validation-first streaming already implemented) - -**Streaming with Output Guardrails (Current Implementation):** -```python -# RAG workflow uses validation-first approach -async for validated_chunk in guardrails_adapter.stream_with_guardrails( - user_message=refined_query, - bot_message_generator=llm_streaming_generator -): - # NeMo buffers tokens (chunk_size=200) - # Validates each buffer before yielding - yield format_sse(chatId, validated_chunk) - -yield format_sse(chatId, "END") -``` - -**Fallback:** -- If no chunks found (`len(relevant_chunks) == 0`) → Layer 4 (OOD) -- If response confidence low → Layer 4 (OOD) - ---- - -## 5.2 Streaming + Output Guardrails Comparison - -### Summary: How Each Workflow Handles Streaming + Validation - -| Workflow | Response Source | Validation Approach | Streaming Method | -|----------|----------------|---------------------|------------------| -| **RAG** | LLM streaming generation | NeMo buffers + validates chunks (chunk_size=200) | `stream_with_guardrails()` wraps bot generator | -| **Service** | External service (complete) | Validate complete response | Stream validated response token-by-token | -| **Context** | LLM structured output (complete) | Validate complete response | Stream validated response token-by-token | -| **OOD** | Fixed message | No validation needed | Stream fixed message token-by-token | - -### Technical Flow for Each Workflow - -#### RAG Workflow (Existing - Validation-First) - -**Non-Streaming:** -```python -response = await response_generator.generate(...) -output_check = await guardrails_adapter.check_output_async(response) -if output_check.allowed: - return OrchestrationResponse(content=response) -``` - -**Streaming:** -```python -# LLM generates via streaming -async def bot_generator(): - async for token in llm.stream(): - yield token - -# NeMo validates in real-time (buffers chunks) -async for validated_chunk in guardrails_adapter.stream_with_guardrails( - user_message=query, - bot_message_generator=bot_generator -): - yield format_sse(chatId, validated_chunk) # Already validated -``` - -#### Service Workflow (New - Validate Then Stream) - -**Non-Streaming:** -```python -service_response = await call_external_service(...) # Complete response -output_check = await guardrails_adapter.check_output_async(service_response) -if output_check.allowed: - return OrchestrationResponse(content=service_response) -else: - return GuardrailViolationResponse() -``` - -**Streaming:** -```python -service_response = await call_external_service(...) # Complete response - -# Validate complete response FIRST -output_check = await guardrails_adapter.check_output_async(service_response) -if not output_check.allowed: - yield format_sse(chatId, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE) - yield format_sse(chatId, "END") - return - -# Validated - now stream to client token-by-token -for token in split_into_tokens(service_response, chunk_size=5): - yield format_sse(chatId, token) - await asyncio.sleep(0.01) -yield format_sse(chatId, "END") -``` - -#### Context Workflow (New - Validate Then Stream) - -**Non-Streaming:** -```python -context_result = await llm.check_context(query, history) # Complete answer -if context_result.can_answer_from_context: - output_check = await guardrails_adapter.check_output_async(context_result.answer) - if output_check.allowed: - return OrchestrationResponse(content=context_result.answer) - else: - return GuardrailViolationResponse() -``` - -**Streaming:** -```python -context_result = await llm.check_context(query, history) # Complete answer - -if context_result.can_answer_from_context: - # Validate complete answer FIRST - output_check = await guardrails_adapter.check_output_async(context_result.answer) - if not output_check.allowed: - yield format_sse(chatId, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE) - yield format_sse(chatId, "END") - return - - # Validated - stream to client token-by-token - for token in split_into_tokens(context_result.answer, chunk_size=5): - yield format_sse(chatId, token) - await asyncio.sleep(0.01) - yield format_sse(chatId, "END") -``` - -### Key Differences - -**RAG Workflow:** -- **Real-time validation**: LLM generates → NeMo validates chunks → Stream to client -- **Buffered approach**: Tokens buffered in chunks of 200 characters -- **Bi-directional**: Generator feeding into NeMo, NeMo yielding validated chunks -- **Cost**: Inline (no separate validation call) - -**Service/Context Workflows:** -- **Pre-validation**: Get complete response → Validate → Stream to client -- **Complete response**: Already have full text before streaming starts -- **Uni-directional**: Simply chunk and send validated response -- **Cost**: Separate validation call tracked in `costs_metric["output_guardrails"]` -- **UX Consistency**: Simulates streaming to match RAG workflow behavior - -### Why Different Approaches? - -1. **RAG**: LLM streaming is inherently token-by-token, so NeMo can validate in real-time -2. **Service**: External API returns complete response, no streaming generation occurs -3. **Context**: LLM returns structured JSON with complete answer, not streaming - -### Common Pattern: Validation-First - -All three workflows share the **validation-first principle**: -- Content is validated BEFORE reaching the user -- Blocked content never sent to client -- Consistent safety guarantees across all workflows -- Streaming provides smooth UX even with complete responses (Service/Context) - ---- - -## 6. Layer 4: OOD (Out of Domain) Response - -### 6.1 Trigger Conditions - -- No service detected (Layer 1 failed) -- No context match (Layer 2 failed) -- No relevant knowledge chunks (Layer 3 failed) - -### 6.2 Response Generation - -**Return localized OOD message:** -```python -return OrchestrationResponse( - chatId=request.chatId, - llmServiceActive=True, - questionOutOfLLMScope=True, # Flag as out of scope - inputGuardFailed=False, - content=get_localized_message(OUT_OF_SCOPE_MESSAGES, detected_language) -) -``` - -**Existing Constants (Reuse):** -```python -# From: src/llm_orchestrator_config/llm_ochestrator_constants.py -OUT_OF_SCOPE_MESSAGES = { - "et": "Vabandust, ma ei suuda sellele küsimusele vastata...", - "en": "I apologize, but I cannot answer this question..." -} -``` - ---- - -## 7. Data Schemas - -### 7.1 Database Schema - -**Table: `services`** - -```sql --- Location: DSL/Liquibase/changelog/rag-search-script-v6-services.sql - --- Custom ENUM types -CREATE TYPE ruuter_request_type AS ENUM ('GET', 'POST'); -CREATE TYPE service_state AS ENUM ('active', 'inactive', 'draft'); - -CREATE TABLE public.services ( - -- Primary key - id BIGINT PRIMARY KEY, - - -- Basic service information - name TEXT NOT NULL, -- Service name (e.g., "ExchangeRateService") - description TEXT NOT NULL, -- Human-readable description - service_id TEXT NOT NULL UNIQUE, -- Unique identifier (e.g., "exchange-rate-001") - - -- Service classification - ruuter_type ruuter_request_type DEFAULT 'GET', -- HTTP method: 'GET' or 'POST' - current_state service_state DEFAULT 'draft', -- State: 'active', 'inactive', 'draft' - is_common BOOLEAN NOT NULL DEFAULT FALSE, -- Is this a common/shared service? - deleted BOOLEAN NOT NULL DEFAULT FALSE, -- Soft delete flag - - -- Intent classification data (for LLM) - slot TEXT NOT NULL DEFAULT '', -- Reserved for future use - entities text[] NOT NULL DEFAULT '{}', -- Expected entity names ["entity1", "entity2"] - examples text[] NOT NULL DEFAULT '{}', -- Example queries - - -- Service configuration - structure JSON NOT NULL DEFAULT '{}', -- Service schema/structure - endpoints JSON NOT NULL DEFAULT '[]', -- Endpoint configurations - - -- Timestamps - created_at TIMESTAMP WITH TIME ZONE DEFAULT CURRENT_TIMESTAMP, - updated_at TIMESTAMP WITH TIME ZONE DEFAULT CURRENT_TIMESTAMP -); - --- Indexes for performance -CREATE UNIQUE INDEX idx_services_service_id ON public.services(service_id); -CREATE INDEX idx_services_active ON public.services(current_state, deleted) - WHERE deleted = FALSE; -CREATE INDEX idx_services_name ON public.services(name); -``` - -**Update Master Changelog:** -```yaml -# Location: DSL/Liquibase/master.yml - -databaseChangeLog: - - include: - file: changelog/rag-search-script-v1-llm-connections.sql - - include: - file: changelog/rag-search-script-v2-user-management.sql - - include: - file: changelog/rag-search-script-v3-configuration.sql - - include: - file: changelog/rag-search-script-v4-authority-data.xml - - include: - file: changelog/rag-search-script-v5-prompt-config.sql - - include: - file: changelog/rag-search-script-v6-services.sql # NEW -``` - -### 7.2 Qdrant Collection Schema - -**Collection Name:** `intent_collection` - -**Configuration:** -```python -{ - "collection_name": "intent_collection", - "vectors_config": { - "size": 3072, # text-embedding-3-large - "distance": "Cosine" - } -} -``` - -**Document Schema:** -```json -{ - "id": "common_service_companies_workforce_taxes", - "name": "Ettevõtte tööjõumaksud", - "description": "Kasutaja soovib infot ettevõtte poolt tasutud tööjõumaksude kohta, näiteks palgamaksud ja sotsiaalmaks.", - "examples": [ - "ettevõtte tasutud tööjõumaksud", - "kui palju maksis ettevõte tööjõumakse", - "firma poolt tasutud tööjõumaksud" - ], - "entities": ["company_name"], - "text_for_embedding": "Kasutaja soovib infot ettevõtte poolt tasutud tööjõumaksude kohta, näiteks palgamaksud ja sotsiaalmaks.\nettevõtte tasutud tööjõumaksud\nkui palju maksis ettevõte tööjõumakse\nfirma poolt tasutud tööjõumaksud", - - "service_id": "common_service_companies_workforce_taxes", - "ruuter_type": "POST", - "current_state": "active" -} -``` - -**Field Mapping:** -| Qdrant Field | Source | Purpose | -|--------------|--------|---------| -| `id` | `services.service_id` | Unique identifier | -| `name` | `services.name` | Service display name | -| `description` | `services.description` | Service description | -| `examples` | `services.examples` | Example queries | -| `entities` | `services.entities` | Expected parameters | -| `text_for_embedding` | Computed | Concatenated text for vector embedding | -| `service_id` | `services.service_id` | Link to database record | -| `ruuter_type` | `services.ruuter_type` | HTTP method | -| `current_state` | `services.current_state` | Service status | - -**Embedding Text Construction:** -```python -def construct_embedding_text(service: ServiceRecord) -> str: - """ - Construct text for embedding from service data. - Format: description + examples (newline-separated) - """ - parts = [service.description] - parts.extend(service.examples) - return "\n".join(parts) -``` - -### 7.3 Database → Qdrant Synchronization - -**Trigger Mechanism:** -```sql --- PostgreSQL NOTIFY/LISTEN pattern or polling -CREATE OR REPLACE FUNCTION notify_service_change() -RETURNS TRIGGER AS $$ -BEGIN - IF TG_OP = 'INSERT' OR TG_OP = 'UPDATE' THEN - PERFORM pg_notify( - 'service_sync', - json_build_object( - 'action', TG_OP, - 'service_id', NEW.service_id, - 'current_state', NEW.current_state - )::text - ); - ELSIF TG_OP = 'DELETE' THEN - PERFORM pg_notify( - 'service_sync', - json_build_object( - 'action', 'DELETE', - 'service_id', OLD.service_id - )::text - ); - END IF; - RETURN NEW; -END; -$$ LANGUAGE plpgsql; - -CREATE TRIGGER service_sync_trigger -AFTER INSERT OR UPDATE OR DELETE ON services -FOR EACH ROW EXECUTE FUNCTION notify_service_change(); -``` - -**Sync Service:** -```python -# Location: src/tool_classifier/intent_sync_service.py - -class IntentCollectionSyncService: - """Synchronizes services table with Qdrant intent_collection.""" - - async def handle_service_change(self, event: Dict): - action = event['action'] - service_id = event['service_id'] - - if action in ['INSERT', 'UPDATE']: - # Fetch service from database - service = await self.db.fetch_service(service_id) - - # Generate embedding - embedding_text = self.construct_embedding_text(service) - embedding_vector = await self.embed(embedding_text) - - # Upsert to Qdrant - await self.qdrant_client.upsert( - collection_name="intent_collection", - points=[{ - "id": service.service_id, - "vector": embedding_vector, - "payload": { - "name": service.name, - "description": service.description, - "examples": service.examples, - "entities": service.entities, - "text_for_embedding": embedding_text, - "service_id": service.service_id, - "ruuter_type": service.ruuter_type, - "current_state": service.current_state - } - }] - ) - - elif action == 'DELETE': - await self.qdrant_client.delete( - collection_name="intent_collection", - points_selector={"points": [service_id]} - ) -``` - ---- - -## 8. Error Messages & Constants - -### 8.1 New Error Messages - -**Location:** `src/llm_orchestrator_config/llm_ochestrator_constants.py` - -```python -# Service Workflow Errors -SERVICE_NOT_FOUND_MESSAGES = { - "et": "Vabandust, ma ei leidnud sobivat teenust teie päringu jaoks.", - "en": "Sorry, I couldn't find a matching service for your request.", -} - -SERVICE_VALIDATION_FAILED_MESSAGES = { - "et": "Teenus ei ole hetkel saadaval.", - "en": "The requested service is currently unavailable.", -} - -SERVICE_TIMEOUT_ERROR_MESSAGES = { - "et": "Teenuse vastus võttis liiga kaua aega. Palun proovige hiljem uuesti.", - "en": "The service took too long to respond. Please try again later.", -} - -SERVICE_EXECUTION_ERROR_MESSAGES = { - "et": "Teenuse kutsumine ebaõnnestus. Palun proovige hiljem uuesti.", - "en": "Service execution failed. Please try again later.", -} - -ENTITY_EXTRACTION_FAILED_MESSAGES = { - "et": "Ma ei suutnud teie päringust vajalikku infot tuvastada.", - "en": "I couldn't extract the required information from your query.", -} - -# Context Workflow Errors -INSUFFICIENT_CONTEXT_MESSAGES = { - "et": "Ma ei leia vastust meie eelmisest vestlusest. Kas saate täpsustada?", - "en": "I can't find the answer in our previous conversation. Can you clarify?", -} - -NO_CONTEXT_AVAILABLE_MESSAGES = { - "et": "Mul pole piisavalt konteksti teie küsimusele vastamiseks.", - "en": "I don't have enough context to answer your question.", -} - -# Greeting Responses -GREETING_HELLO_MESSAGES = { - "et": "Tere! Kuidas saan teid aidata?", - "en": "Hello! How can I help you?", -} - -GREETING_GOODBYE_MESSAGES = { - "et": "Head aega! Kui vajate abi, olen siin.", - "en": "Goodbye! If you need help, I'm here.", -} - -GREETING_THANKS_MESSAGES = { - "et": "Pole tänu väärt! Kas saan veel kuidagi aidata?", - "en": "You're welcome! Can I help you with anything else?", -} - -GREETING_CASUAL_MESSAGES = { - "et": "Tere! Mida te soovite teada?", - "en": "Hi there! What would you like to know?", -} -``` - -**Helper Function for Default Greeting Responses:** - -```python -def get_default_greeting_response(greeting_type: str, language: str) -> str: - """ - Get default greeting response based on type and language. - - Args: - greeting_type: Type of greeting ('hello', 'goodbye', 'thanks', 'casual') - language: Language code ('et', 'en') - - Returns: - Localized greeting response - """ - greeting_map = { - "hello": GREETING_HELLO_MESSAGES, - "goodbye": GREETING_GOODBYE_MESSAGES, - "thanks": GREETING_THANKS_MESSAGES, - "casual": GREETING_CASUAL_MESSAGES - } - - messages = greeting_map.get(greeting_type, GREETING_HELLO_MESSAGES) - return messages.get(language, messages["en"]) -``` - -### 8.2 Reused Constants - -```python -# Already defined - reuse for consistency -OUT_OF_SCOPE_MESSAGE -TECHNICAL_ISSUE_MESSAGE -INPUT_GUARDRAIL_VIOLATION_MESSAGE -OUTPUT_GUARDRAIL_VIOLATION_MESSAGE -``` - ---- - -## 9. API Integration - -### 9.1 Entry Points (No Changes) - -The tool classifier is transparent to API consumers. All existing endpoints continue to work: - -**Non-Streaming:** -```http -POST /orchestrate -Content-Type: application/json - -{ - "chatId": "session-123", - "message": "What is the EUR to USD exchange rate?", - "authorId": "user-456", - "conversationHistory": [], - "url": "https://example.com", - "environment": "production", - "connection_id": "conn-789" -} -``` - -**Streaming:** -```http -POST /orchestrate/stream -Content-Type: application/json - -(Same request body as /orchestrate) -``` - -**Testing:** -```http -POST /orchestrate/test -Content-Type: application/json - -{ - "message": "Convert 100 EUR to USD", - "environment": "testing", - "connectionId": 1 -} -``` - -### 9.2 Response Format (No Changes) - -**Success Response:** -```json -{ - "chatId": "session-123", - "llmServiceActive": true, - "questionOutOfLLMScope": false, - "inputGuardFailed": false, - "content": "The current EUR to USD exchange rate is 1.08." -} -``` - -**Service Workflow Response:** -```json -{ - "chatId": "session-123", - "llmServiceActive": true, - "questionOutOfLLMScope": false, - "inputGuardFailed": false, - "content": "Based on the ExchangeRateService: EUR/USD = 1.0850" -} -``` - -The response format remains unchanged. The workflow selection is internal and transparent to the API consumer. - ---- - -## 10. Implementation Considerations - -### 10.1 Performance Optimization - -**Service Discovery Caching:** -```python -# Cache active service count for 5 minutes -@cached(ttl=300) -async def get_active_service_count() -> int: - return await db.count_active_services() -``` - -**Intent Collection Warm-up:** -```python -# Pre-load intent collection on startup -async def warmup_intent_collection(): - """Ensure intent_collection is ready before processing requests.""" - collection_info = await qdrant_client.get_collection("intent_collection") - logger.info(f"Intent collection ready: {collection_info.points_count} services") -``` - -### 10.2 Monitoring & Analytics - -**Tool Classifier Decisions Table:** -```sql --- Track classifier decisions for analytics -CREATE TABLE tool_classifier_decisions ( - id SERIAL PRIMARY KEY, - chat_id TEXT NOT NULL, - author_id TEXT, - user_query TEXT NOT NULL, - detected_workflow VARCHAR(20) NOT NULL, -- 'service', 'context', 'rag', 'ood' - classifier_confidence NUMERIC(5,4), - service_id VARCHAR(100), -- If service workflow - execution_time_ms INTEGER, - created_at TIMESTAMP DEFAULT NOW() -); - -CREATE INDEX idx_classifier_decisions_workflow - ON tool_classifier_decisions(detected_workflow); -``` - -### 10.3 Cost Tracking - -**Add tracking for new LLM calls:** -# Service workflow - intent detection -costs_metric["intent_detection"] = { - "total_prompt_tokens": usage.prompt_tokens, - "total_completion_tokens": usage.completion_tokens, - "total_cost": calculate_cost(usage) -} - -# Context workflow - context availability check -costs_metric["context_check -costs_metric["intent_detection"] = { - "total_prompt_tokens": usage.prompt_tokens, - "total_completion_tokens": usage.completion_tokens, - "total_cost": calculate_cost(usage) -} -``` - -### 10.4 Guardrails Strategy - -**Output Guardrails Application:** -```python -# Apply output guardrails to ALL workflows for consistency -WORKFLOWS_WITH_OUTPUT_GUARDRAILS = [ - WorkflowType.SERVICE, # Check service responses (may contain PII/sensitive data) - WorkflowType.CONTEXT, # Check context-based responses (conversation history may have PII) - WorkflowType.RAG # Existing behavior (knowledge base responses) -] - -# OOD responses skip guardrails (fixed message) -WORKFLOWS_WITHOUT_OUTPUT_GUARDRAILS = [ - WorkflowType.OOD -] -``` - -**Validation-First Approach:** - -All workflows use the **validation-first** approach where content is validated BEFORE streaming to the client: - -1. **RAG Workflow** (existing): - - LLM generates response via streaming - - NeMo buffers tokens (chunk_size=200) - - Each buffer validated before yielding - - Uses `stream_with_guardrails()` method - -2. **Service Workflow** (new): - - External service returns complete response - - Apply output guardrails validation - - Stream validated response token-by-token to client - - Consistent UX with RAG workflow - -3. **Context Workflow** (new): - - LLM returns complete answer from history - - Apply output guardrails validation - - Stream validated response token-by-token to client - - Consistent UX with RAG workflow - -**Streaming + Output Guardrails Integration:** - -```python -# For Service and Context workflows -async def stream_validated_response( - response_text: str, - guardrails_adapter: NeMoRailsAdapter, - request: OrchestrationRequest, - costs_metric: Dict -) -> AsyncIterator[str]: - """ - Apply output guardrails and stream validated response. - - Flow: - 1. Validate complete response with guardrails - 2. If allowed: Stream token-by-token to client - 3. If blocked: Send guardrail violation message - """ - # Check output guardrails (non-streaming validation) - output_check = await guardrails_adapter.check_output_async(response_text) - - # Track costs - costs_metric["output_guardrails"] = output_check.usage - - if not output_check.allowed: - logger.warning(f"[{request.chatId}] Output blocked by guardrails") - # Send violation message - yield format_sse(request.chatId, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE) - yield format_sse(request.chatId, "END") - return - - # Response validated - stream to client - logger.info(f"[{request.chatId}] Streaming validated response") - for token in split_into_tokens(response_text): - yield format_sse(request.chatId, token) - await asyncio.sleep(0.01) # Maintain streaming pace - - yield format_sse(request.chatId, "END") -``` - -**Utility Function for Token Streaming:** -```python -def split_into_tokens(text: str, chunk_size: int = 5) -> List[str]: - """ - Split text into token-like chunks for streaming simulation. - - Used by Service and Context workflows to provide streaming UX - even though the complete response is already available. - - Args: - text: Complete response text - chunk_size: Number of words per chunk - - Returns: - List of text chunks - """ - words = text.split() - tokens = [] - for i in range(0, len(words), chunk_size): - chunk = " ".join(words[i:i + chunk_size]) - tokens.append(chunk + " " if i + chunk_size < len(words) else chunk) - return tokens -``` - -### 10.5 Streaming Implementation Summary - -| Aspect | RAG Workflow | Service Workflow | Context Workflow | -|--------|--------------|------------------|------------------| -| **Response Type** | Streaming (token-by-token) | Complete (all at once) | Complete (all at once) | -| **Validation Timing** | Real-time (buffered chunks) | Pre-validation | Pre-validation | -| **Guardrail Method** | `stream_with_guardrails()` | `check_output_async()` | `check_output_async()` | -| **Streaming Reason** | Natural (LLM streams) | UX consistency | UX consistency | -| **Token Buffering** | NeMo 200-char chunks | Manual 5-word chunks | Manual 5-word chunks | -| **Cost Tracking** | Inline (timing = 0.0) | Separate call | Separate call | -| **Blocked Handling** | Stop mid-stream | Pre-check, don't stream | Pre-check, don't stream | -| **Client Experience** | Progressive reveal | Progressive reveal | Progressive reveal | - -**Implementation Status:** -- RAG streaming + guardrails: **Already implemented** (production-ready) -- Service streaming + guardrails: **To be implemented** (spec complete) -- Context streaming + guardrails: **To be implemented** (spec complete) - ---- - -## 11. Testing Strategy - -### 11.1 Unit Tests -async def test_context_detection_with_llm(): - query = "What did you say earlier?" - history = [ - ConversationItem(authorRole="bot", message="The EUR to USD rate is 1.08"), - ConversationItem(authorRole="user", message="Thanks") - ] - result = await context_analyzer.check_context_availability(query, history) - assert result.can_answer_from_context == True - assert "1.08" in result.answer - -async def test_context_detection_no_reference(): - query = "What are digital signatures?" - history = [ConversationItem(message="The rate is 1.08", ...)] - result = await context_analyzer.check_context_availability(query, history) - assert result.can_answer_from_context == False - -def test_rag_fallback(): - query = "What are digital signatures?" - result = classifier.classify(query, []) - assert result.workflow == WorkflowType.RAG - -async def test_context_streaming(): - """Test that context workflow supports streaming.""" - query = "What was the rate?" - history = [ConversationItem(message="The rate is 1.08", ...)] - - tokens = [] - async for token in context_workflow.execute_streaming(query, history): - tokens.append(token) - - assert len(tokens) > 0 - assert tokens[-1] == "END" - query = "What did you say earlier?" - history = [ConversationItem(message="The rate is 1.08", ...)] - result = classifier.classify(query, history) - assert result.workflow == WorkflowType.CONTEXT - -def test_rag_fallback(): - query = "What are digital signatures?" - result = classifier.classify(query, []) - assert result.workflow == WorkflowType.RAG -``` - -### 11.2 Integration Tests - -```python -# tests/integration_tests/test_service_workflow.py -async def test_full_service_workflow(): - request = OrchestrationRequest( - message="Convert 100 EUR to USD", - chatId="test-123", - ... - ) - response = await orchestration_service.process_orchestration_request(request) - assert response.llmServiceActive == True - assert "exchange rate" in response.content.lower() -``` - -### 11.3 Load `ContextAnalyzer` with LLM-based context checking -- Create context check prompt template with structured output -- Implement `ContextWorkflowExecutor` with streaming support -- Add conversation history formatting utilities -- Integration tests for context workflow (streaming + non-streaming) -- Cost tracking for context check LLM calls>50 services -locust -f tests/load/test_classifier_load.py --users 100 --spawn-rate 10 -``` - ---- - -## 12. Migration Path - -### 12.1 Phase 1: -- Create database migration for `services` table -- Create Qdrant `intent_collection` -- Relocate input guardrails before tool classifier -- Define error message constants - -### 12.2 Phase 2: -- Implement `ToolClassifier` with rule-based logic -- Implement workflow routing in `LLMOrchestrationService` -- Add classifier decision logging -- Unit tests for classifier - -### 12.3 Phase 3: Service Workflow -- Implement `ServiceDiscoveryManager` (Qdrant semantic search) -- Implement `IntentEntityExtractor` (LLM-based) -- Implement `ServiceWorkflowExecutor` (validation & triggering) -- Implement `IntentCollectionSyncService` (DB → Qdrant) -- Integration tests for service workflow - -### 12.4 Phase 4: Context Workflow -- ✅ ImpleHECK_TEMPERATURE=0.0 # Deterministic for classification -CONTEXT_CHECK_MAX_TOKENS=300tection -- Implement conversation history semantic search -- Implement `ContextWorkflowExecutor` -- Integration tests for context workflow - -### 12.5 Phase 5: Finalization -- Extend output guardrails to service & context workflows -- Implement fallback chain (service → context → rag → ood) -- Add comprehensive error handling -- Performance optimization (caching, async) -- End-to-end testing -- Production deployment - ---- - -## 13. Configuration - -### 13.1 Environment Variables - -```bash -# Service Workflow Configuration -RUUTER_BASE_URL=http://ruuter:8086 -SERVICE_DISCOVERY_TIMEOUT=2 # seconds -SERVICE_CALL_TIMEOUT=10 # seconds -MAX_SERVICES_FOR_LLM_CONTEXT=50 - -# Qdrant Configuration -QDRANT_INTENT_COLLECTION=intent_collection -INTENT_SEARCH_TOP_K=20 -INTENT_SEARCH_THRESHOLD=0.5 - -# Context Workflow Configuration -CONTEXT_WINDOW_SIZE=10 -CONTEXT_CONFIDENCE_THRESHOLD=0.7 -``` - -### 13.2 Feature Flags - -```python -# src/llm_orchestrator_config/feature_flags.py - -class FeatureFlags: - # Enable/disable tool classifier (rollback switch) - TOOL_CLASSIFIER_ENABLED = os.getenv("TOOL_CLASSIFIER_ENABLED", "true").lower() == "true" - - # Enable/disable specific workflows - SERVICE_WORKFLOW_ENABLED = os.getenv("SERVICE_WORKFLOW_ENABLED", "true").lower() == "true" - CONTEXT_WORKFLOW_ENABLED = os.getenv("CONTEXT_WORKFLOW_ENABLED", "true").lower() == "true" - - # Fallback to RAG if tool classifier fails - FALLBACK_TO_RAG_ON_ERROR = True -``` - ---- - -## 14. Rollback Strategy - -### 14.1 Graceful Degradation - -```python -def process_orchestration_request(self, request: OrchestrationRequest): - """Process with tool classifier or fallback to RAG.""" - - if not FeatureFlags.TOOL_CLASSIFIER_ENABLED: - # Fallback: Use existing RAG-only pipeline - logger.info("Tool classifier disabled - using RAG pipeline") - return self._execute_rag_workflow(request, None) - - try: - # New: Tool classifier routing - classifier_result = self.tool_classifier.classify(...) - return self._route_to_workflow(request, classifier_result) - - except Exception as e: - logger.error(f"Tool classifier failed: {e}") - if FeatureFlags.FALLBACK_TO_RAG_ON_ERROR: - logger.info("Falling back to RAG workflow") - return self._execute_rag_workflow(request, None) - raise -``` - -## 15. Success Metrics - -### 15.1 Performance Metrics - -| Metric | Target | Measurement | -|--------|--------|-------------| -| Tool Classifier Latency | < 200ms | p95 response time | -| Service Discovery (>50 services) | < 500ms | Qdrant search + LLM intent | -| Service Call Success Rate | > 95% | Successful service executions | -| Context Match Accuracy | > 80% | Correct context-based responses | -| End-to-End Latency | < 3s | Request to response | - -### 15.2 Quality Metrics - -| Metric | Target | Measurement | -|--------|--------|-------------| -| Workflow Classification Accuracy | > 90% | Manual evaluation sample | -| Service Intent Accuracy | > 85% | Correct service selection | -| Entity Extraction Accuracy | > 90% | Correct entity values | -| False Positive Rate (Service) | < 5% | Incorrect service routing | -| User Satisfaction | > 4.0/5.0 | User feedback surveys | - ---- diff --git a/docs/TOOL_CLASSIFIER_SKELETON_USAGE.md b/docs/TOOL_CLASSIFIER_SKELETON_USAGE.md deleted file mode 100644 index 38ce1f53..00000000 --- a/docs/TOOL_CLASSIFIER_SKELETON_USAGE.md +++ /dev/null @@ -1,542 +0,0 @@ -# Tool Classifier Skeleton - Usage Guide - -**Version**: 1.0 -**Date**: February 17, 2026 -**Status**: Skeleton Implementation - ---- - -## Overview - -This skeleton implements the **framework** for a multi-workflow routing system based on the [TOOL_CLASSIFIER_EXTENSION_SPEC.md](./TOOL_CLASSIFIER_EXTENSION_SPEC.md) specification. - -### Current Status - - **Implemented (Skeleton)**: -- Abstract base classes and interfaces -- Workflow executor skeletons (Service, Context, RAG, OOD) -- Tool classifier with classification and routing logic -- Feature flags for safe deployment -- Integration into LLMOrchestrationService - - **Not Implemented (Separate Tasks)**: -- Service discovery logic (Layer 1) -- Context analysis logic (Layer 2) -- Actual LLM calls in workflows -- Output guardrails integration for new workflows -- Database schema changes - -### Current Behavior - -When `TOOL_CLASSIFIER_ENABLED=false` (default): -- System works exactly as before (RAG-only pipeline) -- No changes to existing functionality - -When `TOOL_CLASSIFIER_ENABLED=true`: -- Classifier routes queries (currently always to RAG) -- Service and Context workflows return `None` (fallback to RAG) -- RAG workflow wraps existing pipeline -- All queries ultimately handled by RAG - ---- - -## Architecture - -### Layer-Wise Workflow Routing - -``` -User Query - ↓ -Input Guardrails - ↓ -Tool Classifier - ↓ -┌────────────────┐ -│ Classification │ -└────────┬───────┘ - ↓ - ┌─────┴──────┐ - │ Routing │ - └─────┬──────┘ - ↓ - ╔═══════════════════════════════════╗ - ║ Layer 1: Service Workflow ║ → (returns None - not implemented) - ╚═══════════════════════════════════╝ - ↓ (fallback) - ╔═══════════════════════════════════╗ - ║ Layer 2: Context Workflow ║ → (returns None - not implemented) - ╚═══════════════════════════════════╝ - ↓ (fallback) - ╔═══════════════════════════════════╗ - ║ Layer 3: RAG Workflow ║ → Handles query (existing pipeline) - ╚═══════════════════════════════════╝ - ↓ - Response to User -``` - -### Component Structure - -``` -src/tool_classifier/ -├── __init__.py # Module exports -├── enums.py # WorkflowType enum -├── models.py # ClassificationResult models -├── base_workflow.py # Abstract BaseWorkflow class -├── classifier.py # Main ToolClassifier -└── workflows/ - ├── __init__.py - ├── service_workflow.py # Layer 1 (skeleton) - ├── context_workflow.py # Layer 2 (skeleton) - ├── rag_workflow.py # Layer 3 (complete) - └── ood_workflow.py # Layer 4 (skeleton) -``` - -### Abstract Base Class Pattern - -The system uses **BaseWorkflow** as an abstract base class to ensure all workflows follow the same contract. - -#### How It Works - -1. **BaseWorkflow defines the contract**: - - Every workflow MUST implement two methods: `execute_async()` and `execute_streaming()` - - Both methods return `Optional[...]` to support the fallback pattern (return `None` → next layer) - - Python's `@abstractmethod` decorator enforces this at instantiation time - -2. **All workflows inherit from BaseWorkflow**: - - ServiceWorkflowExecutor extends BaseWorkflow → implements both methods - - ContextWorkflowExecutor extends BaseWorkflow → implements both methods - - RAGWorkflowExecutor extends BaseWorkflow → implements both methods - - OODWorkflowExecutor extends BaseWorkflow → implements both methods - -3. **Classifier treats all workflows uniformly**: - - The `ToolClassifier.route_to_workflow()` method doesn't need to know which specific workflow it's calling - - It just calls `workflow.execute_async()` or `workflow.execute_streaming()` - - This is **polymorphism** - same interface, different behavior - -4. **Benefits**: - - **Consistency**: All workflows have the same interface - - **Enforcement**: Can't create a workflow without implementing required methods - - **Flexibility**: Easy to add new workflows - just extend BaseWorkflow - - **Testability**: Each workflow can be tested independently - - **Fallback Pattern**: `Optional` return type enables layer chaining - -#### Example Flow - -``` -ToolClassifier needs to execute a workflow - ↓ -Gets workflow object (could be Service, Context, RAG, or OOD) - ↓ -Calls workflow.execute_async(request, context) - ↓ -BaseWorkflow contract guarantees this method exists - ↓ -Each workflow implements its own logic - ↓ -Returns OrchestrationResponse or None (fallback to next layer) -``` - -The abstract class is like a **blueprint** that says: "Any workflow in this system MUST be able to do these two things: execute normally and execute with streaming. I don't care *how* you do it, but you must provide these capabilities." - ---- - -## Feature Flags - -### Environment Variables - -```bash -# Master switch (default: false for safe deployment) -TOOL_CLASSIFIER_ENABLED=false - -# Individual workflow toggles (only apply when classifier enabled) -SERVICE_WORKFLOW_ENABLED=true -CONTEXT_WORKFLOW_ENABLED=true -``` - -### Configuration Class - -```python -from src.llm_orchestrator_config.feature_flags import FeatureFlags - -# Check if classifier is enabled -if FeatureFlags.TOOL_CLASSIFIER_ENABLED: - # Use tool classifier - pass - -# Check specific workflow -if FeatureFlags.is_workflow_enabled("service"): - # Service workflow logic - pass - -# Log current configuration -FeatureFlags.log_configuration() -``` - ---- - -## How It Works - -### 1. Non-Streaming Endpoint (`/orchestrate`) - -#### Current Flow (TOOL_CLASSIFIER_ENABLED=false) - -```python -POST /orchestrate - ↓ -LLMOrchestrationService.process_orchestration_request() - ↓ -Initialize components (LLM, guardrails, retriever, generator) - ↓ -Execute RAG pipeline - ↓ -Return OrchestrationResponse -``` - -#### With Classifier (TOOL_CLASSIFIER_ENABLED=true) - -```python -POST /orchestrate - ↓ -LLMOrchestrationService.process_orchestration_request() - ↓ -Initialize components - ↓ -Tool Classifier Integration: - 1. Initialize ToolClassifier (if first time) - 2. Classify query → ClassificationResult - - Currently always returns: WorkflowType.RAG - 3. Route to workflow: - - ServiceWorkflow.execute_async() → returns None - - ContextWorkflow.execute_async() → returns None - - RAGWorkflow.execute_async() → returns response - ↓ -Return OrchestrationResponse -``` - -### 2. Streaming Endpoint (`/orchestrate/stream`) - -#### Current Flow (TOOL_CLASSIFIER_ENABLED=false) - -```python -POST /orchestrate/stream - ↓ -LLMOrchestrationService.stream_orchestration_response() - ↓ -Initialize components - ↓ -Check input guardrails - ↓ -Refine prompt → Retrieve chunks → Stream through NeMo - ↓ -Yield SSE strings -``` - -#### With Classifier (TOOL_CLASSIFIER_ENABLED=true) - -```python -POST /orchestrate/stream - ↓ -LLMOrchestrationService.stream_orchestration_response() - ↓ -Initialize components - ↓ -Check input guardrails - ↓ -Tool Classifier Integration: - 1. Initialize ToolClassifier (if first time) - 2. Classify query → ClassificationResult - 3. Route to streaming workflow: - - ServiceWorkflow.execute_streaming() → returns None - - ContextWorkflow.execute_streaming() → returns None - - RAGWorkflow.execute_streaming() → yields SSE - ↓ -Yield SSE strings -``` - -### 3. Test Endpoint (`/orchestrate/test`) - -Works identically to `/orchestrate`: -- Converts `TestOrchestrationRequest` → `OrchestrationRequest` -- Routes through classifier (if enabled) -- Converts response back to `TestOrchestrationResponse` - ---- - -## Code Examples - -### Using the Classification System - -```python -from src.tool_classifier import ToolClassifier, WorkflowType, ClassificationResult - -# Initialize classifier -classifier = ToolClassifier( - llm_manager=llm_manager, - orchestration_service=service, -) - -# Classify a query -classification = await classifier.classify( - query="Hello, how are you?", - conversation_history=[], - language="en", -) - -# Check result -print(classification.workflow) # WorkflowType.RAG (in skeleton) -print(classification.confidence) # 1.0 -print(classification.reasoning) # "Default to RAG workflow..." - -# Route to workflow -response = await classifier.route_to_workflow( - classification=classification, - request=request, - is_streaming=False, -) -``` - -### Implementing a Workflow (Example) - -```python -from src.tool_classifier.base_workflow import BaseWorkflow -from models.request_models import OrchestrationRequest, OrchestrationResponse - -class MyCustomWorkflow(BaseWorkflow): - """Custom workflow implementation.""" - - async def execute_async( - self, - request: OrchestrationRequest, - context: Dict[str, Any], - ) -> Optional[OrchestrationResponse]: - """Handle query in non-streaming mode.""" - - # Check if this workflow can handle the query - can_handle = await self._check_if_applicable(request.message) - - if not can_handle: - # Return None to trigger fallback to next layer - return None - - # Execute workflow logic - result = await self._process_query(request.message) - - # Validate with output guardrails (TODO) - # is_safe = await guardrails.check_output_async(result) - # if not is_safe: - # return None or violation_response - - # Return response - return OrchestrationResponse( - chatId=request.chatId, - llmServiceActive=True, - questionOutOfLLMScope=False, - inputGuardFailed=False, - content=result, - ) - - async def execute_streaming( - self, - request: OrchestrationRequest, - context: Dict[str, Any], - ) -> Optional[AsyncIterator[str]]: - """Handle query in streaming mode.""" - - # Check if applicable - can_handle = await self._check_if_applicable(request.message) - - if not can_handle: - return None # Fallback - - # Get complete result - result = await self._process_query(request.message) - - # Validate with guardrails (TODO) - # is_safe = await guardrails.check_output_async(result) - # if not is_safe: - # yield format_sse(chatId, VIOLATION_MESSAGE) - # yield format_sse(chatId, "END") - # return - - # Stream result token-by-token - async def stream_result(): - for chunk in self._split_into_tokens(result): - yield self.format_sse(request.chatId, chunk) - await asyncio.sleep(0.01) - yield self.format_sse(request.chatId, "END") - - return stream_result() -``` - ---- - -## Deployment Strategy - -### Phase 1: Testing (Current State) - -```bash -# Keep classifier disabled -TOOL_CLASSIFIER_ENABLED=false -``` - -**Result**: System works exactly as before (RAG-only) - -### Phase 2: Enable Classifier (No Impact) - -```bash -# Enable classifier (but workflows not implemented) -TOOL_CLASSIFIER_ENABLED=true -SERVICE_WORKFLOW_ENABLED=true -CONTEXT_WORKFLOW_ENABLED=true -``` - -**Result**: -- Classifier runs but always routes to RAG -- Service/Context return `None` → fallback to RAG -- Functionally identical to Phase 1 -- Validates integration works - -### Phase 3: Implement Service Workflow - -1. Implement service discovery logic (separate task) -2. Deploy with `SERVICE_WORKFLOW_ENABLED=true` -3. Monitor service routing behavior -4. Rollback flag if issues occur - -### Phase 4: Implement Context Workflow - -1. Implement context analysis logic (separate task) -2. Deploy with `CONTEXT_WORKFLOW_ENABLED=true` -3. Monitor greeting/context detection -4. Rollback flag if issues occur - -### Phase 5: Production - -All workflows operational, full layer-wise routing active. - ---- - -## Extending the System - -### Adding a New Workflow - -1. **Create Workflow Executor**: - -```python -# src/tool_classifier/workflows/custom_workflow.py - -from src.tool_classifier.base_workflow import BaseWorkflow - -class CustomWorkflowExecutor(BaseWorkflow): - """Your custom workflow.""" - - async def execute_async(self, request, context): - # Implement logic - pass - - async def execute_streaming(self, request, context): - # Implement streaming logic - pass -``` - -2. **Register in Classifier**: - -```python -# src/tool_classifier/enums.py - -class WorkflowType(Enum): - SERVICE = "service" - CONTEXT = "context" - RAG = "rag" - CUSTOM = "custom" # Add new type - OOD = "ood" - -# Update layer order -WORKFLOW_LAYER_ORDER = [ - WorkflowType.SERVICE, - WorkflowType.CONTEXT, - WorkflowType.CUSTOM, # Add to chain - WorkflowType.RAG, - WorkflowType.OOD, -] -``` - -3. **Initialize in ToolClassifier**: - -```python -# src/tool_classifier/classifier.py - -def __init__(self, ...): - # ... existing workflows ... - self.custom_workflow = CustomWorkflowExecutor(...) -``` - -4. **Add Feature Flag**: - -```python -# src/llm_orchestrator_config/feature_flags.py - -CUSTOM_WORKFLOW_ENABLED = ( - os.getenv("CUSTOM_WORKFLOW_ENABLED", "true").lower() == "true" -) -``` - ---- - -## Key Concepts - -### 1. None Return Pattern - -Workflows return `None` when they cannot handle a query: - -```python -if not can_handle: - return None # Triggers fallback to next layer -``` - -This enables the fallback chain: Service → Context → RAG → OOD - -### 2. Validation-First Streaming - -For Service and Context workflows (complete responses): - -```python -# 1. Get complete response -response = await call_service(...) - -# 2. Validate BEFORE streaming -is_safe = await guardrails.check_output_async(response) - -if not is_safe: - yield format_sse(chatId, VIOLATION_MESSAGE) - yield format_sse(chatId, "END") - return - -# 3. Stream validated response -for chunk in split_into_tokens(response): - yield format_sse(chatId, chunk) -yield format_sse(chatId, "END") -``` - -### 3. Two Execution Methods - -Every workflow implements both: -- `execute_async()` → For `/orchestrate` (returns complete response) -- `execute_streaming()` → For `/orchestrate/stream` (yields SSE strings) - ---- - -## Summary - -This skeleton provides: - - **Complete framework** for multi-workflow routing - **Safe deployment** with feature flags - **Extensible architecture** using OOP patterns - **Backward compatibility** (disabled by default) - **Clear contracts** via abstract base classes - **Documentation** for implementation tasks - -The system is ready for workflow implementation in separate, independent tasks. - ---- diff --git a/docs/VAULT_SECURITY_ARCHITECTURE.md b/docs/VAULT_SECURITY_ARCHITECTURE.md index fe6fd741..2c2c6836 100644 --- a/docs/VAULT_SECURITY_ARCHITECTURE.md +++ b/docs/VAULT_SECURITY_ARCHITECTURE.md @@ -197,9 +197,12 @@ Day 0+: Automatic Token Renewal: Container Restart: vault-init: Check if Vault is sealed ↓ - If unsealed: Regenerate secret_id only + If unsealed: Validate existing secret_ids ↓ - vault-agent: Re-authenticate with new secret_id + If valid: Reuse existing secret_id (no churn) + If invalid: Mint new secret_id and write to disk + ↓ + vault-agent: Re-authenticate with secret_id ↓ New token issued and cached ``` @@ -413,8 +416,9 @@ Connected Services: - GUI (React Frontend) Token Lifecycle: - - Default Lease: 768h (32 days) - - Auto-renewal: Before expiration + - Token type: periodic (token_period 20m, no max-TTL) + - Auto-renewal: Every ~13 minutes (~2/3 of period) + - Re-auth: only on agent restart (never in steady state) ``` #### Agent 2: vault-agent-cron @@ -429,8 +433,9 @@ Connected Services: - CronManager (Python worker) Token Lifecycle: - - Default Lease: 768h (32 days) - - Auto-renewal: Before expiration + - Token type: periodic (token_period 30m, no max-TTL) + - Auto-renewal: Every ~20 minutes (~2/3 of period) + - Re-auth: only on agent restart (never in steady state) ``` #### Agent 3: vault-agent-llm @@ -445,8 +450,9 @@ Connected Services: - LLM Orchestration Service (FastAPI) Token Lifecycle: - - Default Lease: 1h (shorter for higher security) - - Auto-renewal: Every ~45 minutes + - Token type: periodic (token_period 1h, no max-TTL) + - Auto-renewal: Every ~40 minutes (~2/3 of period) + - Re-auth: only on agent restart (never in steady state) ``` ### Token Caching and Auto-Renewal @@ -464,29 +470,31 @@ T=0: Initial Authentication ├─► POST /v1/auth/approle/login │ Body: { role_id, secret_id } │ - └─► Receives: { token, ttl: 3600s, renewable: true } + └─► Receives: { token, period: 3600s, renewable: true } ← periodic token, no max-TTL │ └─► Cache token in: /agent/llm-token/token -T=45min: Proactive Renewal (75% of TTL) +T≈40min: Proactive Renewal (~2/3 of period) vault-agent monitors expiration │ ├─► POST /v1/auth/token/renew-self │ Header: X-Vault-Token: │ - └─► Receives: { token, ttl: 3600s } (same token, extended) + └─► Receives: { token, period: 3600s } (same token, period reset) │ └─► Update cache: /agent/llm-token/token + │ + └─► Repeats forever — a periodic token never hits a max-TTL, + so steady-state operation never needs approle/login again. -T=59min: Renewal Failed (fallback) - If renewal fails: +On agent restart only: + vault-agent re-reads role_id + secret_id from disk │ - ├─► Re-authenticate from scratch - │ POST /v1/auth/approle/login + ├─► POST /v1/auth/approle/login (secret_id must still be valid) │ - └─► New token issued and cached + └─► New periodic token issued and cached Application Request (anytime): @@ -856,15 +864,16 @@ Step 12: Check Vault Seal Status └─► GET /v1/sys/seal-status └─► If unsealed: Skip unseal steps -Step 13: Regenerate Secret IDs Only - └─► POST /v1/auth/approle/role/gui-service/secret-id - └─► POST /v1/auth/approle/role/cron-manager-service/secret-id - └─► POST /v1/auth/approle/role/llm-orchestration-service/secret-id - └─► Write new secret_ids to /agent/credentials/ +Step 13: Validate and Reconcile Secret IDs + └─► For each role (gui, cron-manager, llm-orchestration): + ├─► Test existing on-disk secret_id via AppRole login + ├─► If valid: Reuse (no change to credential file) + └─► If invalid/missing: Mint new secret_id and write to disk Note: role_ids remain unchanged (static identifiers) Note: Existing secrets and policies preserved Note: RSA keypair NOT regenerated (preserved) +Note: Stable secret_ids across restarts reduce credential churn ═══════════════════════════════════════════════════════════════════ COMPLETION @@ -1128,13 +1137,14 @@ Startup Order: vault-init Behavior: - Detects Vault already initialized - Skips initialization steps - - Regenerates secret_ids only - - Updates credential files + - Validates existing secret_ids (reuses if still valid) + - Mints new secret_ids only if existing ones are invalid Result: - All services start with fresh credentials + All services start with validated credentials Existing secrets preserved No manual intervention needed + Stable secret_ids reduce unnecessary credential churn ``` ### Token Regeneration Strategy @@ -1143,22 +1153,23 @@ Result: Current Implementation: 1. On Every Container Restart: - └─► vault-init regenerates secret_ids - └─► Vault agents get new tokens - └─► Old tokens remain valid until expiration + └─► vault-init validates existing secret_ids + ├─► If valid: Reuse (agents continue with same credentials) + └─► If invalid: Mint new secret_id, agents re-authenticate 2. Token Lifecycle: - └─► Issue: vault-agent authenticates + └─► Issue: vault-agent authenticates (periodic token, token_period per role) └─► Use: Application makes requests - └─► Renew: vault-agent extends TTL - └─► Expire: Automatic renewal failed - └─► Re-issue: vault-agent re-authenticates + └─► Renew: vault-agent renews within the period (~2/3 of period) + └─► No max-TTL: renewal continues indefinitely + └─► Re-issue: only on agent restart, via secret_id login 3. Security Benefits: - Short-lived tokens (1 hour for LLM, 32 days for others) - Automatic rotation on agent restart - No manual token management - Compromised tokens have limited lifetime + Periodic tokens (period 1h LLM, 30m Cron, 20m GUI), renewed continuously + Steady-state operation never re-runs approle/login (a stale secret_id + cannot strand a running agent) + Stable secret_ids (no unnecessary churn on restart) + Compromised tokens limited to one un-renewed period ``` ### Audit Logging Capabilities diff --git a/docs/VAULT_SETUP_AND_USAGE.md b/docs/VAULT_SETUP_AND_USAGE.md new file mode 100644 index 00000000..e61d362b --- /dev/null +++ b/docs/VAULT_SETUP_AND_USAGE.md @@ -0,0 +1,355 @@ +# Vault Setup & Usage Guide + +A single reference for how HashiCorp Vault is deployed, initialized, and consumed in the +RAG-Module. It covers the topology, the three Vault Agents, the secret layout, and — in +depth — **how each agent renews its token and how secrets are rotated**. + +Source files this document describes: + +- `docker-compose.yml` — service/topology definition +- `vault/config/vault.hcl` — Vault server config +- `vault-init.sh` — one-time bootstrap + per-restart reconcile +- `vault/agents/{gui,cron,llm}/*.hcl` — the three Vault Agent configs +- `DSL/CronManager/script/store_secrets_in_vault.sh` — writes/rotates secrets +- `DSL/CronManager/script/delete_secrets_from_vault.sh` — deletes secrets + +For the security rationale (threat model, defense-in-depth, access matrix) see the +companion `docs/VAULT_SECURITY_ARCHITECTURE.md`. This guide focuses on the *operational* +mechanics. + +--- + +## 1. Topology at a glance + +``` + bykstack (application network) vault-network (internal: true) + ┌───────────────────────────────────────────────┐ ┌──────────────────────────────┐ + │ gui ──────────────► vault-agent-gui :8202 ───┼────────┤ │ + │ cron-manager ─────► vault-agent-cron :8203 ───┼────────┤ vault :8200 │ + │ llm-orchestration ► vault-agent-llm :8201 ───┼────────┤ (Raft storage, KV v2, │ + │ │ │ AppRole auth) │ + │ vault-init (also on vault-network) ───────────┼────────┤ │ + └───────────────────────────────────────────────┘ └──────────────────────────────┘ +``` + +- **`vault`** runs only on `vault-network`, which is `internal: true` — it has **no route to + or from the host or the internet**. Port 8200 is never published. +- **Vault Agents** straddle both networks: they reach `vault` on `vault-network` and are + reachable by their owning application on `bykstack`. +- **Applications** talk *only* to their agent (`VAULT_ADDR=http://vault-agent-*:820x`) and + never hold a Vault token themselves. The agent injects the token transparently. + +| Service | Agent it uses | Agent address | AppRole | Policy | +|---|---|---|---|---| +| `gui` | `vault-agent-gui` | `:8202` | `gui-service` | `gui-policy` | +| `cron-manager` | `vault-agent-cron` | `:8203` | `cron-manager-service` | `cron-manager-policy` | +| `llm-orchestration-service` | `vault-agent-llm` | `:8201` | `llm-orchestration-service` | `llm-orchestration-policy` | + +--- + +## 2. Vault server (`vault/config/vault.hcl`) + +- **Storage:** Raft, single node (`node_id = vault-node-1`, path `/vault/file`, persisted in + the `vault-data` volume). No `retry_join` — a lone node self-bootstraps; adding a self- + pointing join was found to cause "Vault is sealed" boot loops. +- **Listener:** `0.0.0.0:8200`, `tls_disable = true` (TLS is terminated at the network + boundary; the network itself is the isolation layer here). Port `8201` is *not* given its + own listener because Vault uses it as the internal cluster port automatically. +- **Lease defaults:** `default_lease_ttl = 168h` (7 days), `max_lease_ttl = 720h` (30 days). + These are *system ceilings*; the per-AppRole token TTLs (below) are much shorter and are + what actually governs agent renewal cadence. +- `disable_mlock = false`, `ui = false`, JSON logs at INFO. + +Vault boots **sealed**. It must be unsealed before any operation — that is `vault-init`'s +first job. + +--- + +## 3. Bootstrap & reconcile (`vault-init.sh`) + +`vault-init` is a **run-once-then-exit** container (`restart: "no"`). The agents declare +`depends_on: vault-init: condition: service_completed_successfully`, so they only start +after init has finished cleanly. It runs `su vault -s /bin/sh /vault-init.sh` after creating +and `chown`ing the shared agent directories. + +The script has two branches, selected by the presence of `/vault/data/.initialized`. + +### 3.1 First-time deployment + +1. Wait for `/v1/sys/health` to respond. +2. **Initialize** with Shamir's Secret Sharing: `secret_shares=5`, `secret_threshold=3`. + The full response (5 unseal keys + root token) is written to + `/vault/data/unseal-keys.json`. +3. **Unseal** by submitting 3 of the 5 keys. +4. **Enable engines:** KV v2 at `secret/`, and the AppRole auth method. +5. **Create three ACL policies** (see §5). +6. **Create three AppRoles** issuing periodic tokens (see §4 — this is the heart of renewal), + via the `ensure_approles` helper. The same helper re-runs on subsequent deploys, so AppRole + config changes land without re-initializing Vault. +7. **Issue credentials:** for each role, fetch the static `role_id` and mint a `secret_id`, + writing both to `/agent/credentials/_role_id` and `_secret_id` (`chmod 640`). +8. **Generate an RSA-2048 keypair** with `openssl` and store it in Vault at + `secret/encryption/public_key` and `secret/encryption/private_key` + (algorithm `RSA-OAEP`, with `key_id` and `created_at` metadata). +9. Seed a test LLM secret, then `touch /vault/data/.initialized`. + +### 3.2 Subsequent deployment (restart) + +1. Check `/v1/sys/seal-status`; if sealed, reload the 3 unseal keys from + `unseal-keys.json` and unseal. +2. **Reconcile each secret_id** via `reconcile_secret_id`: + - `ensure_role_id` — make sure the `role_id` file exists (re-fetch from Vault if missing). + - `validate_secret_id` — attempt an AppRole login with the on-disk `role_id` + `secret_id`. + If it returns a `client_token`, the credential is still good. + - **Valid → reuse** the existing `secret_id` (no churn). + - **Invalid/missing → `mint_secret_id`** writes a fresh one. + +This is deliberate: because the AppRoles are created with `secret_id_ttl=0` and +`secret_id_num_uses=0` (non-expiring, unlimited-use), a single long-lived `secret_id` +survives normal restarts instead of being regenerated every boot. The RSA keypair, policies, +and stored secrets are all preserved across restarts. + +> **Note on file permissions:** `vault-init.sh` writes credential files with `chmod 640`. +> (The older architecture doc mentions `644`; the script is the source of truth — `640`.) + +--- + +## 4. The three Vault Agents — auth, renewal & rotation + +This is the core of the question. All three agents are the same Vault binary +(`hashicorp/vault:1.20.3`) run as `vault agent -config=...`. They differ only in which +credentials they read, which token sink they write, and their listener port. + +### 4.1 What an agent config actually does + +Example (`vault/agents/llm/agent.hcl`; gui/cron are identical in shape): + +```hcl +vault { address = "http://vault:8200"; retry { num_retries = 5 } } + +auto_auth { + method "approle" { + mount_path = "auth/approle" + config = { + role_id_file_path = "/agent/credentials/llm_role_id" + secret_id_file_path = "/agent/credentials/llm_secret_id" + remove_secret_id_file_after_reading = false + } + } + sink "file" { config = { path = "/agent/llm-token/token"; mode = 0640 } } +} + +cache { default_lease_duration = "1h" } +listener "tcp" { address = "0.0.0.0:8201"; tls_disable = true } +api_proxy { use_auto_auth_token = true } +``` + +Three mechanisms are at work: + +1. **`auto_auth` (authentication + renewal):** On startup the agent reads `role_id` + + `secret_id` and calls `POST /v1/auth/approle/login`. Vault returns a **periodic token** + (the AppRoles set `token_period`, defined in `vault-init.sh`, *not* in the HCL). The agent + then runs Vault's **auto-auth lifecycle manager**, which **renews the token automatically + in the background** before each period elapses. A periodic token has **no max-TTL**, so the + agent renews it indefinitely and — during normal operation — **never has to call + `approle/login` again**. The agent only re-authenticates (and thus only needs the + `secret_id` again) if it is **restarted** or if a renewal is missed long enough for the + token to lapse. `remove_secret_id_file_after_reading = false` keeps the `secret_id` on disk + so the agent can re-auth after a restart without `vault-init` re-minting. + + > **Why periodic tokens?** An earlier design issued tokens with `token_ttl`/`token_max_ttl`, + > which forced a full re-login every time `token_max_ttl` was reached. If the `secret_id` + > had become invalid by then (expiry, clock skew, server re-init), the agent got stuck in an + > `invalid role or secret ID` 400 backoff loop with no way to self-heal. Periodic tokens + > remove that re-login from the steady state, so a stale `secret_id` can no longer strand a + > running agent. +2. **`sink "file"` (token hand-off):** Every time the agent obtains/renews a token it writes + it to a file (`/agent/-token/token`, mode `0640`). The compose **health check** for + each agent is simply `test -f && test -s ` — a non-empty token file means + the agent has authenticated successfully. +3. **`api_proxy { use_auto_auth_token = true }` (transparent injection):** The agent also + listens as an HTTP proxy on its port. When the application sends a token-less request, the + agent injects `X-Vault-Token: ` and forwards it to `vault:8200`. + This is why application code never sets `VAULT_TOKEN`. + +> **`cache.default_lease_duration` is not the token TTL.** It is the agent's cache lease +> hint. The authoritative token lifetime comes from the AppRole's `token_period` in +> `vault-init.sh`. The per-agent cache hint is set to match the period. + +### 4.2 Per-agent renewal parameters + +AppRole token settings are created in `vault-init.sh`; all three use +`token_period` (periodic token, **no max-TTL**), `secret_id_ttl=0`, `secret_id_num_uses=0`, +`token_num_uses=0`, `bind_secret_id=true`. + +| Agent | AppRole | `token_period` | Proactive renewal (~⅔ of period) | Re-login (`approle/login`) | +|---|---|---|---|---| +| `vault-agent-gui` | `gui-service` | **20m** | ~every 13 min | only on agent restart | +| `vault-agent-cron` | `cron-manager-service` | **30m** | ~every 20 min | only on agent restart | +| `vault-agent-llm` | `llm-orchestration-service` | **1h** | ~every 40 min | only on agent restart | + +Reading the lifecycle for, e.g., the LLM agent: + +``` +T=0 login → periodic token (period 1h) → written to /agent/llm-token/token +T≈40m renew-self → period resets to 1h → token file refreshed +... renew repeats forever; token never hits a max-TTL +(restart) agent re-runs approle/login with the on-disk secret_id → fresh token +``` + +The periods are tuned per service (shorter for the GUI, which only reads the public key; +longer for the high-traffic LLM read path), but functionally all three behave the same: +**renew forever, re-login only on restart.** + +### 4.3 Two distinct "rotation" concepts — keep them separate + +1. **Token rotation (automatic, continuous):** Handled entirely by the agent's `auto_auth` + loop as described above — the periodic token is renewed indefinitely with no human action + and no `vault-init` involvement. +2. **`secret_id` rotation (rare):** The `secret_id` is the long-lived credential the agent + uses to *log in* (at startup/restart only, now that tokens are periodic). It is configured + non-expiring (`secret_id_ttl=0`, `secret_id_num_uses=0`) and is only replaced by + `vault-init` on a restart when the existing one fails validation (§3.2). To force rotation, + delete the `secret_id` file (or invalidate it in Vault) and re-run `vault-init`, then + restart the agent so it logs in with the freshly minted one. + + > **Operational caveat (learned the hard way):** if a `secret_id` ever does become invalid + > while an agent is running, the periodic-token design means a *running* agent keeps working + > (it only renews, never re-logs-in). But a **restarted** agent needs a valid `secret_id` to + > log in. Recovery is always: re-run `vault-init` (mints a fresh `secret_id` via the §3.2 + > reconcile) → restart the affected agent. See `docs/` runbook / the troubleshooting note + > below. + +### 4.4 Restart behavior + +- **Restart an agent:** It re-reads `role_id`/`secret_id` from the (read-only) creds volume + and re-authenticates. New token, written to the sink. App sees a brief blip. +- **Restart `vault`:** Data persists; `vault-init` (or the existing agent tokens, if still + valid) handle re-unseal/re-auth. Existing tokens remain valid if not expired. +- **Full `down && up`:** Order is `vault → vault-init → agents → apps`. `vault-init` detects + the `.initialized` flag, skips first-time setup, reconciles secret_ids, and the agents + start with validated credentials. + +--- + +## 5. Authorization — policies (who can touch what) + +Created in `vault-init.sh`. Paths are KV v2, so data lives under `secret/data/...` and +listing/metadata under `secret/metadata/...`. + +| Path | `gui-policy` | `cron-manager-policy` | `llm-orchestration-policy` | +|---|---|---|---| +| `secret/data/encryption/public_key` | **read** | read | — | +| `secret/data/encryption/private_key` | **deny** | **read** | — | +| `secret/data/encryption/*` | — | — | **deny** | +| `secret/data/llm/connections/*` | deny | **create/read/update/delete** | **read, list** | +| `secret/data/embeddings/connections/*` | deny | **create/read/update/delete** | **read, list** | +| `auth/token/lookup-self` | — | read | read | + +The intent, by tier: + +- **GUI** — can read *only* the public key, to encrypt user-entered credentials in the + browser before they ever leave it. Everything else is explicitly denied. +- **CronManager** — the only writer. Reads the **private key** to decrypt what the GUI + encrypted, then writes plaintext credentials into Vault. Full CRUD on connection secrets. +- **LLM Orchestration** — read-only consumer of connection secrets. **Explicitly denied** all + encryption keys, so a compromise of this hot-path service cannot exfiltrate the private key. + +--- + +## 6. Secret layout (KV v2 under `secret/`) + +``` +secret/ +├── llm/connections// ← e.g. aws_bedrock, azure_openai +├── embeddings/connections// +└── encryption/ + ├── public_key { key, algorithm: RSA-OAEP, key_size: 2048, key_id, created_at } + └── private_key { key, algorithm: RSA-OAEP, key_size: 2048, key_id, created_at } +``` + +The current write/delete scripts key connection secrets by a stable **`vaultUuid`** as the +final path segment (environment is tracked in the DB, not the path). KV v2 versions every +write, so updating a credential keeps prior versions for audit/rollback. + +LLM secret shape (AWS): `{ connection_id, access_key, secret_key, model, tags }`. +Azure: `{ connection_id, endpoint, api_key, deployment_name, model, api_version, tags }`. + +--- + +## 7. Usage flows + +### 7.1 Storing / rotating a credential (`store_secrets_in_vault.sh`, via cron-manager) + +1. GUI encrypts the raw key with the RSA **public** key and submits it. +2. The cron-manager job runs the script against `vault-agent-cron:8203` (no token — the agent + injects it). +3. The script **fetches the private key** (`GET secret/data/encryption/private_key`), then + decrypts each sensitive field in-memory via `decrypt_vault_secrets.py` (RSA-OAEP). +4. It builds the JSON payload with `jq` and `POST`s plaintext to + `secret/data//connections//`. Re-posting the same path + = a KV v2 version bump = credential rotation. +5. Sensitive shell variables are `unset` immediately after use. + +### 7.2 Deleting a credential (`delete_secrets_from_vault.sh`) + +`DELETE`s both `secret/data/...` and `secret/metadata/...` for the connection (404 treated as +success), again through `vault-agent-cron` with no explicit token. + +### 7.3 Reading a credential (LLM orchestration) + +The LLM service issues a token-less `GET http://vault-agent-llm:8201/v1/secret/data/llm/...`. +`vault-agent-llm` injects its cached token, Vault validates it against +`llm-orchestration-policy`, and returns the secret. The service then calls AWS/Azure with it. + +--- + +## 8. Operational notes & known trade-offs + +- **Unseal keys + root token sit in the `vault-data` volume** (`unseal-keys.json`). This makes + auto-unseal on restart trivial but is a **dev/test convenience**. For production, switch to + auto-unseal backed by a cloud KMS/HSM and remove the keys from the volume. +- **Root token** is used only by `vault-init` and is never injected into app containers. Best + practice for production is to revoke it after bootstrap and use scoped admin policies. +- **TLS is disabled** on the Vault listener and agent listeners; isolation relies on the + `internal: true` `vault-network`. Add TLS for any non-local deployment. +- **Audit logging is available but not enabled.** Turn it on with + `vault audit enable file file_path=/vault/logs/audit.log` (the `./vault/logs` mount already + exists) for a full request trail. +- **Credential files are world-readable within the shared volume** (mode 640, single owner, + but all agents mount the same `vault-agent-creds` volume read-only) — isolation is at the + volume level, not per-file. Fine for this trust boundary; note it if the threat model + tightens. + +--- + +## 9. Troubleshooting: agents looping on `invalid role or secret ID` + +**Symptom:** an agent logs `lifetime watcher done channel triggered, re-authenticating` +followed by repeating `PUT .../auth/approle/login → Code: 400 ... invalid role or secret ID` +with growing backoff. Token *renewals* had been succeeding up to that point. + +**Cause:** the agent's `secret_id` became invalid server-side (expiry, clock skew, or a Vault +re-init), and the agent reached a point where it had to do a full `approle/login`. With the +old `token_ttl`/`token_max_ttl` design this happened on every `token_max_ttl` cycle; the +switch to **periodic tokens** (§4) removes re-login from steady state, so a *running* agent no +longer hits this — but a **restarted** agent still needs a valid `secret_id`. + +**Recovery:** + +```bash +# Mint fresh secret_ids (vault-init's reconcile detects the invalid ones and replaces them) +docker compose up -d --force-recreate vault-init +docker wait vault-init +# Restart the affected agents so they log in with the fresh secret_id +docker compose restart vault-agent-gui vault-agent-cron vault-agent-llm +``` + +**Confirm root cause (read-only):** + +```bash +ROOT=$(docker exec vault sh -c "grep -o '\"root_token\":\"[^\"]*\"' /vault/file/unseal-keys.json | cut -d: -f2 | tr -d '\"'") +docker exec -e VAULT_TOKEN=$ROOT -e VAULT_ADDR=http://127.0.0.1:8200 vault \ + vault read auth/approle/role/gui-service # expect token_period set, secret_id_ttl=0 +echo "host: $(date -u)"; docker exec vault date -u # check for WSL2/Docker clock drift +``` diff --git a/docs/images/LLM Module App Diagram (Current).png b/docs/images/LLM Module App Diagram (Current).png new file mode 100644 index 00000000..12a44379 Binary files /dev/null and b/docs/images/LLM Module App Diagram (Current).png differ diff --git a/docs/images/LLM Module Context Diagram (Current).png b/docs/images/LLM Module Context Diagram (Current).png new file mode 100644 index 00000000..4e4b7ca7 Binary files /dev/null and b/docs/images/LLM Module Context Diagram (Current).png differ diff --git a/docs/images/LLM Orchestration Service Component Diagram (Current).png b/docs/images/LLM Orchestration Service Component Diagram (Current).png new file mode 100644 index 00000000..d011b9e5 Binary files /dev/null and b/docs/images/LLM Orchestration Service Component Diagram (Current).png differ diff --git a/endpoints.md b/endpoints.md deleted file mode 100644 index 262e81a3..00000000 --- a/endpoints.md +++ /dev/null @@ -1,683 +0,0 @@ -# LLM Connections API Endpoints - -## Base URL -``` -/ruuter-private/llm/connections -``` - ---- - -## 1. Create LLM Connection - -### Endpoint -```http -POST /ruuter-private/llm/connections/create -``` - -### Request Body -```json -{ - "llmPlatform": "OpenAI", - "llmModel": "GPT-4o", - "embeddingPlatform": "OpenAI", - "embeddingModel": "text-embedding-3-small", - "monthlyBudget": 1000.00, - "deploymentEnvironment": "Testing", - // Azure credentials (optional) - "deploymentName": "my-deployment", - "targetUri": "https://my-endpoint.azure.com", - "apiKey": "azure-api-key", - // AWS Bedrock credentials (optional) - "secretKey": "aws-secret-key", - "accessKey": "aws-access-key", - // Embedding model credentials (optional) - "embeddingModelApiKey": "embedding-api-key" -} -``` - -### Response (201 Created) -```json -{ - "id": 1, - "llmPlatform": "OpenAI", - "llmModel": "GPT-4o", - "embeddingPlatform": "OpenAI", - "embeddingModel": "text-embedding-3-small", - "monthlyBudget": 1000.00, - "usedBudget": 0.00, - "deploymentEnvironment": "Testing", - "status": "active", - "createdAt": "2025-09-02T10:15:30.000Z", - // Azure credentials (if provided) - "deploymentName": "my-deployment", - "targetUri": "https://my-endpoint.azure.com", - "apiKey": "azure-api-key", - // AWS Bedrock credentials (if provided) - "secretKey": "aws-secret-key", - "accessKey": "aws-access-key", - // Embedding model credentials (if provided) - "embeddingModelApiKey": "embedding-api-key" -} -``` - ---- - -## 2. Update LLM Connection - -### Endpoint -```http -POST /ruuter-private/llm/connections/update -``` - -### Request Body -```json -{ - "connectionId": 1, - "llmPlatform": "Azure AI", - "llmModel": "GPT-4o-mini", - "embeddingPlatform": "Azure AI", - "embeddingModel": "text-embedding-ada-002", - "monthlyBudget": 2000.00, - "deploymentEnvironment": "Production", - // Azure credentials (optional) - "deploymentName": "updated-deployment", - "targetUri": "https://updated-endpoint.azure.com", - "apiKey": "updated-azure-api-key", - // AWS Bedrock credentials (optional) - "secretKey": "updated-aws-secret-key", - "accessKey": "updated-aws-access-key", - // Embedding model credentials (optional) - "embeddingModelApiKey": "updated-embedding-api-key" -} -``` - -### Response (200 OK) -```json -{ - "id": 1, - "llmPlatform": "Azure AI", - "llmModel": "GPT-4o-mini", - "embeddingPlatform": "Azure AI", - "embeddingModel": "text-embedding-ada-002", - "monthlyBudget": 2000.00, - "usedBudget": 150.75, - "deploymentEnvironment": "Production", - "status": "active", - "createdAt": "2025-09-02T10:15:30.000Z", - // Azure credentials (if provided) - "deploymentName": "updated-deployment", - "targetUri": "https://updated-endpoint.azure.com", - "apiKey": "updated-azure-api-key", - // AWS Bedrock credentials (if provided) - "secretKey": "updated-aws-secret-key", - "accessKey": "updated-aws-access-key", - // Embedding model credentials (if provided) - "embeddingModelApiKey": "updated-embedding-api-key" -} -``` - ---- - -## 3. Get LLM Connections (Paginated List) - -### Endpoint -```http -POST /ruuter-private/rag-search/llm-connections/list -``` - -### Request Body -```json -{ - "page": 1, - "page_size": 10, - "sorting": "created_at desc" -} -``` - -### Request Parameters -| Parameter | Type | Required | Description | Default | -|-----------|------|----------|-------------|---------| -| `page` | number | No | Page number (1-based) | 1 | -| `page_size` | number | No | Number of items per page | 10 | -| `sorting` | string | No | Sorting criteria | "created_at desc" | - -### Sorting Options -- `llm_platform asc/desc` -- `llm_model asc/desc` -- `embedding_platform asc/desc` -- `embedding_model asc/desc` -- `monthly_budget asc/desc` -- `environment asc/desc` -- `status asc/desc` -- `created_at asc/desc` -- `updated_at asc/desc` - -### Response (200 OK) -```json -[ - { - "id": 1, - "llmPlatform": "OpenAI", - "llmModel": "GPT-4o", - "embeddingPlatform": "OpenAI", - "embeddingModel": "text-embedding-3-small", - "monthlyBudget": 1000.00, - "environment": "Testing", - "status": "active", - "createdAt": "2025-09-02T10:15:30.000Z", - "updatedAt": "2025-09-02T10:15:30.000Z", - "totalPages": 3 - }, - { - "id": 2, - "llmPlatform": "Azure AI", - "llmModel": "GPT-4o-mini", - "embeddingPlatform": "Azure AI", - "embeddingModel": "Ada-200-1", - "monthlyBudget": 2000.00, - "environment": "Production", - "status": "active", - "createdAt": "2025-09-02T09:30:15.000Z", - "updatedAt": "2025-09-02T11:00:00.000Z", - "totalPages": 3 - } -] -``` - ---- - -## 4. Get Single LLM Connection - -### Endpoint -```http -POST /ruuter-private/rag-search/llm-connections/get -``` - -### Request Body -```json -{ - "connection_id": 1 -} -``` - -### Response (200 OK) -```json -{ - "id": 1, - "llmPlatform": "OpenAI", - "llmModel": "GPT-4o", - "embeddingPlatform": "OpenAI", - "embeddingModel": "text-embedding-3-small", - "monthlyBudget": 1000.00, - "environment": "Testing", - "status": "active", - "createdAt": "2025-09-02T10:15:30.000Z", - "updatedAt": "2025-09-02T10:15:30.000Z" -} -``` - -### Response (404 Not Found) -```json -"error: connection not found" -``` - ---- - -## 5. Add New LLM Connection - -### Endpoint -```http -POST /ruuter-private/rag-search/llm-connections/add -``` - -### Request Body -```json -{ - "llm_platform": "OpenAI", - "llm_model": "GPT-4o", - "embedding_platform": "OpenAI", - "embedding_model": "text-embedding-3-small", - "monthly_budget": 1000.00, - "environment": "Testing" -} -``` - -### Request Parameters -| Parameter | Type | Required | Description | -|-----------|------|----------|-------------| -| `llm_platform` | string | Yes | LLM platform (e.g., "Azure AI", "OpenAI") | -| `llm_model` | string | Yes | LLM model (e.g., "GPT-4o") | -| `embedding_platform` | string | Yes | Embedding platform | -| `embedding_model` | string | Yes | Embedding model | -| `monthly_budget` | number | Yes | Monthly budget amount | -| `environment` | string | Yes | "Testing" or "Production" | - -### Response (200 OK) -```json -{ - "id": 3, - "llm_platform": "OpenAI", - "llm_model": "GPT-4o", - "embedding_platform": "OpenAI", - "embedding_model": "text-embedding-3-small", - "monthly_budget": 1000.00, - "environment": "Testing", - "status": "active", - "created_at": "2025-09-02T12:00:00.000Z", - "updated_at": "2025-09-02T12:00:00.000Z" -} -``` - -### Response (400 Bad Request) -```json -"error: environment must be 'Testing' or 'Production'" -``` - ---- - -## 6. Update LLM Connection - -### Endpoint -```http -POST /ruuter-private/rag-search/llm-connections/edit -``` - -### Request Body -```json -{ - "connection_id": 1, - "llm_platform": "Azure AI", - "llm_model": "GPT-4o-mini", - "embedding_platform": "Azure AI", - "embedding_model": "Ada-200-1", - "monthly_budget": 2000.00, - "environment": "Production" -} -``` - -### Response (200 OK) -```json -{ - "id": 1, - "llm_platform": "Azure AI", - "llm_model": "GPT-4o-mini", - "embedding_platform": "Azure AI", - "embedding_model": "Ada-200-1", - "monthly_budget": 2000.00, - "environment": "Production", - "status": "active", - "created_at": "2025-09-02T10:15:30.000Z", - "updated_at": "2025-09-02T12:30:00.000Z" -} -``` - -### Response (404 Not Found) -```json -"error: connection not found" -``` - ---- - -## 7. Delete LLM Connection - -### Endpoint -```http -POST /ruuter-private/rag-search/llm-connections/delete -``` - -### Request Body -```json -{ - "connection_id": 1 -} -``` - -### Response (200 OK) -```json -"LLM connection deleted successfully" -``` - -### Response (404 Not Found) -```json -"error: connection not found" -``` - ---- - -## 4. List All LLM Connections - -### Endpoint -```http -GET /ruuter-private/llm/connections/list -``` - -### Query Parameters (Optional for filtering) -| Parameter | Type | Description | -|-----------|------|-------------| -| `llmPlatform` | `string` | Filter by LLM platform | -| `llmModel` | `string` | Filter by LLM model | -| `deploymentEnvironment` | `string` | Filter by environment (Testing / Production) | -| `pageNumber` | `number` | Page number (1-based) | -| `pageSize` | `number` | Number of items per page | -| `sortBy` | `string` | Field to sort by | -| `sortOrder` | `string` | Sort order: 'asc' or 'desc' | - -### Example Request -```http -GET /ruuter-private/llm/connections/list?llmPlatform=OpenAI&deploymentEnvironment=Testing&model=GPT4 -``` - ---- - -## 5. Get Production LLM Connection (with filters) - -### Endpoint -```http -GET /ruuter-private/llm/connections/production -``` - -### Query Parameters (Optional for filtering) -| Parameter | Type | Description | -|-----------|------|-------------| -| `llmPlatform` | `string` | Filter by LLM platform | -| `llmModel` | `string` | Filter by LLM model | -| `embeddingPlatform` | `string` | Filter by embedding platform | -| `embeddingModel` | `string` | Filter by embedding model | -| `connectionStatus` | `string` | Filter by connection status | -| `sortBy` | `string` | Field to sort by | -| `sortOrder` | `string` | Sort order: 'asc' or 'desc' | - -### Example Request -```http -GET /ruuter-private/llm/connections/production?llmPlatform=OpenAI&connectionStatus=active -``` - -### Response (200 OK) -```json -[ - { - "id": 1, - "llmPlatform": "OpenAI", - "llmModel": "GPT-4o", - "embeddingPlatform": "OpenAI", - "embeddingModel": "text-embedding-3-small", - "monthlyBudget": 1000.00, - "deploymentEnvironment": "Testing", - "status": "active", - "createdAt": "2025-09-02T10:15:30.000Z", - "updatedAt": "2025-09-02T10:15:30.000Z" - } -] -``` - ---- - -## 5. Get Single LLM Connection - -### Endpoint -```http -GET /ruuter-private/llm/connections/overview -``` - -### Response (200 OK) -```json -{ - "id": 1, - "llmPlatform": "OpenAI", - "llmModel": "GPT-4o", - "embeddingPlatform": "OpenAI", - "embeddingModel": "text-embedding-3-small", - "monthlyBudget": 1000.00, - "deploymentEnvironment": "Testing", - "status": "active", - "createdAt": "2025-09-02T10:15:30.000Z", - "updatedAt": "2025-09-02T10:15:30.000Z" -} -``` - ---- -# Inference Results API Endpoints - -## Base URL -``` -/ruuter-private/inference/results -``` - ---- - -## 1. Store Test Inference Result - -### Endpoint -```http -POST /ruuter-private/inference/results/test/store -``` - -### Request Body -```json -{ - "llm_connection_id": 1, - "user_question": "What are the benefits of using LLMs?", - "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." -} -``` - -### Request Parameters -| Parameter | Type | Required | Description | -|-----------|------|----------|-------------| -| `llm_connection_id` | number | Yes | ID of the LLM connection | -| `user_question` | string | Yes | User's raw question/input | -| `final_answer` | string | Yes | LLM's final generated answer | - -### Response (200 OK) -```json -{ - "data": { - "id": 10, - "llm_connection_id": 1, - "chat_id": null, - "user_question": "What are the benefits of using LLMs?", - "refined_questions": null, - "conversation_history": null, - "ranked_chunks": null, - "embedding_scores": null, - "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", - "environment": "testing", - "created_at": "2025-09-25T12:15:00.000Z" - }, - "operationSuccess": true, - "statusCode": 200 -} -``` - -### Response (400 Bad Request) -```json -{ - "data": "[]", - "operationSuccess": false, - "statusCode": 400 -} -``` - -### Response (404 Not Found) -```json -"error: LLM connection not found" -``` - ---- - -## 2. Store Production Inference Result - -### Endpoint -```http -POST /ruuter-private/inference/results/production/store -``` - -### Request Body -```json -{ - "chat_id": "chat-12345", - "user_question": "What are the benefits of using LLMs?", - "refined_questions": [ - "How do LLMs improve productivity?", - "What are practical use cases of LLMs?" - ], - "conversation_history": [ - { "role": "user", "content": "Hello" }, - { "role": "assistant", "content": "Hi! How can I help you?" } - ], - "ranked_chunks": [ - { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, - { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } - ], - "embedding_scores": [0.92, 0.85, 0.78], - "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." -} -``` - -### Request Parameters -| Parameter | Type | Required | Description | -|-----------|------|----------|-------------| -| `chat_id` | string | No | Optional chat session ID | -| `user_question` | string | Yes | User's raw question/input | -| `refined_questions` | object | No | List of refined questions (LLM-generated) | -| `conversation_history` | object | No | Prior messages array of {role, content} | -| `ranked_chunks` | object | No | Retrieved chunks ranked with metadata | -| `embedding_scores` | object | No | Distance scores for each chunk | -| `final_answer` | string | Yes | LLM's final generated answer | - -### Response (200 OK) -```json -{ - "data": { - "id": 15, - "llm_connection_id": null, - "chat_id": "chat-12345", - "user_question": "What are the benefits of using LLMs?", - "refined_questions": [ - "How do LLMs improve productivity?", - "What are practical use cases of LLMs?" - ], - "conversation_history": [ - { "role": "user", "content": "Hello" }, - { "role": "assistant", "content": "Hi! How can I help you?" } - ], - "ranked_chunks": [ - { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, - { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } - ], - "embedding_scores": [0.92, 0.85, 0.78], - "final_answer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", - "environment": "production", - "created_at": "2025-09-25T12:15:00.000Z" - }, - "operationSuccess": true, - "statusCode": 200 -} -``` - -### Response (400 Bad Request) -```json -{ - "data": "[]", - "operationSuccess": false, - "statusCode": 400 -} -``` - ---- - -## 3. View/get Inference Result - -### Endpoint -```http -POST /ruuter-private/inference/results/test/store -``` - -### Request Body -```json -{ - "llmConnectionId": 1, - "userQuestion": "What are the benefits of using LLMs?", - "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." -} -``` - -### Response (201 Created) -```json -{ - "data": { - "id": 15, - "llmConnectionId": 1, - "userQuestion": "What are the benefits of using LLMs?", - "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", - "environment": "testing", - "createdAt": "2025-09-25T10:15:30.000Z" - }, - "operationSuccess": true, - "statusCode": 200 -} -``` - -## 4. Inquiry from chatbot to llm orchestration service - -### Endpoint -```http -POST /ruuter-private/inference/results/production/store -``` - -### Request Body -```json -{ - "llmConnectionId": 1, - "chatId": "chat-session-12345", - "userQuestion": "What are the benefits of using LLMs?", - "refinedQuestions": [ - "How do LLMs improve productivity?", - "What are practical use cases of LLMs?" - ], - "conversationHistory": [ - { "role": "user", "content": "Hello" }, - { "role": "assistant", "content": "Hi! How can I help you?" } - ], - "rankedChunks": [ - { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, - { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } - ], - "embeddingScores": { - "chunk_1": 0.92, - "chunk_2": 0.85 - }, - "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation." -} -``` - -### Response (201 Created) -```json -{ - "id": 20, - "llmConnectionId": 1, - "chatId": "chat-session-12345", - "userQuestion": "What are the benefits of using LLMs?", - "refinedQuestions": [ - "How do LLMs improve productivity?", - "What are practical use cases of LLMs?" - ], - "conversationHistory": [ - { "role": "user", "content": "Hello" }, - { "role": "assistant", "content": "Hi! How can I help you?" } - ], - "rankedChunks": [ - { "id": "chunk_1", "content": "LLMs help in summarization", "rank": 1 }, - { "id": "chunk_2", "content": "They improve Q&A systems", "rank": 2 } - ], - "embeddingScores": { - "chunk_1": 0.92, - "chunk_2": 0.85 - }, - "finalAnswer": "LLMs can improve productivity by summarizing large documents, enabling Q&A, and enhancing automation.", - "environment": "production", - "createdAt": "2025-09-25T10:15:30.000Z" -} -``` - ---- \ No newline at end of file diff --git a/generate-changelog.sh b/generate-changelog.sh new file mode 100755 index 00000000..5f35f494 --- /dev/null +++ b/generate-changelog.sh @@ -0,0 +1,105 @@ +#!/bin/bash + +append=${1:-false}; [ "$append" = "true" ] && append=true || append=false + +env_file=$(@%an](https://www.github.com/%an) in [#%h]($REPO_URL/commit/%h)") + +while read -r line; do + pattern="^([^(:]+)\(([^)]+)\): (.*)" + + if [[ $line =~ $pattern ]]; then + type="${BASH_REMATCH[1]}" + scope="${BASH_REMATCH[2]}" + description="${BASH_REMATCH[3]}" + rest_of_line="**$scope**: $description" + else + type="others" + rest_of_line="$line" + fi + + author_link=$(echo "$rest_of_line" | grep -o 'https://www.github.com/[[:alnum:][:space:]]*' | tr -d '[:space:]') + rest_of_line=$(echo "$rest_of_line" | awk -v replacement="$author_link" '{gsub(/https:\/\/www\.github\.com\/[[:alnum:][:space:]]*/, replacement); print}') + + case $type in + "feat") features+=("- $rest_of_line");; + "fix") fixes+=("- $rest_of_line");; + "docs") docs+=("- $rest_of_line");; + "style") styles+=("- $rest_of_line");; + "refactor") refactors+=("- $rest_of_line");; + "test") tests+=("- $rest_of_line");; + "chore") chores+=("- $rest_of_line");; + *) others+=("- $rest_of_line");; + esac +done <<< "$commit_log" + +[[ ${#features[@]} -gt 0 ]] && changelog_content+="\n## Features\n$(printf "%s\n" "${features[@]}")" +[[ ${#fixes[@]} -gt 0 ]] && changelog_content+="\n## Fixes\n$(printf "%s\n" "${fixes[@]}")" +[[ ${#docs[@]} -gt 0 ]] && changelog_content+="\n## Documentation\n$(printf "%s\n" "${docs[@]}")" +[[ ${#styles[@]} -gt 0 ]] && changelog_content+="\n## Style\n$(printf "%s\n" "${styles[@]}")" +[[ ${#refactors[@]} -gt 0 ]] && changelog_content+="\n## Refactor\n$(printf "%s\n" "${refactors[@]}")" +[[ ${#tests[@]} -gt 0 ]] && changelog_content+="\n## Tests\n$(printf "%s\n" "${tests[@]}")" +[[ ${#chores[@]} -gt 0 ]] && changelog_content+="\n## Chores\n$(printf "%s\n" "${chores[@]}")" +[[ ${#others[@]} -gt 0 ]] && changelog_content+="\n## Others\n$(printf "%s\n" "${others[@]}")" + +# Append or overwrite the changelog file based on the append variable +if [ "$append" = "true" ]; then + echo -e "$changelog_content" >> "CHANGELOG.md" +else + echo -e "$changelog_content" > "CHANGELOG.md" +fi diff --git a/generate_presigned_url.py b/generate_presigned_url.py deleted file mode 100644 index 61028beb..00000000 --- a/generate_presigned_url.py +++ /dev/null @@ -1,64 +0,0 @@ -import boto3 -from botocore.client import Config -from typing import List, Dict -from loguru import logger - -# Create S3 client for MinIO -s3_client = boto3.client( - "s3", - endpoint_url="http://minio:9000", # Replace with your MinIO URL - aws_access_key_id="minioadmin", # Replace with your access key - aws_secret_access_key="minioadmin", # Replace with your secret key - config=Config(signature_version="s3v4"), # Hardcoded signature version - region_name="us-east-1", # MinIO usually works with any region -) - -# List of files to process -files_to_process: List[Dict[str, str]] = [ - {"bucket": "ckb", "key": "ID.ee/ID.zip"}, -] - -# Generate presigned URLs -presigned_urls: List[str] = [] - -logger.info("Generating presigned URLs...") -for file_info in files_to_process: - try: - url = s3_client.generate_presigned_url( - ClientMethod="get_object", - Params={"Bucket": file_info["bucket"], "Key": file_info["key"]}, - ExpiresIn=24 * 3600, # 4 hours in seconds - ) - presigned_urls.append(url) - logger.success(f"Generated URL for: {file_info['key']}") - logger.info(f" URL: {url}") - except Exception as e: - logger.error(f"Failed to generate URL for: {file_info['key']}") - logger.error(f" Error: {str(e)}") - -output_file: str = "minio_presigned_urls.txt" - -try: - with open(output_file, "w") as f: - # Write URLs separated by ||| delimiter (for your script) - url_string: str = "|||".join(presigned_urls) - f.write(url_string) - f.write("\n\n") - - # Also write each URL on separate lines for readability - f.write("Individual URLs:\n") - f.write("=" * 50 + "\n") - for i, url in enumerate(presigned_urls, 1): - f.write(f"URL {i}:\n{url}\n\n") - - logger.success(f"Presigned URLs saved to: {output_file}") - logger.info(f"Total URLs generated: {len(presigned_urls)}") - - # Display the combined URL string for easy copying - if presigned_urls: - logger.info("Combined URL string (for signedUrls environment variable):") - logger.info("=" * 60) - logger.info("|||".join(presigned_urls)) - -except Exception as e: - logger.error(f"Failed to save URLs to file: {str(e)}") diff --git a/grafana-configs/README.md b/grafana-configs/README.md index 6feba7bb..dcdb3a9a 100644 --- a/grafana-configs/README.md +++ b/grafana-configs/README.md @@ -87,12 +87,17 @@ from grafana_configs.loki_logger import LokiLogger logger = LokiLogger(service_name="model-deployment-orchestrator") # Log with model context -logger.info("Starting deployment", model_id="model123", - current_env="testing", target_env="production") +logger.info( + "Starting deployment", + model_id="model123", + current_env="testing", + target_env="production", +) # Log errors with extra context -logger.error("Deployment failed", model_id="model123", - error_code=500, step="model_loading") +logger.error( + "Deployment failed", model_id="model123", error_code=500, step="model_loading" +) ``` ### Accessing Grafana diff --git a/grafana-configs/grafana-dashboard-deployment.json b/grafana-configs/grafana-dashboard-deployment.json index a1e469f2..549b8035 100644 --- a/grafana-configs/grafana-dashboard-deployment.json +++ b/grafana-configs/grafana-dashboard-deployment.json @@ -1,7 +1,7 @@ { "id": null, "title": "RAG Module Orchestrator", - "tags": ["deployment", "models", "triton"], + "tags": ["rag-module", "loki-logs", "llm-orchestration", "vector-search"], "timezone": "browser", "refresh": "30s", "time": { @@ -15,7 +15,7 @@ "type": "query", "label": "Service Name", "refresh": 1, - "query": "label_values(service)", + "query": "label_values({}, service)", "datasource": { "type": "loki", "uid": "loki-datasource" diff --git a/grafana-configs/loki_logger.py b/grafana-configs/loki_logger.py index e90dd059..a86a969c 100644 --- a/grafana-configs/loki_logger.py +++ b/grafana-configs/loki_logger.py @@ -1,19 +1,31 @@ #!/usr/bin/env python3 """ -Loki Logger for Global Classifier +Loki Logger for RAG Module Sends logs directly to Loki API for centralized logging """ import json -import socket +import sys import time from datetime import datetime +from threading import Thread +from queue import Full, Queue import requests class LokiLogger: - """Simple logger that sends logs directly to Loki API""" + """Simple logger that sends logs directly to Loki API with async background thread""" + + _instances: dict[str, "LokiLogger"] = {} + + def __new__( + cls, loki_url: str = "http://loki:3100", service_name: str = "default" + ) -> "LokiLogger": + key = f"{loki_url}:{service_name}" + if key not in cls._instances: + cls._instances[key] = super().__new__(cls) + return cls._instances[key] def __init__( self, loki_url: str = "http://loki:3100", service_name: str = "default" @@ -25,15 +37,40 @@ def __init__( loki_url: URL for Loki service (default: container URL in bykstack network) service_name: Name of the service for labeling logs """ + if hasattr(self, "_initialized"): + return + self._initialized = True self.loki_url = loki_url self.service_name = service_name - self.hostname = socket.gethostname() self.session = requests.Session() # Set default timeout for all requests self.timeout = 5 - def _send_to_loki(self, level: str, message: str) -> None: - """Send log entry directly to Loki API""" + # Queue for async log processing (bounded to avoid unbounded memory growth under load) + self.log_queue: Queue[tuple[str, str]] = Queue(maxsize=10_000) + + # Start background worker thread + self.worker_thread = Thread(target=self._process_logs, daemon=True) + self.worker_thread.start() + + def _process_logs(self) -> None: + """Background worker that processes log queue""" + while True: + try: + # Get log entry from queue (blocking) + level, message = self.log_queue.get() + + # Send to Loki + self._send_to_loki_sync(level, message) + + # Mark task as done + self.log_queue.task_done() + except Exception: + # Silently ignore errors in background thread + pass + + def _send_to_loki_sync(self, level: str, message: str) -> None: + """Send log entry directly to Loki API (called from background thread)""" try: # Create timestamp in nanoseconds (Loki requirement) timestamp_ns = str(int(time.time() * 1_000_000_000)) @@ -42,15 +79,12 @@ def _send_to_loki(self, level: str, message: str) -> None: labels = { "service": self.service_name, "level": level, - "hostname": self.hostname, } # Create log entry log_entry = { - "timestamp": datetime.now().isoformat(), "level": level, "message": message, - "hostname": self.hostname, "service": self.service_name, } @@ -64,7 +98,7 @@ def _send_to_loki(self, level: str, message: str) -> None: ] } - # Send to Loki (non-blocking, fire-and-forget) + # Send to Loki self.session.post( f"{self.loki_url}/loki/api/v1/push", json=payload, @@ -76,18 +110,67 @@ def _send_to_loki(self, level: str, message: str) -> None: # Silently ignore logging errors to not affect main application pass - # Also print to console for immediate feedback + def _log(self, level: str, message: str) -> None: + """Queue log entry for async processing (non-blocking)""" + # Print to console immediately for real-time feedback. Written to + # stderr (not stdout) so callers that capture a subprocess's stdout + # for its return value (e.g. decrypt_vault_secrets.py) never pick up + # log lines mixed in with the actual output. timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - print(f"[{timestamp}] {level: <8} | {message}") # noqa: T201 + print(f"[{timestamp}] {level: <8} | {message}", file=sys.stderr) # noqa: T201 + + # Queue for async Loki sending (non-blocking) + try: + self.log_queue.put_nowait((level, message)) + except Full: + # Queue full (Loki may be slow/unreachable) - drop log to avoid blocking + pass - def info(self, message: str) -> None: - self._send_to_loki("INFO", message) + def info(self, message: str, **kwargs: object) -> None: + """Log info message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("INFO", message) + + def error(self, message: str, **kwargs: object) -> None: + """Log error message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("ERROR", message) + + def warning(self, message: str, **kwargs: object) -> None: + """Log warning message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("WARNING", message) + + def debug(self, message: str, **kwargs: object) -> None: + """Log debug message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("DEBUG", message) + + def success(self, message: str, **kwargs: object) -> None: + """Log success message (loguru compatibility). Extra kwargs ignored.""" + self._log("SUCCESS", message) + + def critical(self, message: str, **kwargs: object) -> None: + """Log critical message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("CRITICAL", message) + + def exception(self, message: str, **kwargs: object) -> None: + """Log exception message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("EXCEPTION", message) + + def add(self, *args: object, **kwargs: object) -> None: + """ + No-op method for loguru compatibility. + + LokiLogger sends logs to Loki/console only, not to files. + This method exists for backward compatibility with loguru code. + """ + pass # Silently ignore - logs go to Loki instead of files - def error(self, message: str) -> None: - self._send_to_loki("ERROR", message) + def remove(self, *args: object, **kwargs: object) -> None: + """No-op method for loguru compatibility.""" + pass # Silently ignore - def warning(self, message: str) -> None: - self._send_to_loki("WARNING", message) + def bind(self, **kwargs: object) -> "LokiLogger": + """No-op method for loguru compatibility. Returns self for chaining.""" + return self # Allow method chaining - def debug(self, message: str) -> None: - self._send_to_loki("DEBUG", message) + def opt(self, **kwargs: object) -> "LokiLogger": + """No-op method for loguru compatibility. Returns self for chaining.""" + return self # Allow method chaining diff --git a/kubernetes/Chart.yaml b/kubernetes/Chart.yaml index eb9a316a..ef6d5a2c 100644 --- a/kubernetes/Chart.yaml +++ b/kubernetes/Chart.yaml @@ -113,4 +113,7 @@ dependencies: version: 0.1.0 repository: "file://./charts/Notifications-Node" condition: Notifications-Node.enabled - + - name: OpenSearch + version: 0.1.0 + repository: "file://./charts/OpenSearch" + condition: OpenSearch.enabled \ No newline at end of file diff --git a/kubernetes/LANGFUSE_SETUP.md b/kubernetes/LANGFUSE_SETUP.md index 6c0f11bd..c54d91af 100644 --- a/kubernetes/LANGFUSE_SETUP.md +++ b/kubernetes/LANGFUSE_SETUP.md @@ -51,9 +51,12 @@ kubectl cp store-langfuse-secrets.sh rag-module/vault-0:/tmp/store-langfuse-secr kubectl exec -n your-namespace vault-0 -- sh -c \ "LANGFUSE_INIT_PROJECT_PUBLIC_KEY=pk-lf-YOUR_KEY \ LANGFUSE_INIT_PROJECT_SECRET_KEY=sk-lf-YOUR_KEY \ + LANGFUSE_HOST=http://langfuse-web:3005 \ sh /tmp/store-langfuse-secrets.sh" ``` Replace `pk-lf-YOUR_KEY` and `sk-lf-YOUR_KEY` with the actual keys from step 3. -The script stores them at `secret/data/langfuse/config` in Vault, where the LLM Orchestration Service reads them. +> **Note:** In Kubernetes, the Langfuse-Web service port is `3005` (mapped to container port 3000), so `LANGFUSE_HOST` must be set explicitly. In Docker Compose, the default (`http://langfuse-web:3000`) is used automatically. + +The script stores them at `secret/data/langfuse/config` in Vault, where the LLM Orchestration Service reads them. \ No newline at end of file diff --git a/kubernetes/charts/CronManager/templates/deployment-byk-cronmanager.yaml b/kubernetes/charts/CronManager/templates/deployment-byk-cronmanager.yaml index 15dc9615..bdf53252 100644 --- a/kubernetes/charts/CronManager/templates/deployment-byk-cronmanager.yaml +++ b/kubernetes/charts/CronManager/templates/deployment-byk-cronmanager.yaml @@ -36,6 +36,14 @@ spec: mountPath: /app/scripts - name: vector-indexer mountPath: /app/src/vector_indexer + - name: tool-classifier + mountPath: /app/src/tool_classifier + - name: intent-data-enrichment + mountPath: /app/src/intent_data_enrichment + - name: api-tool-indexer + mountPath: /app/src/api_tool_indexer + - name: src-utils + mountPath: /app/src/utils command: - sh - -c @@ -45,12 +53,20 @@ spec: mkdir -p /app/src/vector_indexer && mkdir -p /app/scripts && mkdir -p /DSL && - mkdir -p /app/src/utils + mkdir -p /app/src/utils && + mkdir -p /app/src/tool_classifier && + mkdir -p /app/src/intent_data_enrichment && + mkdir -p /app/src/api_tool_indexer cp -r /tmp/rag/DSL/CronManager/DSL/* /DSL/ && cp -r /tmp/rag/DSL/CronManager/script/* /app/scripts/ && cp -r /tmp/rag/src/vector_indexer/* /app/src/vector_indexer/ && - cp -r /tmp/rag/src/utils/decrypt_vault_secrets.py /app/src/utils/ && + cp -r /tmp/rag/src/tool_classifier/* /app/src/tool_classifier/ && + cp -r /tmp/rag/src/intent_data_enrichment/* /app/src/intent_data_enrichment/ && + cp -r /tmp/rag/src/api_tool_indexer/* /app/src/api_tool_indexer/ && + cp /tmp/rag/src/utils/decrypt_vault_secrets.py /app/src/utils/ && + cp /tmp/rag/src/__init__.py /app/src/__init__.py && + cp /tmp/rag/grafana-configs/loki_logger.py /app/src/vector_indexer/loki_logger.py && # Set execute permissions on all shell scripts chmod +x /app/scripts/*.sh && @@ -91,7 +107,7 @@ spec: value: {{ .Values.cronmanager.environment.pythonPath | quote }} {{- if .Values.vaultAgent.enabled }} # Vault Agent proxy URL (localhost sidecar) - - name: VAULT_AGENT_URL + - name: vaultAgentUrl value: "http://localhost:8203" {{- end }} - name: RAG_MODULE_RUUTER_PRIVATE @@ -112,6 +128,14 @@ spec: mountPath: /app/scripts - name: vector-indexer mountPath: /app/src/vector_indexer + - name: tool-classifier + mountPath: /app/src/tool_classifier + - name: intent-data-enrichment + mountPath: /app/src/intent_data_enrichment + - name: api-tool-indexer + mountPath: /app/src/api_tool_indexer + - name: src-utils + mountPath: /app/src/utils - name: datasets mountPath: /app/datasets @@ -122,8 +146,16 @@ spec: emptyDir: {} - name: vector-indexer emptyDir: {} + - name: tool-classifier + emptyDir: {} + - name: intent-data-enrichment + emptyDir: {} + - name: api-tool-indexer + emptyDir: {} - name: datasets emptyDir: {} + - name: src-utils + emptyDir: {} - name: cronmanager-data persistentVolumeClaim: claimName: "{{ .Values.release_name }}-data" diff --git a/kubernetes/charts/CronManager/values.yaml b/kubernetes/charts/CronManager/values.yaml index df8013cf..4db03e3e 100644 --- a/kubernetes/charts/CronManager/values.yaml +++ b/kubernetes/charts/CronManager/values.yaml @@ -11,7 +11,7 @@ cronmanager: environment: containerPort: "8080" - pythonPath: "/app:/app/src/vector_indexer" + pythonPath: "/app:/app/src:/app/src/vector_indexer:/app/src/intent_data_enrichment:/app/src/api_tool_indexer" VAULT_ADDR: "http://vault:8200" service: diff --git a/kubernetes/charts/GUI/templates/configmap-vite-config.yaml b/kubernetes/charts/GUI/templates/configmap-vite-config.yaml index 7110554b..fe4e29cc 100644 --- a/kubernetes/charts/GUI/templates/configmap-vite-config.yaml +++ b/kubernetes/charts/GUI/templates/configmap-vite-config.yaml @@ -45,6 +45,21 @@ data: 'Content-Security-Policy': process.env.REACT_APP_CSP, }), }, + proxy: { + '/vault-agent-gui': { + target: 'http://localhost:8202', + changeOrigin: true, + rewrite: (path) => path.replace(/^\/vault-agent-gui/, ''), + }, + '/sse': { + target: 'http://notifications-node:4040', + changeOrigin: true, + }, + '/channels': { + target: 'http://notifications-node:4040', + changeOrigin: true, + }, + }, }, resolve: { alias: { @@ -53,4 +68,4 @@ data: }, }, }); -{{- end }} +{{- end }} \ No newline at end of file diff --git a/kubernetes/charts/GUI/templates/ingress-byk-gui.yaml b/kubernetes/charts/GUI/templates/ingress-byk-gui.yaml index cda59a70..15deedde 100644 --- a/kubernetes/charts/GUI/templates/ingress-byk-gui.yaml +++ b/kubernetes/charts/GUI/templates/ingress-byk-gui.yaml @@ -18,3 +18,10 @@ spec: name: gui port: number: 3001 + - path: /vault-agent-gui + pathType: Prefix + backend: + service: + name: gui + port: + number: 3001 diff --git a/kubernetes/charts/GUI/values.yaml b/kubernetes/charts/GUI/values.yaml index 03257659..50c1ac9a 100644 --- a/kubernetes/charts/GUI/values.yaml +++ b/kubernetes/charts/GUI/values.yaml @@ -14,10 +14,10 @@ gui: #service URLs services: - ruuterPublic: "http:///ruuter-public" - ruuterPrivate: "http:///ruuter-private" - authenticationLayer: "http://" - notificationNode: "http://notifications-node:4040" + ruuterPublic: "http://localhost:8086" + ruuterPrivate: "http://localhost:8088" + authenticationLayer: "http://localhost:3004" + notificationNode: "http://localhost:3003" datasetGenerator: "http://dataset-gen-service:8000" # Content Security Policy - Updated for browser access @@ -33,7 +33,7 @@ gui: # Ingress host ingress: - host: "" # Update with actual domain + host: "localhost" # Update with actual domain resources: limits: @@ -52,20 +52,4 @@ gui: # Vault Agent sidecar configuration vaultAgent: - enabled: true - - - # ingress: - # enabled: true - # className: nginx - # annotations: - # nginx.ingress.kubernetes.io/rewrite-target: / - # nginx.ingress.kubernetes.io/proxy-read-timeout: "3600" - # nginx.ingress.kubernetes.io/proxy-send-timeout: "3600" - # nginx.ingress.kubernetes.io/proxy-body-size: "50m" - # hosts: - # - host: rag.local - # paths: - # - path: / - # pathType: Prefix - # tls: [] \ No newline at end of file + enabled: true \ No newline at end of file diff --git a/kubernetes/charts/LLM-Orchestration-Service/templates/deployment-byk-llm-orchestration.yaml b/kubernetes/charts/LLM-Orchestration-Service/templates/deployment-byk-llm-orchestration.yaml index 3e0feacb..fdfca98b 100644 --- a/kubernetes/charts/LLM-Orchestration-Service/templates/deployment-byk-llm-orchestration.yaml +++ b/kubernetes/charts/LLM-Orchestration-Service/templates/deployment-byk-llm-orchestration.yaml @@ -21,6 +21,7 @@ spec: initContainers: - name: volume-init image: "{{ .Values.initContainer.image.repository }}:{{ .Values.initContainer.image.tag }}" + imagePullPolicy: {{ .Values.initContainer.image.pullPolicy }} command: - sh - -c @@ -146,6 +147,11 @@ spec: - name: logs-volume mountPath: {{ .Values.volumes.logs.mountPath }} {{- end }} + {{- if .Values.vaultAgent.enabled }} + - name: vault-agent-llm-token + mountPath: /agent/llm-token + readOnly: true + {{- end }} resources: requests: diff --git a/kubernetes/charts/LLM-Orchestration-Service/values.yaml b/kubernetes/charts/LLM-Orchestration-Service/values.yaml index 7e457d22..6db9bcfe 100644 --- a/kubernetes/charts/LLM-Orchestration-Service/values.yaml +++ b/kubernetes/charts/LLM-Orchestration-Service/values.yaml @@ -50,6 +50,7 @@ initContainer: image: repository: "ghcr.io/buerokratt/llm-orchestration-service" # Update with actual llm-orchestration image repository tag: "latest" + pullPolicy: "IfNotPresent" # InitContainer will prepare the runtime volumes prepareVolumes: true @@ -73,9 +74,22 @@ healthcheck: # Additional readiness checks readinessPath: "/ready" +# Environment variables injected into the LLM container +# Redis defaults match the in-cluster Redis service (see Redis chart) +env: + REDIS_HOST: "redis" + REDIS_PORT: "6379" + REDIS_AUTH: "myredissecret" + REDIS_SESSION_DB: "0" + VAULT_AGENT_PROXY: "true" + TOOL_CLASSIFIER_ENABLED: "true" + SERVICE_WORKFLOW_ENABLED: "true" + API_TOOL_CALLING_WORKFLOW_ENABLED: "true" + CONTEXT_WORKFLOW_ENABLED: "true" + MULTI_INTENT_ENABLED: "true" + # Vault Agent sidecar configuration # WHY: LLM Orchestration needs read access to encrypted LLM API keys # Security: Agent enforces policy - read-only access to LLM secrets vaultAgent: - enabled: true - + enabled: true \ No newline at end of file diff --git a/kubernetes/charts/Langfuse-Web/templates/deployment-byk-langfuse-web.yaml b/kubernetes/charts/Langfuse-Web/templates/deployment-byk-langfuse-web.yaml index 18d14804..f7e3e054 100644 --- a/kubernetes/charts/Langfuse-Web/templates/deployment-byk-langfuse-web.yaml +++ b/kubernetes/charts/Langfuse-Web/templates/deployment-byk-langfuse-web.yaml @@ -50,6 +50,10 @@ spec: - name: http containerPort: {{ .Values.service.targetPort }} protocol: TCP + {{- if .Values.envFrom }} + envFrom: + {{- toYaml .Values.envFrom | nindent 12 }} + {{- end }} env: {{- range $key, $value := .Values.env }} - name: {{ $key }} diff --git a/kubernetes/charts/Langfuse-Web/values.yaml b/kubernetes/charts/Langfuse-Web/values.yaml index 6dfaf1cf..ce4da5b7 100644 --- a/kubernetes/charts/Langfuse-Web/values.yaml +++ b/kubernetes/charts/Langfuse-Web/values.yaml @@ -17,6 +17,7 @@ service: # Environment variables env: # Non-sensitive configuration + HOSTNAME: "0.0.0.0" NEXTAUTH_URL: "http://localhost:3000" TELEMETRY_ENABLED: "true" LANGFUSE_ENABLE_EXPERIMENTAL_FEATURES: "true" @@ -53,9 +54,9 @@ env: REDIS_HOST: "redis" REDIS_PORT: "6379" REDIS_TLS_ENABLED: "false" - REDIS_TLS_CA: "" - REDIS_TLS_CERT: "" - REDIS_TLS_KEY: "" + REDIS_TLS_CA: "/certs/ca.crt" + REDIS_TLS_CERT: "/certs/redis.crt" + REDIS_TLS_KEY: "/certs/redis.key" # Email configuration EMAIL_FROM_ADDRESS: "" @@ -90,7 +91,7 @@ resources: pullPolicy: IfNotPresent healthcheck: - enabled: true + enabled: false initialDelaySeconds: 60 periodSeconds: 30 timeoutSeconds: 10 diff --git a/kubernetes/charts/Langfuse-Worker/values.yaml b/kubernetes/charts/Langfuse-Worker/values.yaml index 0a7343eb..b02bfe4a 100644 --- a/kubernetes/charts/Langfuse-Worker/values.yaml +++ b/kubernetes/charts/Langfuse-Worker/values.yaml @@ -16,6 +16,7 @@ service: # Environment variables env: # Non-sensitive configuration + HOSTNAME: "0.0.0.0" NEXTAUTH_URL: "http://localhost:3000" TELEMETRY_ENABLED: "true" LANGFUSE_ENABLE_EXPERIMENTAL_FEATURES: "true" @@ -52,9 +53,9 @@ env: REDIS_HOST: "redis" REDIS_PORT: "6379" REDIS_TLS_ENABLED: "false" - REDIS_TLS_CA: "" - REDIS_TLS_CERT: "" - REDIS_TLS_KEY: "" + REDIS_TLS_CA: "/certs/ca.crt" + REDIS_TLS_CERT: "/certs/redis.crt" + REDIS_TLS_KEY: "/certs/redis.key" # Email configuration EMAIL_FROM_ADDRESS: "" @@ -77,7 +78,7 @@ resources: pullPolicy: IfNotPresent healthcheck: - enabled: true + enabled: false initialDelaySeconds: 60 periodSeconds: 30 timeoutSeconds: 10 diff --git a/kubernetes/charts/Loki/templates/deployment-loki.yaml b/kubernetes/charts/Loki/templates/deployment-loki.yaml index 7967b8a3..9bf85dd3 100644 --- a/kubernetes/charts/Loki/templates/deployment-loki.yaml +++ b/kubernetes/charts/Loki/templates/deployment-loki.yaml @@ -25,6 +25,7 @@ spec: volumeMounts: - name: config mountPath: /etc/loki/local-config.yaml + subPath: loki.yaml {{- if .Values.persistence.enabled }} - name: storage mountPath: /loki diff --git a/kubernetes/charts/Notifications-Node/Chart.yaml b/kubernetes/charts/Notifications-Node/Chart.yaml new file mode 100644 index 00000000..a7756889 --- /dev/null +++ b/kubernetes/charts/Notifications-Node/Chart.yaml @@ -0,0 +1,6 @@ +apiVersion: v2 +name: Notifications-Node +description: A Helm chart for Notifications server +type: application +version: 0.1.0 +appVersion: "1.0" \ No newline at end of file diff --git a/kubernetes/charts/Notifications-Node/templates/deployment-byk-notifications.yaml b/kubernetes/charts/Notifications-Node/templates/deployment-byk-notifications.yaml new file mode 100644 index 00000000..81141337 --- /dev/null +++ b/kubernetes/charts/Notifications-Node/templates/deployment-byk-notifications.yaml @@ -0,0 +1,69 @@ +apiVersion: apps/v1 +kind: Deployment +metadata: + name: {{ .Values.release_name }} + labels: + app: {{ .Values.release_name }} +spec: + replicas: {{ .Values.replicas }} + selector: + matchLabels: + app: {{ .Values.release_name }} + template: + metadata: + labels: + app: {{ .Values.release_name }} + spec: + containers: + - name: {{ .Values.release_name }} + image: "{{ .Values.notifications.image.repository }}:{{ .Values.notifications.image.tag }}" + imagePullPolicy: {{ .Values.notifications.image.pullPolicy }} + ports: + - containerPort: {{ .Values.notifications.port }} + protocol: TCP + env: + # Node.js application configuration + - name: NODE_ENV + value: {{ .Values.notifications.nodeEnv | quote }} + - name: PORT + value: {{ .Values.notifications.port | quote }} + - name: REFRESH_INTERVAL + value: {{ .Values.notifications.refreshInterval | quote }} + + # OpenSearch configuration + {{- if .Values.notifications.opensearch.enabled }} + - name: OPENSEARCH_PROTOCOL + value: {{ .Values.notifications.opensearch.protocol | quote }} + - name: OPENSEARCH_HOST + value: {{ .Values.notifications.opensearch.host | quote }} + - name: OPENSEARCH_PORT + value: {{ .Values.notifications.opensearch.port | quote }} + - name: OPENSEARCH_USERNAME + valueFrom: + secretKeyRef: + name: notifications-env-secret + key: OPENSEARCH_USERNAME + - name: OPENSEARCH_PASSWORD + valueFrom: + secretKeyRef: + name: notifications-env-secret + key: OPENSEARCH_PASSWORD + {{- end }} + + # CORS configuration + - name: CORS_WHITELIST_ORIGINS + value: {{ .Values.notifications.cors.whitelistOrigins | quote }} + + # BYK Stack integration + - name: RUUTER_URL + value: {{ .Values.notifications.services.ruuterUrl | quote }} + + resources: + limits: + cpu: {{ .Values.notifications.resources.limits.cpu }} + memory: {{ .Values.notifications.resources.limits.memory }} + requests: + cpu: {{ .Values.notifications.resources.requests.cpu }} + memory: {{ .Values.notifications.resources.requests.memory }} + + restartPolicy: Always \ No newline at end of file diff --git a/kubernetes/charts/Notifications-Node/templates/secret-byk-notifications.yaml b/kubernetes/charts/Notifications-Node/templates/secret-byk-notifications.yaml new file mode 100644 index 00000000..5119a1c7 --- /dev/null +++ b/kubernetes/charts/Notifications-Node/templates/secret-byk-notifications.yaml @@ -0,0 +1,13 @@ +{{- if .Values.notifications.opensearch.enabled }} +apiVersion: v1 +kind: Secret +metadata: + name: notifications-env-secret + labels: + app: {{ .Values.release_name }} + +type: Opaque +data: + OPENSEARCH_USERNAME: {{ .Values.notifications.opensearch.username | b64enc }} + OPENSEARCH_PASSWORD: {{ .Values.notifications.opensearch.password | b64enc }} +{{- end }} \ No newline at end of file diff --git a/kubernetes/charts/Notifications-Node/templates/service-byk-notifications.yaml b/kubernetes/charts/Notifications-Node/templates/service-byk-notifications.yaml new file mode 100644 index 00000000..fbd8da49 --- /dev/null +++ b/kubernetes/charts/Notifications-Node/templates/service-byk-notifications.yaml @@ -0,0 +1,16 @@ +apiVersion: v1 +kind: Service +metadata: + name: {{ .Values.release_name }} + labels: + app: {{ .Values.release_name }} +spec: + type: {{ .Values.notifications.service.type }} + ports: + - port: {{ .Values.notifications.service.port }} + targetPort: {{ .Values.notifications.service.targetPort }} + protocol: TCP + name: http + selector: + app: {{ .Values.release_name }} + \ No newline at end of file diff --git a/kubernetes/charts/Notifications-Node/values.yaml b/kubernetes/charts/Notifications-Node/values.yaml new file mode 100644 index 00000000..9ad67631 --- /dev/null +++ b/kubernetes/charts/Notifications-Node/values.yaml @@ -0,0 +1,52 @@ +replicas: 1 + +podAnnotations: {} +podSecurityContext: {} +securityContext: {} + +release_name: "notifications-node" + +notifications: + image: + repository: public.ecr.aws/e7g9l0j0/rag-module/notification-server + tag: latest + pullPolicy: IfNotPresent + + # Node.js application configuration + port: 4040 + refreshInterval: 1000 + nodeEnv: production + + # OpenSearch configuration + opensearch: + enabled: true + protocol: http + host: opensearch-node + port: 9200 + username: admin + password: admin + + # CORS configuration for frontend access + cors: + whitelistOrigins: "http://gui:3001,http://gui:3002,http://gui:3003,http://authentication-layer:3004,http://ruuter-public:8086,http://ruuter-private:8088" + + # BYK Stack integration + services: + ruuterUrl: "http://ruuter-public:8086" + + resources: + limits: + cpu: 500m + memory: 512Mi + requests: + cpu: 100m + memory: 128Mi + + service: + type: ClusterIP + port: 4040 + targetPort: 4040 + + # Security configuration + security: + csrfEnabled: true \ No newline at end of file diff --git a/kubernetes/charts/OpenSearch/Chart.yaml b/kubernetes/charts/OpenSearch/Chart.yaml new file mode 100644 index 00000000..58322f7c --- /dev/null +++ b/kubernetes/charts/OpenSearch/Chart.yaml @@ -0,0 +1,6 @@ +apiVersion: v2 +name: OpenSearch +description: A Helm chart for OpenSearch search and analytics engine +type: application +version: 0.1.0 +appVersion: "2.11.1" \ No newline at end of file diff --git a/kubernetes/charts/OpenSearch/templates/deployment-byk-opensearch.yaml b/kubernetes/charts/OpenSearch/templates/deployment-byk-opensearch.yaml new file mode 100644 index 00000000..48239ee6 --- /dev/null +++ b/kubernetes/charts/OpenSearch/templates/deployment-byk-opensearch.yaml @@ -0,0 +1,77 @@ +apiVersion: apps/v1 +kind: Deployment +metadata: + name: {{ .Values.release_name }} + labels: + app: {{ .Values.release_name }} +spec: + replicas: {{ .Values.replicas }} + selector: + matchLabels: + app: {{ .Values.release_name }} + template: + metadata: + labels: + app: {{ .Values.release_name }} + spec: + containers: + - name: {{ .Values.release_name }} + image: "{{ .Values.opensearch.image.repository }}:{{ .Values.opensearch.image.tag }}" + imagePullPolicy: {{ .Values.opensearch.image.pullPolicy }} + securityContext: + capabilities: + add: ["IPC_LOCK", "SYS_RESOURCE"] + ports: + - containerPort: {{ .Values.opensearch.ports.api }} + name: api + protocol: TCP + - containerPort: {{ .Values.opensearch.ports.performance }} + name: performance + protocol: TCP + env: + # Cluster configuration + - name: node.name + value: {{ .Values.opensearch.cluster.nodeName | quote }} + - name: cluster.name + value: {{ .Values.opensearch.cluster.name | quote }} + - name: discovery.type + value: {{ .Values.opensearch.cluster.discoveryType | quote }} + - name: discovery.seed_hosts + value: {{ .Values.opensearch.cluster.seed_hosts | quote }} + + # Java memory configuration + - name: OPENSEARCH_JAVA_OPTS + value: {{ .Values.opensearch.javaOpts | quote }} + + # Performance configuration + - name: bootstrap.memory_lock + value: {{ .Values.opensearch.bootstrapMemoryLock | quote }} + + # Security configuration + {{- if not .Values.opensearch.security.enabled }} + - name: plugins.security.disabled + value: "true" + {{- end }} + + {{- if .Values.opensearch.persistence.enabled }} + volumeMounts: + - name: opensearch-data + mountPath: {{ .Values.opensearch.persistence.mountPath }} + {{- end }} + + resources: + limits: + cpu: {{ .Values.opensearch.resources.limits.cpu }} + memory: {{ .Values.opensearch.resources.limits.memory }} + requests: + cpu: {{ .Values.opensearch.resources.requests.cpu }} + memory: {{ .Values.opensearch.resources.requests.memory }} + + {{- if .Values.opensearch.persistence.enabled }} + volumes: + - name: opensearch-data + persistentVolumeClaim: + claimName: {{ .Values.release_name }}-data-pvc + {{- end }} + + restartPolicy: Always \ No newline at end of file diff --git a/kubernetes/charts/OpenSearch/templates/ingress.yaml b/kubernetes/charts/OpenSearch/templates/ingress.yaml new file mode 100644 index 00000000..201079f1 --- /dev/null +++ b/kubernetes/charts/OpenSearch/templates/ingress.yaml @@ -0,0 +1,43 @@ +{{- if .Values.opensearch.ingress.enabled -}} +apiVersion: networking.k8s.io/v1 +kind: Ingress +metadata: + name: {{ .Values.release_name }}-ingress + labels: + app: {{ .Values.release_name }} + {{- with .Values.opensearch.ingress.annotations }} + annotations: + {{- toYaml . | nindent 4 }} + {{- end }} +spec: + {{- if .Values.opensearch.ingress.className }} + ingressClassName: {{ .Values.opensearch.ingress.className }} + {{- end }} + {{- if .Values.opensearch.ingress.tls }} + tls: + {{- range .Values.opensearch.ingress.tls }} + - hosts: + {{- range .hosts }} + - {{ . | quote }} + {{- end }} + secretName: {{ .secretName }} + {{- end }} + {{- end }} + rules: + {{- range .Values.opensearch.ingress.hosts }} + - host: {{ .host | quote }} + http: + paths: + {{- range .paths }} + - path: {{ .path }} + {{- if .pathType }} + pathType: {{ .pathType }} + {{- end }} + backend: + service: + name: {{ $.Values.release_name }} + port: + number: {{ $.Values.opensearch.service.apiPort }} + {{- end }} + {{- end }} +{{- end }} \ No newline at end of file diff --git a/kubernetes/charts/OpenSearch/templates/pvc-opensearch.yaml b/kubernetes/charts/OpenSearch/templates/pvc-opensearch.yaml new file mode 100644 index 00000000..09417ffc --- /dev/null +++ b/kubernetes/charts/OpenSearch/templates/pvc-opensearch.yaml @@ -0,0 +1,22 @@ +{{- if .Values.opensearch.persistence.enabled }} +apiVersion: v1 +kind: PersistentVolumeClaim +metadata: + name: {{ .Values.release_name }}-data-pvc + namespace: {{ .Release.Namespace }} + labels: + app: {{ .Values.release_name }} +spec: + accessModes: + - {{ .Values.opensearch.persistence.accessMode }} + resources: + requests: + storage: {{ .Values.opensearch.persistence.size }} + {{- if .Values.opensearch.persistence.storageClass }} + {{- if (eq "-" .Values.opensearch.persistence.storageClass) }} + storageClassName: "" + {{- else }} + storageClassName: {{ .Values.opensearch.persistence.storageClass | quote }} + {{- end }} + {{- end }} +{{- end }} \ No newline at end of file diff --git a/kubernetes/charts/OpenSearch/templates/service-byk-opensearch.yaml b/kubernetes/charts/OpenSearch/templates/service-byk-opensearch.yaml new file mode 100644 index 00000000..cf40517d --- /dev/null +++ b/kubernetes/charts/OpenSearch/templates/service-byk-opensearch.yaml @@ -0,0 +1,19 @@ +apiVersion: v1 +kind: Service +metadata: + name: {{ .Values.release_name }} + labels: + app: {{ .Values.release_name }} +spec: + type: {{ .Values.opensearch.service.type }} + ports: + - port: {{ .Values.opensearch.service.api }} + targetPort: {{ .Values.opensearch.ports.api }} + protocol: TCP + name: api + - port: {{ .Values.opensearch.service.performance }} + targetPort: {{ .Values.opensearch.ports.performance }} + protocol: TCP + name: performance + selector: + app: {{ .Values.release_name }} \ No newline at end of file diff --git a/kubernetes/charts/OpenSearch/values.yaml b/kubernetes/charts/OpenSearch/values.yaml new file mode 100644 index 00000000..185806c7 --- /dev/null +++ b/kubernetes/charts/OpenSearch/values.yaml @@ -0,0 +1,80 @@ +replicas: 1 + +podAnnotations: {} +podSecurityContext: {} +securityContext: {} + +release_name: "opensearch-node" +dashboards_release_name: "opensearch-dashboards" + +opensearch: + image: + repository: opensearchproject/opensearch + tag: "2.11.1" + pullPolicy: IfNotPresent + + + cluster: + name: "opensearch-cluster" + nodeName: "opensearch-node" + discoveryType: "single-node" + seed_hosts: "opensearch" + + # Java memory configuration + javaOpts: "-Xms512m -Xmx512m" + + # Security configuration + security: + enabled: false + + # Performance settings + bootstrapMemoryLock: false + + # Ports configuration for pod + ports: + api: 9200 + performance: 9600 + + # Persistent storage configuration + persistence: + enabled: true + size: 10Gi + storageClass: "" + accessMode: ReadWriteOnce + mountPath: /usr/share/opensearch/data + + resources: + limits: + cpu: 1000m + memory: 2Gi + requests: + cpu: 200m + memory: 1Gi + + # Service configuration + service: + type: ClusterIP + api: 9200 + performance: 9600 + + # Ingress configuration + ingress: + enabled: false + className: nginx + annotations: + nginx.ingress.kubernetes.io/rewrite-target: / + hosts: + - host: opensearch.global-classifier.local + paths: + - path: / + pathType: Prefix + tls: [] + + # Ulimits (required for OpenSearch memory locking) + ulimits: + memlock: + soft: -1 + hard: -1 + nofile: + soft: 65536 + hard: 65536 \ No newline at end of file diff --git a/kubernetes/charts/Ruuter-Private/templates/configmap-byk-ruuter-private.yaml b/kubernetes/charts/Ruuter-Private/templates/configmap-byk-ruuter-private.yaml index 6f84c283..6b158fb9 100644 --- a/kubernetes/charts/Ruuter-Private/templates/configmap-byk-ruuter-private.yaml +++ b/kubernetes/charts/Ruuter-Private/templates/configmap-byk-ruuter-private.yaml @@ -7,12 +7,13 @@ metadata: data: constants.ini: | [DSL] - RAG_SEARCH_RUUTER_PUBLIC=http://ruuter-public:8086/rag-search - RAG_SEARCH_RUUTER_PRIVATE=http://ruuter-private:8088/rag-search - RAG_SEARCH_DMAPPER=http://data-mapper:3000 + RAG_SEARCH_RUUTER_PUBLIC=http://ruuter-public:8086 + RAG_SEARCH_RUUTER_PRIVATE=http://ruuter-private:8088 + RAG_SEARCH_DMAPPER=http://data-mapper:3001 RAG_SEARCH_RESQL=http://resql:8082/rag-search RAG_SEARCH_PROJECT_LAYER=rag-search RAG_SEARCH_TIM=http://tim:8085 RAG_SEARCH_CRON_MANAGER=http://cron-manager:9010 RAG_SEARCH_LLM_ORCHESTRATOR=http://llm-orchestration-service:8100/orchestrate - DOMAIN=localhost \ No newline at end of file + DOMAIN=localhost + DB_PASSWORD=dbadmin \ No newline at end of file diff --git a/kubernetes/charts/Ruuter-Private/values.yaml b/kubernetes/charts/Ruuter-Private/values.yaml index 3a48ac17..f3d5a644 100644 --- a/kubernetes/charts/Ruuter-Private/values.yaml +++ b/kubernetes/charts/Ruuter-Private/values.yaml @@ -6,7 +6,7 @@ images: scope: registry: "ghcr.io" repository: "buerokratt/ruuter" - tag: "v2.2.1" + tag: "v2.2.8" service: type: ClusterIP @@ -15,7 +15,7 @@ service: env: - APPLICATION_CORS_ALLOWEDORIGINS: "http://gui:3001,http://ruuter-private:8088,http://ruuter-public:8086,http://authentication-layer:3004,http://notifications-node:4040,http://dataset-gen-service:8000,http://localhost:3001" + APPLICATION_CORS_ALLOWEDORIGINS: "https://ec2-34-253-140-113.eu-west-1.compute.amazonaws.com:32017,http://gui:3001,http://ruuter-private:8088,http://ruuter-public:8086,http://authentication-layer:3004,http://notifications-node:4040,http://dataset-gen-service:8000,http://localhost:3001" APPLICATION_HTTPCODESALLOWLIST: "200,201,202,400,401,403,500" APPLICATION_INTERNALREQUESTS_ALLOWEDIPS: "127.0.0.1" APPLICATION_LOGGING_DISPLAYREQUESTCONTENT: "true" @@ -42,9 +42,9 @@ resources: ingress: - enabled: false - host: "rag.local" #change this to domain - corsAllowOrigin: "http://localhost:3001,http://localhost:3003,http://localhost:8088,http://localhost:3002,http://localhost:3004,http://localhost:8000" + enabled: true + host: "ec2-34-253-140-113.eu-west-1.compute.amazonaws.com" #change this to domain + corsAllowOrigin: "https://ec2-34-253-140-113.eu-west-1.compute.amazonaws.com:32017,http://localhost:3001,http://localhost:3003,http://localhost:8088,http://localhost:3002,http://localhost:3004,http://localhost:8000" ssl: enabled: false certIssuerName: "letsencrypt-prod" @@ -55,4 +55,3 @@ pullPolicy: IfNotPresent podAnnotations: dsl-checksum: "initial" - diff --git a/kubernetes/charts/Ruuter-Public/templates/configmap-byk-ruuter-public.yaml b/kubernetes/charts/Ruuter-Public/templates/configmap-byk-ruuter-public.yaml index a6a56c0c..67b08aae 100644 --- a/kubernetes/charts/Ruuter-Public/templates/configmap-byk-ruuter-public.yaml +++ b/kubernetes/charts/Ruuter-Public/templates/configmap-byk-ruuter-public.yaml @@ -7,12 +7,13 @@ metadata: data: constants.ini: | [DSL] - RAG_SEARCH_RUUTER_PUBLIC=http://ruuter-public:8086/rag-search - RAG_SEARCH_RUUTER_PRIVATE=http://ruuter-private:8088/rag-search - RAG_SEARCH_DMAPPER=http://data-mapper:3000 + RAG_SEARCH_RUUTER_PUBLIC=http://ruuter-public:8086 + RAG_SEARCH_RUUTER_PRIVATE=http://ruuter-private:8088 + RAG_SEARCH_DMAPPER=http://data-mapper:3001 RAG_SEARCH_RESQL=http://resql:8082/rag-search RAG_SEARCH_PROJECT_LAYER=rag-search RAG_SEARCH_TIM=http://tim:8085 RAG_SEARCH_CRON_MANAGER=http://cron-manager:9010 RAG_SEARCH_LLM_ORCHESTRATOR=http://llm-orchestration-service:8100/orchestrate - DOMAIN=localhost \ No newline at end of file + DOMAIN=localhost + DB_PASSWORD=dbadmin \ No newline at end of file diff --git a/kubernetes/charts/Ruuter-Public/values.yaml b/kubernetes/charts/Ruuter-Public/values.yaml index 320d43f6..22f27d89 100644 --- a/kubernetes/charts/Ruuter-Public/values.yaml +++ b/kubernetes/charts/Ruuter-Public/values.yaml @@ -6,7 +6,7 @@ images: scope: registry: "ghcr.io" repository: "buerokratt/ruuter" - tag: v2.2.1 + tag: "v2.2.8" service: type: ClusterIP @@ -14,7 +14,7 @@ service: targetPort: 8086 env: - APPLICATION_CORS_ALLOWEDORIGINS: "http://localhost:8086,http://localhost:3001,http://localhost:3003,http://localhost:3004,http://localhost:8080,http://localhost:8000,http://localhost:8090" + APPLICATION_CORS_ALLOWEDORIGINS: "https://ec2-34-253-140-113.eu-west-1.compute.amazonaws.com:32017,http://localhost:8086,http://localhost:3001,http://localhost:3003,http://localhost:3004,http://localhost:8080,http://localhost:8000,http://localhost:8090" APPLICATION_HTTPCODESALLOWLIST: "200,201,202,204,400,401,403,500" APPLICATION_INTERNALREQUESTS_ALLOWEDIPS: "127.0.0.1" APPLICATION_LOGGING_DISPLAYREQUESTCONTENT: "true" @@ -40,8 +40,8 @@ resources: ingress: enabled: true - host: "rag.local" # Change this to domain - corsAllowOrigin: "http://localhost:8086,http://localhost:3001,http://localhost:3003,http://localhost:3004,http://localhost:8080,http://localhost:8000,http://localhost:8090" + host: "ec2-34-253-140-113.eu-west-1.compute.amazonaws.com" # EC2 domain + corsAllowOrigin: "https://ec2-34-253-140-113.eu-west-1.compute.amazonaws.com:32017,http://localhost:8086,http://localhost:3001,http://localhost:3003,http://localhost:3004,http://localhost:8080,http://localhost:8000,http://localhost:8090" ssl: enabled: false # Set to true for production with proper certificates certIssuerName: "letsencrypt-prod" @@ -51,4 +51,4 @@ ingress: pullPolicy: IfNotPresent podAnnotations: - dsl-checksum: "94b84bb5ff4d" + dsl-checksum: "ac610bf9ecc0" \ No newline at end of file diff --git a/kubernetes/charts/Vault-Agent-Cron/templates/configmap.yaml b/kubernetes/charts/Vault-Agent-Cron/templates/configmap.yaml index 37a7af6e..4bebfa76 100644 --- a/kubernetes/charts/Vault-Agent-Cron/templates/configmap.yaml +++ b/kubernetes/charts/Vault-Agent-Cron/templates/configmap.yaml @@ -10,7 +10,9 @@ metadata: component: vault-agent data: cron-agent.hcl: | - + # Vault Agent Configuration for CronManager Service + # This agent provides CronManager with access to encryption keys and write access to secrets + vault { address = "http://vault:8200" retry { @@ -29,7 +31,7 @@ data: } } - # Write token to shared volume for agent to use + # Write token to file for CronManager service to use sink "file" { config = { path = "{{ .Values.agent.tokenPath }}/token" @@ -38,12 +40,12 @@ data: } } - # Caching configuration for CronManager + # Caching configuration cache { - default_lease_duration = "{{ .Values.agent.tokenTTL }}" + default_lease_duration = "{{ .Values.agent.tokenTTL }}" # Medium TTL for CronManager } - # API proxy listener - CronManager connects to localhost:{{ .Values.agent.port }} + # API proxy listener for CronManager service listener "tcp" { address = "0.0.0.0:{{ .Values.agent.port }}" tls_disable = true @@ -52,7 +54,5 @@ data: # API proxy configuration api_proxy { use_auto_auth_token = true - enforce_consistency = "always" - when_inconsistent = "forward" } -{{- end }} +{{- end }} \ No newline at end of file diff --git a/kubernetes/charts/Vault-Agent-GUI/templates/configmap.yaml b/kubernetes/charts/Vault-Agent-GUI/templates/configmap.yaml index 72ce877b..11e36dd8 100644 --- a/kubernetes/charts/Vault-Agent-GUI/templates/configmap.yaml +++ b/kubernetes/charts/Vault-Agent-GUI/templates/configmap.yaml @@ -10,7 +10,8 @@ metadata: data: gui-agent.hcl: | # Vault Agent Configuration for GUI Service - + # This agent provides GUI with access to public encryption key only + vault { address = "http://vault:8200" retry { @@ -29,7 +30,7 @@ data: } } - # Write token to shared volume for agent to use + # Write token to file for GUI service to use sink "file" { config = { path = "{{ .Values.agent.tokenPath }}/token" @@ -38,12 +39,12 @@ data: } } - # Caching configuration for GUI + # Caching configuration cache { - default_lease_duration = "{{ .Values.agent.tokenTTL }}" + default_lease_duration = "{{ .Values.agent.tokenTTL }}" # Short-lived tokens for GUI } - # API proxy listener - GUI connects to localhost:{{ .Values.agent.port }} + # API proxy listener for GUI service listener "tcp" { address = "0.0.0.0:{{ .Values.agent.port }}" tls_disable = true @@ -52,7 +53,5 @@ data: # API proxy configuration api_proxy { use_auto_auth_token = true - enforce_consistency = "always" - when_inconsistent = "forward" } -{{- end }} +{{- end }} \ No newline at end of file diff --git a/kubernetes/charts/Vault-Agent-LLM/templates/configmap.yaml b/kubernetes/charts/Vault-Agent-LLM/templates/configmap.yaml index 08c38ea2..1af16338 100644 --- a/kubernetes/charts/Vault-Agent-LLM/templates/configmap.yaml +++ b/kubernetes/charts/Vault-Agent-LLM/templates/configmap.yaml @@ -7,15 +7,13 @@ metadata: component: vault-agent data: agent.hcl: | - # Vault Agent Configuration for LLM Orchestration Service - vault { address = "http://vault:8200" retry { num_retries = 5 } } - + auto_auth { method "approle" { mount_path = "auth/approle" @@ -25,8 +23,7 @@ data: remove_secret_id_file_after_reading = false } } - - # Write token to shared volume for agent to use + sink "file" { config = { path = "/agent/llm-token/token" @@ -34,19 +31,16 @@ data: } } } - - # Caching configuration for LLM (longer TTL) + cache { - default_lease_duration = "1h" + default_lease_duration = "1h" # Longer TTL for LLM service } - + listener "tcp" { - address = "0.0.0.0:8201" + address = "0.0.0.0:8201" # Listen on all interfaces tls_disable = true } - + api_proxy { use_auto_auth_token = true - enforce_consistency = "always" - when_inconsistent = "forward" - } + } \ No newline at end of file diff --git a/kubernetes/charts/Vault-Agent-LLM/templates/deployment.yaml b/kubernetes/charts/Vault-Agent-LLM/templates/deployment.yaml deleted file mode 100644 index 943597a0..00000000 --- a/kubernetes/charts/Vault-Agent-LLM/templates/deployment.yaml +++ /dev/null @@ -1,107 +0,0 @@ -# DEPRECATED: This standalone deployment is no longer used -# WHY: Vault Agent now runs as a SIDECAR in LLM-Orchestration-Service pod -# This ensures LLM cannot bypass the agent and access Vault directly -# Keeping this file for any future reference -{{- if .Values.deployment.standalone }} -apiVersion: apps/v1 -kind: Deployment -metadata: - name: {{ .Values.release_name }} - labels: - app: {{ .Values.release_name }} - component: vault-agent-llm -spec: - replicas: {{ .Values.deployment.replicas }} - selector: - matchLabels: - app: {{ .Values.release_name }} - component: vault-agent-llm - template: - metadata: - labels: - app: {{ .Values.release_name }} - component: vault-agent-llm - spec: - {{- if .Values.affinity.enabled }} - affinity: - podAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - - labelSelector: - matchExpressions: - - key: app - operator: In - values: - - {{ .Values.vault.serviceName }} - topologyKey: kubernetes.io/hostname - {{- end }} - volumes: - {{- if .Values.volumes.agentCredentials.enabled }} - - name: vault-agent-creds - persistentVolumeClaim: - claimName: vault-agent-creds - {{- end }} - {{- if .Values.volumes.agentToken.enabled }} - - name: vault-agent-token - persistentVolumeClaim: - claimName: vault-agent-token - {{- end }} - {{- if .Values.volumes.agentConfig.enabled }} - - name: vault-agent-config - configMap: - name: {{ .Values.release_name }}-config - defaultMode: 0644 - {{- end }} - containers: - - name: vault-agent - image: "{{ .Values.images.vault.registry }}/{{ .Values.images.vault.repository }}:{{ .Values.images.vault.tag }}" - imagePullPolicy: {{ .Values.pullPolicy }} - command: - - vault - - agent - - -config=/agent/config/agent.hcl - - -log-level=info - env: - - name: VAULT_ADDR - value: {{ .Values.vault.addr | quote }} - - name: VAULT_SKIP_VERIFY - value: "true" - volumeMounts: - {{- if .Values.volumes.agentCredentials.enabled }} - - name: vault-agent-creds - mountPath: {{ .Values.volumes.agentCredentials.mountPath }} - readOnly: true - {{- end }} - {{- if .Values.volumes.agentToken.enabled }} - - name: vault-agent-token - mountPath: {{ .Values.volumes.agentToken.mountPath }} - {{- end }} - {{- if .Values.volumes.agentConfig.enabled }} - - name: vault-agent-config - mountPath: {{ .Values.volumes.agentConfig.mountPath }} - readOnly: true - {{- end }} - {{- if .Values.probes.livenessProbe.enabled }} - livenessProbe: - httpGet: - path: {{ .Values.probes.livenessProbe.httpGet.path }} - port: {{ .Values.probes.livenessProbe.httpGet.port }} - initialDelaySeconds: {{ .Values.probes.livenessProbe.initialDelaySeconds }} - periodSeconds: {{ .Values.probes.livenessProbe.periodSeconds }} - {{- end }} - {{- if .Values.probes.readinessProbe.enabled }} - readinessProbe: - httpGet: - path: {{ .Values.probes.readinessProbe.httpGet.path }} - port: {{ .Values.probes.readinessProbe.httpGet.port }} - initialDelaySeconds: {{ .Values.probes.readinessProbe.initialDelaySeconds }} - periodSeconds: {{ .Values.probes.readinessProbe.periodSeconds }} - {{- end }} - {{- if .Values.resources }} - resources: -{{ toYaml .Values.resources | indent 10 }} - {{- end }} - securityContext: - capabilities: - add: - - IPC_LOCK -{{- end }} \ No newline at end of file diff --git a/kubernetes/charts/Vault-Init/templates/configmap.yaml b/kubernetes/charts/Vault-Init/templates/configmap.yaml index 9cc2b12a..036eb1e1 100644 --- a/kubernetes/charts/Vault-Init/templates/configmap.yaml +++ b/kubernetes/charts/Vault-Init/templates/configmap.yaml @@ -9,13 +9,95 @@ data: {{ .Values.initScript.filename }}: | #!/bin/sh set -e - + VAULT_ADDR="${VAULT_ADDR:-http://vault:8200}" UNSEAL_KEYS_FILE="/vault/data/unseal-keys.json" INIT_FLAG="/vault/data/.initialized" - + echo "=== Vault Initialization Script ===" - + + # --------------------------------------------------------------------------- + # Helpers (used by the SUBSEQUENT DEPLOYMENT branch) + # --------------------------------------------------------------------------- + + # Ensure a role_id file exists on disk; fetch from Vault if missing. + # Usage: ensure_role_id + ensure_role_id() { + role="$1"; rid_file="$2" + if [ -f "$rid_file" ] && [ -s "$rid_file" ]; then + return 0 + fi + echo "Fetching role_id for $role..." + rid=$(wget -q -O- \ + --header="X-Vault-Token: $ROOT_TOKEN" \ + "$VAULT_ADDR/v1/auth/approle/role/$role/role-id" | \ + grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') + echo "$rid" > "$rid_file" + chmod 640 "$rid_file" + } + + # Return 0 if the on-disk role_id + secret_id still authenticate, 1 otherwise. + # Usage: validate_secret_id + validate_secret_id() { + rid_file="$1"; sid_file="$2" + [ -f "$rid_file" ] && [ -f "$sid_file" ] || return 1 + rid=$(cat "$rid_file"); sid=$(cat "$sid_file") + [ -n "$rid" ] && [ -n "$sid" ] || return 1 + # wget returns non-zero on HTTP 400 (invalid creds); also confirm a token came back. + resp=$(wget -q -O- \ + --post-data="{\"role_id\":\"$rid\",\"secret_id\":\"$sid\"}" \ + --header='Content-Type: application/json' \ + "$VAULT_ADDR/v1/auth/approle/login" 2>/dev/null) || return 1 + echo "$resp" | grep -q '"client_token"' || return 1 + return 0 + } + + # Mint a fresh secret_id for a role and write it to disk. + # Usage: mint_secret_id + mint_secret_id() { + role="$1"; sid_file="$2" + sid=$(wget -q -O- --post-data='' \ + --header="X-Vault-Token: $ROOT_TOKEN" \ + "$VAULT_ADDR/v1/auth/approle/role/$role/secret-id" | \ + grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') + echo "$sid" > "$sid_file" + chmod 640 "$sid_file" + } + + # Reuse the existing secret_id if it still authenticates; otherwise mint a new one. + # Usage: reconcile_secret_id + reconcile_secret_id() { + role="$1"; rid_file="$2"; sid_file="$3" + ensure_role_id "$role" "$rid_file" + if validate_secret_id "$rid_file" "$sid_file"; then + echo "$role: existing secret_id still valid - reusing" + else + echo "$role: secret_id invalid or missing - minting a new one" + mint_secret_id "$role" "$sid_file" + fi + } + + # Create or update an AppRole that issues a PERIODIC token (no max_ttl): the + # agent renews it forever and never re-runs approle/login in steady state. + # secret_id_ttl=0 + secret_id_num_uses=0 keep the secret_id valid across + # restarts. Idempotent: does not invalidate existing secret_ids, safe per run. + # Usage: upsert_approle + upsert_approle() { + role="$1"; policy="$2"; period="$3" + wget -q -O- --post-data='{"token_policies":["'"$policy"'"],"token_period":"'"$period"'","token_num_uses":0,"secret_id_ttl":"0","secret_id_num_uses":0,"bind_secret_id":true}' \ + --header="X-Vault-Token: $ROOT_TOKEN" \ + --header='Content-Type: application/json' \ + "$VAULT_ADDR/v1/auth/approle/role/$role" >/dev/null + } + + # Apply the current AppRole definitions for all three services. + ensure_approles() { + echo "Ensuring AppRole configs (periodic tokens)..." + upsert_approle "gui-service" "gui-policy" "20m" + upsert_approle "cron-manager-service" "cron-manager-policy" "30m" + upsert_approle "llm-orchestration-service" "llm-orchestration-policy" "1h" + } + # Wait for Vault to be ready echo "Waiting for Vault..." for i in $(seq 1 30); do @@ -26,7 +108,7 @@ data: echo "Waiting... ($i/30)" sleep 2 done - + # Check if this is first time if [ ! -f "$INIT_FLAG" ]; then echo "=== FIRST TIME DEPLOYMENT ===" @@ -115,6 +197,8 @@ data: path "secret/data/embeddings/connections/*" { capabilities = ["read", "list"] } path "secret/metadata/embeddings/connections/*" { capabilities = ["read", "list"] } path "secret/data/encryption/*" { capabilities = ["deny"] } + path "secret/data/langfuse/*" { capabilities = ["read"] } + path "secret/metadata/langfuse/*" { capabilities = ["read", "list"] } path "auth/token/lookup-self" { capabilities = ["read"] }' LLM_POLICY_JSON=$(echo "$LLM_POLICY" | jq -Rs '{"policy":.}') @@ -123,27 +207,9 @@ data: --header='Content-Type: application/json' \ "$VAULT_ADDR/v1/sys/policies/acl/llm-orchestration-policy" >/dev/null - # Create GUI AppRole - echo "Creating gui-service AppRole..." - wget -q -O- --post-data='{"token_policies":["gui-policy"],"token_no_default_policy":true,"token_ttl":"15m","token_max_ttl":"1h","secret_id_ttl":"24h","secret_id_num_uses":0,"bind_secret_id":true}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/auth/approle/role/gui-service" >/dev/null - - # Create CronManager AppRole - echo "Creating cron-manager-service AppRole..." - wget -q -O- --post-data='{"token_policies":["cron-manager-policy"],"token_no_default_policy":true,"token_ttl":"30m","token_max_ttl":"8h","secret_id_ttl":"24h","secret_id_num_uses":0,"bind_secret_id":true}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/auth/approle/role/cron-manager-service" >/dev/null - - # Create LLM Orchestration AppRole - echo "Creating llm-orchestration-service AppRole..." - wget -q -O- --post-data='{"token_policies":["llm-orchestration-policy"],"token_no_default_policy":true,"token_ttl":"1h","token_max_ttl":"24h","secret_id_ttl":"24h","secret_id_num_uses":0,"bind_secret_id":true}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/auth/approle/role/llm-orchestration-service" >/dev/null - + # Create the three AppRoles (periodic tokens - see upsert_approle). + ensure_approles + # Ensure credentials directory exists mkdir -p /agent/credentials @@ -238,18 +304,10 @@ data: rm -rf "$TEMP_KEY_DIR" echo "RSA keypair generated and stored successfully" - # Store test LLM credentials for testing - echo "Creating test LLM credentials..." - wget -q -O- --post-data='{"data":{"access_key":"TEST_AWS_ACCESS_KEY","secret_key":"TEST_AWS_SECRET_KEY","environment":"production","model":"claude-3"}}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/secret/data/llm/connections/aws_bedrock/production/claude-3" >/dev/null - # Mark as initialized touch "$INIT_FLAG" echo "=== First time setup complete ===" - - + else echo "=== SUBSEQUENT DEPLOYMENT ===" @@ -285,65 +343,22 @@ data: # Get root token ROOT_TOKEN=$(grep -o '"root_token":"[^"]*"' "$UNSEAL_KEYS_FILE" | cut -d':' -f2 | tr -d '"') export VAULT_TOKEN="$ROOT_TOKEN" - + + # Re-apply AppRole definitions so config changes (e.g. periodic tokens) + # take effect on redeploy without re-initializing Vault. Idempotent and + # does not invalidate existing secret_ids. + ensure_approles + # Ensure credentials directory exists mkdir -p /agent/credentials - # Always regenerate all secret_ids on restart - echo "Regenerating GUI secret_id..." - GUI_SECRET_ID=$(wget -q -O- --post-data='' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/gui-service/secret-id" | \ - grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$GUI_SECRET_ID" > /agent/credentials/gui_secret_id - - echo "Regenerating CronManager secret_id..." - CRON_SECRET_ID=$(wget -q -O- --post-data='' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/cron-manager-service/secret-id" | \ - grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$CRON_SECRET_ID" > /agent/credentials/cron_secret_id - - echo "Regenerating LLM secret_id..." - LLM_SECRET_ID=$(wget -q -O- --post-data='' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/llm-orchestration-service/secret-id" | \ - grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$LLM_SECRET_ID" > /agent/credentials/llm_secret_id - - # Set permissions - chmod 640 /agent/credentials/*_secret_id - - # Ensure role_ids exist - if [ ! -f /agent/credentials/gui_role_id ]; then - echo "Copying GUI role_id..." - GUI_ROLE_ID=$(wget -q -O- \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/gui-service/role-id" | \ - grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$GUI_ROLE_ID" > /agent/credentials/gui_role_id - chmod 640 /agent/credentials/gui_role_id - fi - - if [ ! -f /agent/credentials/cron_role_id ]; then - echo "Copying CronManager role_id..." - CRON_ROLE_ID=$(wget -q -O- \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/cron-manager-service/role-id" | \ - grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$CRON_ROLE_ID" > /agent/credentials/cron_role_id - chmod 640 /agent/credentials/cron_role_id - fi - - if [ ! -f /agent/credentials/llm_role_id ]; then - echo "Copying LLM role_id..." - LLM_ROLE_ID=$(wget -q -O- \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/llm-orchestration-service/role-id" | \ - grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$LLM_ROLE_ID" > /agent/credentials/llm_role_id - chmod 640 /agent/credentials/llm_role_id - fi + # Reconcile secret_ids: reuse the existing one if it still authenticates, + # mint a new one only if invalid or missing - keeps one stable secret_id + # across restarts instead of rotating every boot. reconcile_secret_id also + # ensures the role_id file exists first (validation needs both). + reconcile_secret_id "gui-service" /agent/credentials/gui_role_id /agent/credentials/gui_secret_id + reconcile_secret_id "cron-manager-service" /agent/credentials/cron_role_id /agent/credentials/cron_secret_id + reconcile_secret_id "llm-orchestration-service" /agent/credentials/llm_role_id /agent/credentials/llm_secret_id fi - + echo "=== Vault init complete ===" \ No newline at end of file diff --git a/kubernetes/charts/Vault-Init/templates/job.yaml b/kubernetes/charts/Vault-Init/templates/job.yaml index 4c1f9811..52758783 100644 --- a/kubernetes/charts/Vault-Init/templates/job.yaml +++ b/kubernetes/charts/Vault-Init/templates/job.yaml @@ -14,6 +14,9 @@ spec: component: vault-init spec: restartPolicy: {{ .Values.job.restartPolicy }} + # Run as root to allow chown of PVC directories (mirrors docker-compose user: "0") + securityContext: + runAsUser: 0 {{- if .Values.affinity.enabled }} affinity: podAffinity: diff --git a/kubernetes/charts/Vault/templates/configmap.yaml b/kubernetes/charts/Vault/templates/configmap.yaml index 1e32fd90..cea5b77c 100644 --- a/kubernetes/charts/Vault/templates/configmap.yaml +++ b/kubernetes/charts/Vault/templates/configmap.yaml @@ -9,24 +9,29 @@ metadata: data: vault.hcl: | # HashiCorp Vault Server Configuration - # Production-ready configuration for LLM Orchestration Service - - # Storage backend - Raft for high availability + # Single-node Raft for the RAG-Module services + + # Storage backend - Raft storage "raft" { path = "/vault/file" node_id = "vault-node-1" - - # Retry join configuration for clustering (single node for now) - retry_join { - leader_api_addr = "http://vault:8200" - } + + # NOTE: No retry_join for a single node. A lone node self-bootstraps. + # A retry_join pointing at itself causes repeated + # "failed to get raft challenge ... Vault is sealed" errors and a + # messy double Raft init on every boot. Add retry_join back only when + # you actually have peer nodes to join. } - - # HTTP listener configuration + + # HTTP API listener. + # Vault automatically uses the next port up (8201) as its internal + # cluster port, so do NOT define a separate listener on 8201 — that + # collides with the cluster listener ("bind: address already in use") + # and degrades the login/request-forwarding path the agents rely on. listener "tcp" { - address = "0.0.0.0:8200" - tls_disable = true - + address = "0.0.0.0:8200" + tls_disable = true + # Enable CORS for web UI access cors_enabled = true cors_allowed_origins = [ @@ -34,33 +39,24 @@ data: "http://vault:8200" ] } - - # Cluster listener for HA (required even for single node) - listener "tcp" { - address = "0.0.0.0:8201" - cluster_addr = "http://0.0.0.0:8201" - tls_disable = true - } - - # API and cluster addresses + + # API and cluster addresses. + # cluster_addr tells Vault where its internal cluster port (8201) is + # reachable; Vault binds that port itself — no listener block needed. api_addr = "http://vault:8200" cluster_addr = "http://vault:8201" - + # Security and performance settings disable_mlock = false disable_cache = false ui = false - + # Default lease and maximum lease durations default_lease_ttl = "168h" # 7 days max_lease_ttl = "720h" # 30 days - + # Logging configuration - log_level = "INFO" + log_level = "INFO" log_format = "json" - - # Development settings (remove in production) - # Note: In production, you should not use dev mode - # and should properly initialize and unseal the vault {{- end }} \ No newline at end of file diff --git a/notification-server/src/connectionManager.js b/notification-server/src/connectionManager.js index a2dee15d..5dba38bd 100644 --- a/notification-server/src/connectionManager.js +++ b/notification-server/src/connectionManager.js @@ -1,5 +1,52 @@ const activeConnections = new Map(); +/** + * Register an AbortController for an in-flight upstream request belonging to a + * connection, so it can be cancelled when the browser disconnects. + * @param {string} connectionId + * @param {AbortController} controller + */ +function registerAbortController(connectionId, controller) { + const connData = activeConnections.get(connectionId); + if (!connData) return; + if (!connData.abortControllers) { + connData.abortControllers = new Set(); + } + connData.abortControllers.add(controller); +} + +/** + * Remove a previously registered AbortController (upstream request finished). + * @param {string} connectionId + * @param {AbortController} controller + */ +function unregisterAbortController(connectionId, controller) { + const connData = activeConnections.get(connectionId); + connData?.abortControllers?.delete(controller); +} + +/** + * Abort every in-flight upstream request for a connection. Called when the + * browser goes away, so we stop paying for generation nobody will read. + * @param {string} connectionId + */ +function abortConnectionRequests(connectionId) { + const connData = activeConnections.get(connectionId); + if (!connData?.abortControllers) return; + + for (const controller of connData.abortControllers) { + try { + controller.abort(); + } catch (error) { + console.error(`Failed to abort upstream request for ${connectionId}:`, error); + } + } + connData.abortControllers.clear(); +} + module.exports = { activeConnections, + registerAbortController, + unregisterAbortController, + abortConnectionRequests, }; diff --git a/notification-server/src/openSearch.js b/notification-server/src/openSearch.js index 1be28b3f..f0262937 100644 --- a/notification-server/src/openSearch.js +++ b/notification-server/src/openSearch.js @@ -39,8 +39,8 @@ async function createLLMOrchestrationStreamRequest({ channelId, message, options authorId: options.authorId || `user-${channelId}`, conversationHistory: options.conversationHistory || [], url: options.url || "sse-stream-context", - environment: "production", // Streaming only works in production - connection_id: options.connection_id || connectionId + environment: options.environment || "production", + connection_id: options.connection_id }; console.log(`Calling LLM orchestration stream for channel ${channelId}`); diff --git a/notification-server/src/server.js b/notification-server/src/server.js index 98e91574..d2f67adb 100644 --- a/notification-server/src/server.js +++ b/notification-server/src/server.js @@ -82,4 +82,12 @@ const server = app.listen(serverConfig.port, () => { console.log(`LLM orchestration streaming at: /channels/:channelId/orchestrate/stream`); }); +// The SSE GET and the trigger POST are both long-lived by design, so Node's +// default 300s requestTimeout would cut them off mid-answer. Disable the +// per-request cap and let the upstream idle watchdog in streamingService.js +// decide when a stream has genuinely stalled. +server.requestTimeout = 0; +server.headersTimeout = 65_000; +server.keepAliveTimeout = 61_000; + module.exports = server; diff --git a/notification-server/src/sseUtil.js b/notification-server/src/sseUtil.js index a2ad0c1a..bd915018 100644 --- a/notification-server/src/sseUtil.js +++ b/notification-server/src/sseUtil.js @@ -1,11 +1,18 @@ const { v4: uuidv4 } = require('uuid'); const streamQueue = require("./streamQueue"); const { createLLMOrchestrationStreamRequest } = require("./streamingService"); -const { activeConnections } = require("./connectionManager"); +const { activeConnections, abortConnectionRequests } = require("./connectionManager"); + +// Comment frames are written this often so that every intermediate proxy sees +// traffic and does not close the connection on its idle timer. EventSource +// ignores comment frames, so this is invisible to the browser. +const HEARTBEAT_INTERVAL_MS = Number( + process.env.SSE_HEARTBEAT_INTERVAL_MS || 15_000 +); function buildSSEResponse({ res, req, buildCallbackFunction, channelId }) { addSSEHeader(req, res); - keepStreamAlive(res); + const heartbeat = keepStreamAlive(res); const connectionId = generateConnectionID(); const sender = buildSender(res); @@ -13,6 +20,7 @@ function buildSSEResponse({ res, req, buildCallbackFunction, channelId }) { res, sender, channelId, + abortControllers: new Set(), }); if (channelId) { @@ -25,6 +33,9 @@ function buildSSEResponse({ res, req, buildCallbackFunction, channelId }) { req.on("close", () => { console.log(`Client disconnected from SSE for channel ${channelId}`); + clearInterval(heartbeat); + // Cancel any in-flight upstream generation - nobody is left to read it. + abortConnectionRequests(connectionId); activeConnections.delete(connectionId); cleanUp?.(); }); @@ -37,6 +48,8 @@ function addSSEHeader(req, res) { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache', 'Connection': 'keep-alive', + // Stops nginx-style proxies buffering the stream into a single response. + 'X-Accel-Buffering': 'no', 'Access-Control-Allow-Origin': origin, 'Access-Control-Allow-Credentials': true, 'Access-Control-Expose-Headers': 'Origin, X-Requested-With, Content-Type, Cache-Control, Connection, Accept' @@ -49,8 +62,33 @@ function extractOrigin(reqOrigin) { return whitelisted ? reqOrigin : '*'; } +/** + * Keep the SSE connection warm with periodic comment frames. + * + * Previously a single `res.write('')`, which did nothing beyond the initial + * flush - any proxy between the browser and this server would still time the + * connection out during a long generation pause. + * + * @returns {NodeJS.Timeout} interval handle; the caller must clear it on close. + */ function keepStreamAlive(res) { res.write(''); + const heartbeat = setInterval(() => { + try { + // A `:` line is an SSE comment: ignored by EventSource, but it is traffic. + res.write(': ping\n\n'); + if (typeof res.flush === "function") { + res.flush(); + } + } catch (error) { + console.error("SSE heartbeat write failed:", error); + clearInterval(heartbeat); + } + }, HEARTBEAT_INTERVAL_MS); + + // Do not hold the event loop open purely for a heartbeat. + heartbeat.unref?.(); + return heartbeat; } function generateConnectionID() { diff --git a/notification-server/src/streamingService.js b/notification-server/src/streamingService.js index 74e58869..cdb1df40 100644 --- a/notification-server/src/streamingService.js +++ b/notification-server/src/streamingService.js @@ -1,6 +1,96 @@ -const { activeConnections } = require("./connectionManager"); +const { + activeConnections, + registerAbortController, + unregisterAbortController, +} = require("./connectionManager"); const streamQueue = require("./streamQueue"); +// Inactivity budget for the upstream SSE body. This is an idle timeout, not a +// total-duration cap: the timer resets on every byte received, so a long answer +// streams fine while a genuinely stalled upstream fails fast with a clear error. +// Kept below undici's 300s default bodyTimeout so we control the failure mode +// (undici's would surface as an opaque `TypeError: terminated`). +const UPSTREAM_IDLE_TIMEOUT_MS = Number( + process.env.LLM_STREAM_IDLE_TIMEOUT_MS || 120_000 +); + +const ORCHESTRATOR_URL = + process.env.LLM_ORCHESTRATOR_URL || "http://llm-orchestration-service:8100"; + +/** + * Translate one SSE `data:` line into a message for the browser. + * @returns {boolean} true if this line ended the stream + */ +function relaySSELine({ line, channelId, sender }) { + if (!line.trim()) return false; + if (!line.startsWith("data: ")) return false; // ignores `: ping` heartbeats + + try { + const data = JSON.parse(line.slice(6)); // Remove 'data: ' prefix + const content = data.payload?.content; + const buttons = data.payload?.buttons; + + if (!content) return false; + + if (content === "END") { + sender({ + type: "stream_end", + streamId: channelId, + channelId, + isComplete: true, + }); + return true; + } + + // Regular token - send to client (include buttons when present) + const chunkMessage = { + type: "stream_chunk", + content: content, + streamId: channelId, + channelId, + isComplete: false, + }; + if (buttons && buttons.length > 0) { + chunkMessage.buttons = buttons; + } + sender(chunkMessage); + return false; + } catch (parseError) { + console.error(`Failed to parse SSE data for channel ${channelId}:`, parseError, line); + return false; + } +} + +/** + * Drain the upstream SSE body, relaying each frame to the browser. + * Returns once the stream ends, the client disconnects, or END is received. + * @returns {Promise} true if an END frame terminated the stream + */ +async function relayUpstreamBody({ response, connectionId, channelId, sender, onActivity }) { + const reader = response.body.getReader(); + const decoder = new TextDecoder(); + let buffer = ""; + + while (activeConnections.has(connectionId)) { + const { done, value } = await reader.read(); + if (done) break; + + onActivity(); + + buffer += decoder.decode(value, { stream: true }); + const lines = buffer.split("\n"); + buffer = lines.pop() || ""; // Keep the incomplete line in buffer + + for (const line of lines) { + // Returning here (rather than `break`) also closes the upstream body — + // a bare `break` only left the inner loop and kept the connection open. + if (relaySSELine({ line, channelId, sender })) return true; + } + } + + return false; +} + /** * Stream LLM orchestration response to connected clients * @param {Object} params - Request parameters @@ -17,7 +107,7 @@ async function createLLMOrchestrationStreamRequest({ channelId, message, options if (connections.length === 0) { streamQueue.addToQueue(channelId, { message, options }); - + if (streamQueue.shouldRetry({ retryCount: 0 })) { throw new Error("No active connections found for this channel - request queued"); } else { @@ -28,116 +118,9 @@ async function createLLMOrchestrationStreamRequest({ channelId, message, options console.log(`Streaming LLM orchestration for channel ${channelId} to ${connections.length} connections`); try { - const responsePromises = connections.map(async ([connectionId, connData]) => { - const { sender } = connData; - - try { - // Construct OrchestrationRequest payload - const orchestrationPayload = { - chatId: channelId, - message: message, - authorId: options.authorId || `user-${channelId}`, - conversationHistory: options.conversationHistory || [], - url: options.url || "sse-stream-context", - environment: "production", // Streaming only works in production - connection_id: options.connection_id || connectionId - }; - - console.log(`Calling LLM orchestration stream for channel ${channelId}`); - - // Call the LLM orchestration streaming endpoint - const response = await fetch(`${process.env.LLM_ORCHESTRATOR_URL || 'http://llm-orchestration-service:8100'}/orchestrate/stream`, { - method: 'POST', - headers: { - 'Content-Type': 'application/json', - }, - body: JSON.stringify(orchestrationPayload), - }); - - if (!response.ok) { - throw new Error(`LLM Orchestration API error: ${response.status} ${response.statusText}`); - } - - if (!activeConnections.has(connectionId)) { - return; - } - - // Send stream start notification - sender({ - type: "stream_start", - streamId: channelId, - channelId, - isComplete:false - }); - - const reader = response.body.getReader(); - const decoder = new TextDecoder(); - let buffer = ''; - - while (true) { - if (!activeConnections.has(connectionId)) break; - - const { done, value } = await reader.read(); - if (done) break; - - buffer += decoder.decode(value, { stream: true }); - const lines = buffer.split('\n'); - buffer = lines.pop() || ''; // Keep the incomplete line in buffer - - for (const line of lines) { - if (!line.trim()) continue; - if (!line.startsWith('data: ')) continue; - - try { - const data = JSON.parse(line.slice(6)); // Remove 'data: ' prefix - const content = data.payload?.content; - const buttons = data.payload?.buttons; - - if (!content) continue; - - if (content === "END") { - // Stream completed - sender({ - type: "stream_end", - streamId: channelId, - channelId, - isComplete:true - }); - break; - } - - // Regular token - send to client (include buttons when present) - const chunkMessage = { - type: "stream_chunk", - content: content, - streamId: channelId, - channelId, - isComplete:false - }; - if (buttons && buttons.length > 0) { - chunkMessage.buttons = buttons; - } - sender(chunkMessage); - - } catch (parseError) { - console.error(`Failed to parse SSE data for channel ${channelId}:`, parseError, line); - } - } - } - - } catch (error) { - console.error(`Streaming error for connection ${connectionId}:`, error); - if (activeConnections.has(connectionId)) { - sender({ - type: "stream_error", - error: error.message, - streamId: channelId, - channelId, - isComplete:true - }); - } - } - }); + const responsePromises = connections.map(([connectionId, connData]) => + streamToConnection({ connectionId, connData, channelId, message, options }) + ); await Promise.all(responsePromises); return { success: true, message: "Stream completed" }; @@ -148,6 +131,138 @@ async function createLLMOrchestrationStreamRequest({ channelId, message, options } } +/** + * Run one upstream orchestration stream and relay it to a single SSE connection. + */ +async function streamToConnection({ connectionId, connData, channelId, message, options }) { + const { sender } = connData; + const abortController = new AbortController(); + let idleTimer = null; + let idleTimedOut = false; + + // Idle watchdog: reset on every chunk received. Only fires when the upstream + // genuinely stops producing, never on a merely long answer. + const resetIdleTimer = () => { + if (idleTimer) clearTimeout(idleTimer); + idleTimer = setTimeout(() => { + idleTimedOut = true; + console.error( + `Upstream idle for ${UPSTREAM_IDLE_TIMEOUT_MS}ms on channel ${channelId} - aborting` + ); + abortController.abort(); + }, UPSTREAM_IDLE_TIMEOUT_MS); + }; + + try { + // Construct OrchestrationRequest payload + const orchestrationPayload = { + chatId: channelId, + message: message, + authorId: options.authorId || `user-${channelId}`, + conversationHistory: options.conversationHistory || [], + url: options.url || "sse-stream-context", + environment: options.environment || "production", + connection_id: options.connection_id + }; + + console.log(`Calling LLM orchestration stream for channel ${channelId}`); + + // The controller serves two purposes: cancelling the upstream request when + // the browser disconnects, and enforcing the idle timeout above. + registerAbortController(connectionId, abortController); + resetIdleTimer(); + + // Call the LLM orchestration streaming endpoint + const response = await fetch(`${ORCHESTRATOR_URL}/orchestrate/stream`, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + }, + body: JSON.stringify(orchestrationPayload), + signal: abortController.signal, + }); + + if (!response.ok) { + throw new Error(`LLM Orchestration API error: ${response.status} ${response.statusText}`); + } + + if (!activeConnections.has(connectionId)) { + return; + } + + // Send stream start notification + sender({ + type: "stream_start", + streamId: channelId, + channelId, + isComplete: false + }); + + const sawEnd = await relayUpstreamBody({ + response, + connectionId, + channelId, + sender, + onActivity: resetIdleTimer, + }); + + // The body closed without an END frame. The browser has a stream_start and + // possibly some chunks, but nothing that ends the stream, so it would spin + // forever. Terminate it explicitly rather than trusting every upstream error + // path to remember the marker. A client that has already disconnected needs + // no notification. + if (!sawEnd && activeConnections.has(connectionId)) { + console.error( + `Upstream body ended without END marker on channel ${channelId} ` + + `(connection ${connectionId})` + ); + sender({ + type: "stream_error", + error: "The response ended unexpectedly. Please try again.", + streamId: channelId, + channelId, + isComplete: true, + }); + } + + } catch (error) { + // A client-disconnect abort is expected teardown, not a failure: the browser + // is gone, there is nobody to notify and nothing to log loudly. + if (error.name === "AbortError" && !idleTimedOut) { + console.log( + `Upstream stream cancelled for connection ${connectionId} (client disconnected)` + ); + return; + } + + if (idleTimedOut) { + // The AbortError here is our own watchdog firing; its stack says nothing + // useful, so report the actual cause instead. + console.error( + `Streaming timed out for connection ${connectionId}: no upstream output ` + + `for ${UPSTREAM_IDLE_TIMEOUT_MS}ms on channel ${channelId}` + ); + } else { + console.error(`Streaming error for connection ${connectionId}:`, error); + } + + if (activeConnections.has(connectionId)) { + sender({ + type: "stream_error", + error: idleTimedOut + ? "The response timed out. Please try again." + : error.message, + streamId: channelId, + channelId, + isComplete: true + }); + } + } finally { + if (idleTimer) clearTimeout(idleTimer); + unregisterAbortController(connectionId, abortController); + } +} + module.exports = { createLLMOrchestrationStreamRequest, }; diff --git a/package-lock.json b/package-lock.json new file mode 100644 index 00000000..c5b87041 --- /dev/null +++ b/package-lock.json @@ -0,0 +1,30 @@ +{ + "name": "LLM-Module", + "version": "0.0.1", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "version": "0.0.1", + "devDependencies": { + "husky": "^9.1.6" + } + }, + "node_modules/husky": { + "version": "9.1.7", + "resolved": "https://registry.npmjs.org/husky/-/husky-9.1.7.tgz", + "integrity": "sha512-5gs5ytaNjBrh5Ow3zrvdUUY+0VxIuWVL4i9irt6friV+BqdCfmV11CQTWMiBYWHbXhco+J1kHfTOUkePhCDvMA==", + "dev": true, + "license": "MIT", + "bin": { + "husky": "bin.js" + }, + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/typicode" + } + } + } +} diff --git a/package.json b/package.json new file mode 100644 index 00000000..3b3057a0 --- /dev/null +++ b/package.json @@ -0,0 +1,10 @@ +{ + "version": "0.0.1", + "devDependencies": { + "husky": "^9.1.6" + }, + "scripts": { + "prepare": "husky", + "changelog": "./generate-changelog.sh" + } +} diff --git a/pyproject.toml b/pyproject.toml index 6e39a1ee..a79351f4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,6 +6,7 @@ readme = "README.md" requires-python = "==3.12.10" dependencies = [ "azure-identity>=1.24.0", + "azure-storage-blob>=12.24.0", "boto3>=1.40.25", "dspy>=3.0.3", "openai>=1.106.1", @@ -104,6 +105,7 @@ unfixable = [] "src/response_generator/response_generate.py" = ["N815", "ANN401"] # Pydantic model fields + DSPy streamify Any type # Library interface patterns - legitimate Any usage +"src/utils/observation_utils.py" = ["ANN401"] "src/contextual_retrieval/contextual_retrieval_api_client.py" = ["ANN401"] # httpx **kwargs pass-through "src/tool_classifier/workflows/service_workflow.py" = ["ANN401"] # LLMManager passed as Any - dynamic multi-provider LLM interface "src/guardrails/dspy_nemo_adapter.py" = ["ANN401"] # LangChain LLM interface + DSPy dynamic types @@ -116,6 +118,7 @@ unfixable = [] "src/utils/api_tool_session_store.py" = ["ANN401"] # Dynamic Pydantic model field updates via **kwargs +"src/utils/atc_cache_store.py" = ["ANN401"] # raw_response / return type are parsed JSON (dict or list) — Any is the correct annotation "src/tool_classifier/workflows/api_tool_workflow.py" = ["ANN401", "N815"] # Dynamic guardrails adapter + orchestration service Any types; camelCase _MinimalRequest field for API contract [tool.ruff.format] diff --git a/release.env b/release.env new file mode 100644 index 00000000..4e84d7fd --- /dev/null +++ b/release.env @@ -0,0 +1,4 @@ +main +MAJOR=1 +MINOR=0 +PATCH=0 diff --git a/src/api_tool_indexer/constants.py b/src/api_tool_indexer/constants.py index 1bdd7b6d..7d84095a 100644 --- a/src/api_tool_indexer/constants.py +++ b/src/api_tool_indexer/constants.py @@ -23,7 +23,9 @@ class ApiToolIndexerConstants: # LLM / Embedding API DEFAULT_API_BASE_URL = "http://llm-orchestration-service:8100" DEFAULT_ENVIRONMENT = "production" - DEFAULT_CONNECTION_ID = "gpt-4o-mini" + # None → orchestration service resolves the embedding model via the + # DB-fetched vault UUID (path: embeddings/connections/{provider}/{vault_uuid}). + DEFAULT_CONNECTION_ID = None # Retry Configuration MAX_RETRIES = 3 diff --git a/src/api_tool_indexer/main_indexer.py b/src/api_tool_indexer/main_indexer.py index 5bddceeb..5ec38565 100644 --- a/src/api_tool_indexer/main_indexer.py +++ b/src/api_tool_indexer/main_indexer.py @@ -28,7 +28,8 @@ import asyncio import argparse from typing import List -from loguru import logger + +from src.loki_logger import LokiLogger from api_tool_indexer.constants import ApiToolIndexerConstants from api_tool_indexer.models import EndpointData, EnrichedEndpoint, IndexingResult @@ -37,8 +38,7 @@ # Reuse LLMAPIClient from intent_data_enrichment. from intent_data_enrichment.api_client import LLMAPIClient -# Reuse sparse encoder from tool_classifier (shared BM25 implementation). -sys.path.insert(0, "/app/src") +logger = LokiLogger(service_name="api-tool-calling") try: from tool_classifier.sparse_encoder import compute_sparse_vector except ImportError: @@ -139,9 +139,8 @@ async def _generate_context_for_endpoint( ) logger.debug( - "Generated context prompt for endpoint '{}': {} chars", - endpoint_data.endpoint_id, - len(context_prompt), + f"Generated context prompt for endpoint '{endpoint_data.endpoint_id}': " + f"{len(context_prompt)} chars" ) # context_type="api_tool" makes context_manager use API_TOOL_CONTEXT_PROMPT, @@ -175,11 +174,10 @@ async def _generate_context_for_endpoint( context = result.get("context", "").strip() - logger.debug( - "context preview: {}{}", - context[:200].replace("\n", "\\n"), - "..." if len(context) > 200 else "", + _preview = context[:200].replace("\n", "\\n") + ( + "..." if len(context) > 200 else "" ) + logger.debug(f"context preview: {_preview}") if not context: raise ValueError("Empty context returned from API") @@ -320,6 +318,8 @@ async def index_endpoint(endpoint_data: EndpointData) -> IndexingResult: params=endpoint_data.params, enriched_context=enriched_context, service_id=endpoint_data.service_id, + cacheable=endpoint_data.cacheable, + cache_ttl_seconds=endpoint_data.cache_ttl_seconds, point_type="example", example_text=example, embedding=ex_embedding, @@ -351,6 +351,8 @@ async def index_endpoint(endpoint_data: EndpointData) -> IndexingResult: params=endpoint_data.params, enriched_context=enriched_context, service_id=endpoint_data.service_id, + cacheable=endpoint_data.cacheable, + cache_ttl_seconds=endpoint_data.cache_ttl_seconds, point_type="summary", embedding=summary_embedding, sparse_indices=summary_sparse.indices, diff --git a/src/api_tool_indexer/models.py b/src/api_tool_indexer/models.py index 6333d2fd..d444df44 100644 --- a/src/api_tool_indexer/models.py +++ b/src/api_tool_indexer/models.py @@ -39,6 +39,21 @@ class EndpointData(BaseModel): ) visibility: str = Field(default="private", description="public or private") type: str = Field(default="custom_endpoint", description="Endpoint type") + cacheable: bool = Field( + default=True, + description=( + "Set False for endpoints returning sensitive/personal data " + "(e.g. document status). Disables all L1/L2 cache writes." + ), + ) + cache_ttl_seconds: Optional[int] = Field( + default=None, + ge=1, + description=( + "Per-endpoint L1 TTL override in seconds. " + "None = use ATC_CACHE_DEFAULT_TTL_SECONDS." + ), + ) class EnrichedEndpoint(BaseModel): @@ -69,6 +84,22 @@ class EnrichedEndpoint(BaseModel): ) service_id: Optional[str] = Field(default=None, description="Parent service UUID") + cacheable: bool = Field( + default=True, + description=( + "Propagated from EndpointData. False disables all L1/L2 cache writes " + "for this endpoint at query time." + ), + ) + cache_ttl_seconds: Optional[int] = Field( + default=None, + ge=1, + description=( + "Per-endpoint L1 TTL override propagated from EndpointData. " + "None = use ATC_CACHE_DEFAULT_TTL_SECONDS." + ), + ) + # Point type — controls which text was embedded for this point point_type: str = Field( default="summary", diff --git a/src/api_tool_indexer/qdrant_manager.py b/src/api_tool_indexer/qdrant_manager.py index de2fa0e8..bd572661 100644 --- a/src/api_tool_indexer/qdrant_manager.py +++ b/src/api_tool_indexer/qdrant_manager.py @@ -4,7 +4,8 @@ import uuid from typing import Any, Dict, List, Optional -from loguru import logger + +from src.loki_logger import LokiLogger from qdrant_client import QdrantClient from qdrant_client.models import ( Distance, @@ -22,6 +23,8 @@ from api_tool_indexer.constants import ApiToolIndexerConstants from api_tool_indexer.models import EnrichedEndpoint +logger = LokiLogger(service_name="api-tool-calling") + # Error messages _CLIENT_NOT_INITIALIZED = "Qdrant client not initialized" @@ -218,7 +221,8 @@ def upsert_endpoint_points(self, enriched_points: List[EnrichedEndpoint]) -> boo Payload fields stored on every point: endpoint_id, name, description, url, method, params, - enriched_context, service_id, point_type, example_text (example only) + enriched_context, service_id, point_type, cacheable, + cache_ttl_seconds, example_text (example only) Args: enriched_points: List of EnrichedEndpoint instances (examples + summary). @@ -256,6 +260,8 @@ def upsert_endpoint_points(self, enriched_points: List[EnrichedEndpoint]) -> boo "enriched_context": enriched.enriched_context, "service_id": enriched.service_id, "point_type": enriched.point_type, + "cacheable": enriched.cacheable, + "cache_ttl_seconds": enriched.cache_ttl_seconds, } if enriched.example_text is not None: payload["example_text"] = enriched.example_text diff --git a/src/contextual_retrieval/bm25_search.py b/src/contextual_retrieval/bm25_search.py index 7ec8ea9a..ef7b3831 100644 --- a/src/contextual_retrieval/bm25_search.py +++ b/src/contextual_retrieval/bm25_search.py @@ -6,10 +6,11 @@ """ from typing import List, Dict, Any, Optional, Set, TYPE_CHECKING -from loguru import logger +from src.loki_logger import LokiLogger from rank_bm25 import BM25Okapi import re import asyncio + from contextual_retrieval.contextual_retrieval_api_client import get_http_client_manager from contextual_retrieval.error_handler import SecureErrorHandler from contextual_retrieval.constants import ( @@ -20,6 +21,9 @@ ) from contextual_retrieval.config import ConfigLoader, ContextualRetrievalConfig +# Initialize Loki logger +logger = LokiLogger(service_name="bm25-search") + if TYPE_CHECKING: from contextual_retrieval.contextual_retrieval_api_client import HTTPClientManager @@ -165,7 +169,7 @@ async def search_bm25( logger.info(f"BM25 search found {len(results)} chunks") - # Detailed results at DEBUG level (loguru filters based on log level config) + # Detailed results at DEBUG level (filters based on log level config) logger.debug("=== BM25 SEARCH RESULTS BREAKDOWN ===") for i, chunk in enumerate(results[:10]): # Show top 10 results content_preview = ( diff --git a/src/contextual_retrieval/config.py b/src/contextual_retrieval/config.py index 49f78ef8..7e811e50 100644 --- a/src/contextual_retrieval/config.py +++ b/src/contextual_retrieval/config.py @@ -9,7 +9,8 @@ from typing import List import yaml from pathlib import Path -from loguru import logger +from src.loki_logger import LokiLogger + from contextual_retrieval.constants import ( HttpClientConstants, SearchConstants, @@ -17,6 +18,9 @@ BM25Constants, ) +# Initialize Loki logger +logger = LokiLogger(service_name="contextual-retrieval-config") + class HttpClientConfig(BaseModel): """HTTP client configuration.""" diff --git a/src/contextual_retrieval/contextual_retrieval.md b/src/contextual_retrieval/contextual_retrieval.md index ce3446c6..3f7c4e70 100644 --- a/src/contextual_retrieval/contextual_retrieval.md +++ b/src/contextual_retrieval/contextual_retrieval.md @@ -134,7 +134,8 @@ async def retrieve_contextual_chunks( ```python class HTTPClientManager: """Centralized HTTP client with connection pooling and resource management""" - + + class ServiceResilienceManager: """Circuit breaker implementation for fault tolerance""" ``` @@ -155,12 +156,14 @@ When the LLM Orchestration Service receives multiple simultaneous requests, the class QdrantContextualSearch: def __init__(self): self.client = httpx.AsyncClient() # New client per instance - + + class SmartBM25Search: def __init__(self): self.client = httpx.AsyncClient() # Another new client -# Result: + +# Result: # - 100+ HTTP connections for 10 concurrent requests # - Connection exhaustion # - Resource leaks @@ -171,19 +174,20 @@ class SmartBM25Search: ```python # GOOD: Shared HTTP client with connection pooling class HTTPClientManager: - _instance: Optional['HTTPClientManager'] = None # Singleton - + _instance: Optional["HTTPClientManager"] = None # Singleton + async def get_client(self) -> httpx.AsyncClient: if self._client is None: self._client = httpx.AsyncClient( limits=httpx.Limits( - max_connections=100, # Total pool size - max_keepalive_connections=20 # Reuse connections + max_connections=100, # Total pool size + max_keepalive_connections=20, # Reuse connections ), - timeout=httpx.Timeout(30.0) + timeout=httpx.Timeout(30.0), ) return self._client + # Result: # - Single connection pool (100 connections max) # - Connection reuse across all components @@ -195,10 +199,10 @@ class HTTPClientManager: ```python class ServiceResilienceManager: def __init__(self, config): - self.failure_threshold = 3 # Open circuit after 3 failures - self.recovery_timeout = 60.0 # Try recovery after 60 seconds - self.state = "CLOSED" # CLOSED → OPEN → HALF_OPEN - + self.failure_threshold = 3 # Open circuit after 3 failures + self.recovery_timeout = 60.0 # Try recovery after 60 seconds + self.state = "CLOSED" # CLOSED → OPEN → HALF_OPEN + def can_execute(self) -> bool: """Prevents cascading failures during high load""" if self.state == "OPEN": @@ -217,16 +221,16 @@ class QdrantContextualSearch: def __init__(self, qdrant_url: str, config: ContextualRetrievalConfig): # Uses shared HTTP client manager self.http_manager = HTTPClientManager() - + async def search_contextual_embeddings(self, embedding, collections, limit): # All Qdrant API calls use managed HTTP client client = await self.http_manager.get_client() - + # Circuit breaker protects against Qdrant downtime response = await self.http_manager.execute_with_circuit_breaker( method="POST", url=f"{self.qdrant_url}/collections/{collection}/points/search", - json=search_payload + json=search_payload, ) ``` @@ -236,12 +240,10 @@ class QdrantContextualSearch: async def get_embedding_for_query(self, query: str): # Uses shared HTTP client for LLM Orchestration API calls client = await self.http_manager.get_client() - + # Resilient embedding generation response = await self.http_manager.execute_with_circuit_breaker( - method="POST", - url="/embeddings", - json={"inputs": [query]} + method="POST", url="/embeddings", json={"inputs": [query]} ) ``` @@ -275,28 +277,28 @@ async def retry_http_request( url: str, max_retries: int = 3, retry_delay: float = 1.0, - backoff_factor: float = 2.0 + backoff_factor: float = 2.0, ) -> Optional[httpx.Response]: """ Handles transient failures gracefully: - Network hiccups during high load - - Temporary service unavailability + - Temporary service unavailability - Rate limiting responses """ for attempt in range(max_retries + 1): try: response = await client.request(method, url, **kwargs) - + # Success - return immediately if response.status_code < 400: return response - + # 4xx errors (client errors) - don't retry if 400 <= response.status_code < 500: return response - + # 5xx errors (server errors) - retry with backoff - + except (httpx.ConnectError, httpx.TimeoutException) as e: if attempt < max_retries: await asyncio.sleep(retry_delay) @@ -312,11 +314,11 @@ def client_stats(self) -> Dict[str, Any]: """Monitor connection pool health during high load""" return { "status": "active", - "pool_connections": 45, # Currently active connections - "keepalive_connections": 15, # Reusable connections + "pool_connections": 45, # Currently active connections + "keepalive_connections": 15, # Reusable connections "circuit_breaker_state": "CLOSED", "total_requests": 1247, - "failed_requests": 3 + "failed_requests": 3, } ``` @@ -356,19 +358,19 @@ class ContextualRetriever: ```python def detect_optimal_collections(query: str) -> List[str]: collections = [] - + # Check Azure keywords if any(keyword in query.lower() for keyword in AZURE_KEYWORDS): collections.append("azure_contextual_collection") - - # Check AWS keywords + + # Check AWS keywords if any(keyword in query.lower() for keyword in AWS_KEYWORDS): collections.append("aws_contextual_collection") - + # Default fallback if not collections: collections = ["azure_contextual_collection", "aws_contextual_collection"] - + return collections ``` @@ -451,9 +453,7 @@ Where: ```python # 1. Initialize ContextualRetriever retriever = ContextualRetriever( - qdrant_url="http://qdrant:6333", - environment="production", - connection_id="user123" + qdrant_url="http://qdrant:6333", environment="production", connection_id="user123" ) # 2. Initialize components @@ -467,7 +467,7 @@ original_question = "How do I set up Azure authentication?" refined_questions = [ "What are the steps to configure Azure Active Directory authentication?", "How to implement OAuth2 with Azure AD?", - "Azure authentication setup guide" + "Azure authentication setup guide", ] ``` @@ -475,8 +475,7 @@ refined_questions = [ ```python # Dynamic provider detection collections = await provider_detection.detect_optimal_collections( - environment="production", - connection_id="user123" + environment="production", connection_id="user123" ) # Result: ["azure_contextual_collection"] (Azure keywords detected) ``` @@ -488,10 +487,8 @@ if config.enable_parallel_search: semantic_task = _semantic_search( original_question, refined_questions, collections, 40, env, conn_id ) - bm25_task = _bm25_search( - original_question, refined_questions, 40 - ) - + bm25_task = _bm25_search(original_question, refined_questions, 40) + semantic_results, bm25_results = await asyncio.gather( semantic_task, bm25_task, return_exceptions=True ) @@ -507,7 +504,7 @@ batch_embeddings = qdrant_search.get_embeddings_for_queries_batch( queries=all_queries, llm_service=cached_llm_service, environment="production", - connection_id="user123" + connection_id="user123", ) # Parallel search execution @@ -542,19 +539,21 @@ deduplicated_bm25 = deduplicate_bm25_results(bm25_results) # Dynamic Rank Fusion fused_results = rank_fusion.fuse_results( semantic_results=semantic_results, # 40 results - bm25_results=bm25_results, # 40 results - final_top_n=12 # Return top 12 + bm25_results=bm25_results, # 40 results + final_top_n=12, # Return top 12 ) # RRF calculation for each document for doc_id in all_document_ids: semantic_rank = get_rank_in_results(doc_id, semantic_results) bm25_rank = get_rank_in_results(doc_id, bm25_results) - + rrf_score = 0 - if semantic_rank: rrf_score += 1 / (60 + semantic_rank) - if bm25_rank: rrf_score += 1 / (60 + bm25_rank) - + if semantic_rank: + rrf_score += 1 / (60 + semantic_rank) + if bm25_rank: + rrf_score += 1 / (60 + bm25_rank) + doc_scores[doc_id] = rrf_score # Sort by RRF score and return top N @@ -574,10 +573,10 @@ for result in fused_results: "retrieval_type": "contextual", "semantic_score": result.get("normalized_score"), "bm25_score": result.get("normalized_bm25_score"), - "fused_score": result.get("fused_score") + "fused_score": result.get("fused_score"), }, "score": result.get("fused_score"), - "id": result.get("chunk_id") + "id": result.get("chunk_id"), } formatted_results.append(formatted_chunk) @@ -613,16 +612,14 @@ selected_collections = ["azure_contextual_collection"] # Batch embedding generation queries = [ "How do I set up Azure authentication?", - "What are the steps to configure Azure Active Directory authentication?", + "What are the steps to configure Azure Active Directory authentication?", "How to implement OAuth2 with Azure AD?", - "Azure authentication setup guide" + "Azure authentication setup guide", ] # LLM API call for batch embeddings embeddings = llm_service.create_embeddings_for_indexer( - texts=queries, - model="text-embedding-3-large", - environment="production" + texts=queries, model="text-embedding-3-large", environment="production" ) # Parallel search across queries @@ -632,7 +629,7 @@ semantic_results = [ "contextual_content": "This section covers Azure Active Directory authentication setup. To configure Azure AD authentication, you need to...", "score": 0.89, "document_url": "azure-auth-guide.pdf", - "source_query": "How do I set up Azure authentication?" + "source_query": "How do I set up Azure authentication?", }, # ... more results ] @@ -643,10 +640,10 @@ semantic_results = [ # BM25 lexical search bm25_results = [ { - "chunk_id": "azure_auth_002", + "chunk_id": "azure_auth_002", "contextual_content": "This guide explains Azure authentication implementation. Follow these steps to set up Azure AD...", "bm25_score": 8.42, - "document_url": "azure-implementation.md" + "document_url": "azure-implementation.md", }, # ... more results ] @@ -675,14 +672,14 @@ final_results = [ "text": "This section covers Azure Active Directory authentication setup. To configure Azure AD authentication, you need to register your application in the Azure portal, configure redirect URIs, and implement the OAuth2 flow...", "meta": { "source_file": "azure-auth-guide.pdf", - "chunk_id": "azure_auth_001", + "chunk_id": "azure_auth_001", "retrieval_type": "contextual", "semantic_score": 0.89, "bm25_score": 0.72, - "fused_score": 0.0323 + "fused_score": 0.0323, }, "score": 0.0323, - "id": "azure_auth_001" + "id": "azure_auth_001", } # ... 11 more chunks (final_top_n = 12) ] @@ -774,14 +771,12 @@ rank_fusion: def _initialize_contextual_retriever( self, environment: str, connection_id: Optional[str] ) -> ContextualRetriever: - qdrant_url = os.getenv('QDRANT_URL', 'http://qdrant:6333') - + qdrant_url = os.getenv("QDRANT_URL", "http://qdrant:6333") + contextual_retriever = ContextualRetriever( - qdrant_url=qdrant_url, - environment=environment, - connection_id=connection_id + qdrant_url=qdrant_url, environment=environment, connection_id=connection_id ) - + return contextual_retriever ``` @@ -791,14 +786,12 @@ def _initialize_contextual_retriever( def _execute_orchestration_pipeline(self, request, components, costs_metric): # Step 1: Refine user prompt refined_output = self._refine_user_prompt(...) - - # Step 2: Retrieve contextual chunks + + # Step 2: Retrieve contextual chunks relevant_chunks = self._safe_retrieve_contextual_chunks( - components["contextual_retriever"], - refined_output, - request + components["contextual_retriever"], refined_output, request ) - + # Step 3: Generate response with chunks response = self._generate_response_with_chunks( relevant_chunks, refined_output, request @@ -810,26 +803,26 @@ def _execute_orchestration_pipeline(self, request, components, costs_metric): def _safe_retrieve_contextual_chunks( self, contextual_retriever: Optional[ContextualRetriever], - refined_output: PromptRefinerOutput, + refined_output: PromptRefinerOutput, request: OrchestrationRequest, ) -> Optional[List[Dict]]: - + async def async_retrieve(): # Initialize if needed if not contextual_retriever.initialized: success = await contextual_retriever.initialize() if not success: return None - + # Retrieve chunks chunks = await contextual_retriever.retrieve_contextual_chunks( original_question=refined_output.original_question, refined_questions=refined_output.refined_questions, environment=request.environment, - connection_id=request.connection_id + connection_id=request.connection_id, ) return chunks - + # Run async in sync context return asyncio.run(async_retrieve()) ``` @@ -905,18 +898,18 @@ Circuit Breaker: CLOSED (no failures) { "total_pool_size": 100, "active_connections": { - "qdrant_searches": 35, # Vector searches - "llm_embeddings": 25, # Embedding generation - "bm25_operations": 10, # Lexical searches - "keepalive_reserved": 20, # Ready for reuse - "available": 10 # Unused capacity + "qdrant_searches": 35, # Vector searches + "llm_embeddings": 25, # Embedding generation + "bm25_operations": 10, # Lexical searches + "keepalive_reserved": 20, # Ready for reuse + "available": 10, # Unused capacity }, "efficiency_metrics": { "connection_reuse_rate": "85%", - "average_connection_lifetime": "45s", + "average_connection_lifetime": "45s", "failed_connections": 0, - "circuit_breaker_activations": 0 - } + "circuit_breaker_activations": 0, + }, } ``` @@ -945,24 +938,24 @@ Total System Downtime: 90 seconds ```python def handle_qdrant_failure_scenario(): """Real-world circuit breaker behavior""" - + # CLOSED → OPEN (after 3 failures) failures = [ "Request 1: Qdrant timeout (30s)", - "Request 2: Qdrant timeout (30s)", - "Request 3: Qdrant timeout (30s)" # Circuit opens here + "Request 2: Qdrant timeout (30s)", + "Request 3: Qdrant timeout (30s)", # Circuit opens here ] - + # OPEN state (60 seconds) blocked_requests = [ "Request 4-47: Immediate failure (0.1s each)", - "Total blocked: 44 requests in 4.4 seconds" + "Total blocked: 44 requests in 4.4 seconds", ] - + # HALF_OPEN → CLOSED (service recovery) recovery = [ "Request 48: Success (200ms) → Circuit CLOSED", - "Request 49-100: Normal operation resumed" + "Request 49-100: Normal operation resumed", ] ``` @@ -1004,14 +997,14 @@ def handle_qdrant_failure_scenario(): "original_question": "How do I set up Azure authentication?", "refined_questions": [ "What are the steps to configure Azure Active Directory authentication?", - "How to implement OAuth2 with Azure AD?", - "Azure authentication setup guide" + "How to implement OAuth2 with Azure AD?", + "Azure authentication setup guide", ], "environment": "production", "connection_id": "user123", - "topk_semantic": 40, # Optional - uses config default - "topk_bm25": 40, # Optional - uses config default - "final_top_n": 12 # Optional - uses config default + "topk_semantic": 40, # Optional - uses config default + "topk_bm25": 40, # Optional - uses config default + "final_top_n": 12, # Optional - uses config default } ``` @@ -1028,16 +1021,15 @@ def handle_qdrant_failure_scenario(): "retrieval_type": "contextual", "primary_source": "azure", "semantic_score": 0.89, - "bm25_score": 0.72, - "fused_score": 0.0323 + "bm25_score": 0.72, + "fused_score": 0.0323, }, - # Legacy compatibility fields "id": "azure_auth_001", "score": 0.0323, "content": "This section covers Azure Active Directory authentication setup...", "document_url": "azure-auth-guide.pdf", - "retrieval_type": "contextual" + "retrieval_type": "contextual", } # ... 11 more chunks ] @@ -1052,15 +1044,15 @@ refined_output = PromptRefinerOutput( original_question="How do I set up Azure authentication?", refined_questions=[...], is_off_topic=False, - reasoning="User asking about Azure authentication setup" + reasoning="User asking about Azure authentication setup", ) # OrchestrationRequest request = OrchestrationRequest( - message="How do I set up Azure authentication?", + message="How do I set up Azure authentication?", environment="production", connection_id="user123", - chatId="chat456" + chatId="chat456", ) ``` @@ -1070,8 +1062,8 @@ request = OrchestrationRequest( contextual_chunks = [ { "text": "contextual content...", # This is what ResponseGenerator uses - "meta": {...}, # Source information and scores - "score": 0.0323 # Final fused score + "meta": {...}, # Source information and scores + "score": 0.0323, # Final fused score } ] ``` @@ -1092,8 +1084,8 @@ class RateLimiter: #### 2. Enhanced Caching ```python class EmbeddingCache: - max_size: int = 1000 # LRU cache for embeddings - ttl_seconds: int = 3600 # 1 hour TTL + max_size: int = 1000 # LRU cache for embeddings + ttl_seconds: int = 3600 # 1 hour TTL ``` #### 3. Connection Pool Optimization diff --git a/src/contextual_retrieval/contextual_retrieval_api_client.py b/src/contextual_retrieval/contextual_retrieval_api_client.py index 0de14558..f962e3d6 100644 --- a/src/contextual_retrieval/contextual_retrieval_api_client.py +++ b/src/contextual_retrieval/contextual_retrieval_api_client.py @@ -8,8 +8,9 @@ import asyncio from typing import Optional, Dict, Any import httpx -from loguru import logger +from src.loki_logger import LokiLogger import time + from contextual_retrieval.error_handler import SecureErrorHandler from contextual_retrieval.constants import ( HttpClientConstants, @@ -20,6 +21,9 @@ ) from contextual_retrieval.config import ConfigLoader, ContextualRetrievalConfig +# Initialize Loki logger +logger = LokiLogger(service_name="contextual-retrieval-api-client") + class ServiceResilienceManager: """Service resilience manager with circuit breaker functionality for HTTP requests.""" diff --git a/src/contextual_retrieval/contextual_retriever.py b/src/contextual_retrieval/contextual_retriever.py index bdb61eb1..7f0f5796 100644 --- a/src/contextual_retrieval/contextual_retriever.py +++ b/src/contextual_retrieval/contextual_retriever.py @@ -11,10 +11,11 @@ """ from typing import List, Dict, Any, Optional, Union, TYPE_CHECKING -from loguru import logger +from src.loki_logger import LokiLogger import asyncio import time from langfuse import observe + from contextual_retrieval.config import ConfigLoader, ContextualRetrievalConfig # Type checking import to avoid circular dependency at runtime @@ -26,6 +27,9 @@ from contextual_retrieval.bm25_search import SmartBM25Search from contextual_retrieval.rank_fusion import DynamicRankFusion +# Initialize Loki logger +logger = LokiLogger(service_name="contextual-retriever") + class ContextualRetriever: """ diff --git a/src/contextual_retrieval/error_handler.py b/src/contextual_retrieval/error_handler.py index 08fac2e7..650f4402 100644 --- a/src/contextual_retrieval/error_handler.py +++ b/src/contextual_retrieval/error_handler.py @@ -8,9 +8,12 @@ import re from typing import Dict, Any, Optional, Union from urllib.parse import urlparse, urlunparse -from loguru import logger +from src.loki_logger import LokiLogger import httpx +# Initialize Loki logger +logger = LokiLogger(service_name="error-handler") + class SecureErrorHandler: """ diff --git a/src/contextual_retrieval/provider_detection.py b/src/contextual_retrieval/provider_detection.py index 8abb4d18..78796b85 100644 --- a/src/contextual_retrieval/provider_detection.py +++ b/src/contextual_retrieval/provider_detection.py @@ -8,7 +8,8 @@ """ from typing import List, Optional, Dict, Any, TYPE_CHECKING -from loguru import logger +from src.loki_logger import LokiLogger + from contextual_retrieval.contextual_retrieval_api_client import get_http_client_manager from contextual_retrieval.error_handler import SecureErrorHandler from contextual_retrieval.constants import ( @@ -18,6 +19,9 @@ ) from contextual_retrieval.config import ConfigLoader, ContextualRetrievalConfig +# Initialize Loki logger +logger = LokiLogger(service_name="provider-detection") + if TYPE_CHECKING: from contextual_retrieval.contextual_retrieval_api_client import HTTPClientManager diff --git a/src/contextual_retrieval/qdrant_search.py b/src/contextual_retrieval/qdrant_search.py index 31515f32..28f45d89 100644 --- a/src/contextual_retrieval/qdrant_search.py +++ b/src/contextual_retrieval/qdrant_search.py @@ -6,8 +6,10 @@ """ from typing import List, Dict, Any, Optional, Protocol, TYPE_CHECKING -from loguru import logger +from src.loki_logger import LokiLogger + import asyncio + from contextual_retrieval.contextual_retrieval_api_client import get_http_client_manager from contextual_retrieval.error_handler import SecureErrorHandler from contextual_retrieval.constants import ( @@ -17,6 +19,9 @@ ) from contextual_retrieval.config import ConfigLoader, ContextualRetrievalConfig +# Initialize Loki logger +logger = LokiLogger(service_name="qdrant-search") + if TYPE_CHECKING: from contextual_retrieval.contextual_retrieval_api_client import HTTPClientManager @@ -151,7 +156,7 @@ async def search_contextual_embeddings_direct( f"Semantic search found {len(all_results)} chunks across {len(collections)} collections" ) - # Detailed results at DEBUG level (loguru filters based on log level config) + # Detailed results at DEBUG level (filters based on log level config) logger.debug("=== SEMANTIC SEARCH RESULTS BREAKDOWN ===") for i, chunk in enumerate(all_results[:10]): # Show top 10 results content_preview = ( diff --git a/src/contextual_retrieval/rank_fusion.py b/src/contextual_retrieval/rank_fusion.py index acea0aa2..d6a18ab9 100644 --- a/src/contextual_retrieval/rank_fusion.py +++ b/src/contextual_retrieval/rank_fusion.py @@ -6,10 +6,14 @@ """ from typing import List, Dict, Any, Optional -from loguru import logger +from src.loki_logger import LokiLogger + from contextual_retrieval.constants import QueryTypeConstants from contextual_retrieval.config import ConfigLoader, ContextualRetrievalConfig +# Initialize Loki logger +logger = LokiLogger(service_name="rank-fusion") + class DynamicRankFusion: """Dynamic score fusion without hardcoded collection weights.""" @@ -65,7 +69,7 @@ def fuse_results( logger.info(f"Fusion completed: {len(final_results)} final results") - # Detailed results at DEBUG level (loguru filters based on log level config) + # Detailed results at DEBUG level (filters based on log level config) logger.debug("=== RANK FUSION FINAL RESULTS ===") for i, chunk in enumerate(final_results): content_preview_len = self._config.rank_fusion.content_preview_length diff --git a/src/guardrails/dspy_nemo_adapter.py b/src/guardrails/dspy_nemo_adapter.py index da2617d7..700f7577 100644 --- a/src/guardrails/dspy_nemo_adapter.py +++ b/src/guardrails/dspy_nemo_adapter.py @@ -7,7 +7,9 @@ from typing import Any, Dict, List, Optional, Union, cast, Iterator, AsyncIterator import asyncio import dspy -from loguru import logger +from src.loki_logger import LokiLogger +from langfuse import observe +from src.utils.observation_utils import update_observation_safe from langchain_core.callbacks.manager import ( CallbackManagerForLLMRun, @@ -16,6 +18,10 @@ from langchain_core.language_models.llms import LLM from langchain_core.outputs import GenerationChunk from src.guardrails.guardrails_llm_configs import TEMPERATURE, MAX_TOKENS, MODEL_NAME +from src.utils.cost_utils import get_lm_usage_since + +# Initialize Loki logger +logger = LokiLogger(service_name="dspy-nemo-adapter") class DSPyNeMoLLM(LLM): @@ -122,6 +128,7 @@ def _extract_chunk_text(self, chunk: Any) -> str: return "" + @observe(name="nemo_guardrail_llm_call", as_type="generation") def _call( self, prompt: str, @@ -135,6 +142,11 @@ def _call( This is the standard path for NeMo Guardrails when streaming is disabled. Call DSPy's LM directly with the prompt. """ + history_length_before = 0 + lm_for_usage = dspy.settings.lm + if lm_for_usage and hasattr(lm_for_usage, "history"): + history_length_before = len(lm_for_usage.history) + try: lm = self._get_dspy_lm() @@ -148,12 +160,32 @@ def _call( # DSPy LM call - returns text directly response = lm(prompt, **call_kwargs) - return self._extract_text_from_response(response) + response_text = self._extract_text_from_response(response) + + usage_info = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "prompt_preview": prompt[:500], + "max_tokens": call_kwargs["max_tokens"], + "temperature": call_kwargs["temperature"], + }, + output_data={"response_preview": response_text[:500]}, + metadata={ + "model": self.model_name, + "usage": usage_info, + "num_calls": usage_info.get("num_calls", 0), + "streaming": False, + "guardrail_provider": "nemo", + }, + ) + + return response_text except Exception as e: logger.error(f"Error in DSPyNeMoLLM._call: {str(e)}") raise RuntimeError(f"LLM generation failed: {str(e)}") from e + @observe(name="nemo_guardrail_llm_call", as_type="generation") async def _acall( self, prompt: str, @@ -167,6 +199,11 @@ async def _acall( Uses asyncio.to_thread to prevent blocking the event loop. This is critical because DSPy's LM is synchronous and makes network calls. """ + history_length_before = 0 + lm_for_usage = dspy.settings.lm + if lm_for_usage and hasattr(lm_for_usage, "history"): + history_length_before = len(lm_for_usage.history) + try: lm = self._get_dspy_lm() @@ -180,7 +217,26 @@ async def _acall( # Run in thread to avoid blocking response = await asyncio.to_thread(lm, prompt, **call_kwargs) - return self._extract_text_from_response(response) + response_text = self._extract_text_from_response(response) + + usage_info = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "prompt_preview": prompt[:500], + "max_tokens": call_kwargs["max_tokens"], + "temperature": call_kwargs["temperature"], + }, + output_data={"response_preview": response_text[:500]}, + metadata={ + "model": self.model_name, + "usage": usage_info, + "num_calls": usage_info.get("num_calls", 0), + "streaming": False, + "guardrail_provider": "nemo", + }, + ) + + return response_text except Exception as e: logger.error(f"Error in DSPyNeMoLLM._acall: {str(e)}") diff --git a/src/guardrails/nemo_rails_adapter.py b/src/guardrails/nemo_rails_adapter.py index 17f6585e..53749149 100644 --- a/src/guardrails/nemo_rails_adapter.py +++ b/src/guardrails/nemo_rails_adapter.py @@ -1,8 +1,7 @@ from typing import Any, Dict, Optional, AsyncIterator, cast, Type import asyncio -from loguru import logger +from src.loki_logger import LokiLogger from pydantic import BaseModel, Field - from nemoguardrails import LLMRails, RailsConfig from nemoguardrails.llm.providers import register_llm_provider from langchain_core.language_models.llms import BaseLLM @@ -13,6 +12,9 @@ import dspy import re +# Initialize Loki logger +logger = LokiLogger(service_name="nemo-rails-adapter") + class GuardrailCheckResult(BaseModel): """Result from a guardrail check.""" @@ -92,8 +94,24 @@ def _ensure_initialized(self) -> None: from llm_orchestrator_config.llm_manager import LLMManager + # Resolve vault_uuid from DB if not provided + connection_id = self.connection_id + if not connection_id: + from src.utils.connection_id_fetcher import get_connection_id_fetcher + + fetcher = get_connection_id_fetcher() + connection_id = fetcher.fetch_vault_uuid_sync(self.environment) + if not connection_id: + raise ValueError( + f"No {self.environment} connection found in database. " + f"Cannot initialize guardrails without a configured LLM connection." + ) + logger.debug( + f"Resolved {self.environment} vault_uuid for guardrails: {connection_id}" + ) + llm_manager = LLMManager( - environment=self.environment, connection_id=self.connection_id + environment=self.environment, connection_id=connection_id ) llm_manager.ensure_global_config() @@ -112,6 +130,8 @@ def _ensure_initialized(self) -> None: rails_config.streaming = True + self._validate_output_streaming_config(rails_config) + if metadata.get("optimized", False): version = metadata.get("version", "unknown") metrics = metadata.get("metrics", {}) @@ -151,6 +171,44 @@ def _ensure_initialized(self) -> None: logger.exception("Full traceback:") raise + @staticmethod + def _validate_output_streaming_config(rails_config: RailsConfig) -> None: + """ + Guard against a buffer configuration that degenerates into one guardrail + LLM call per streamed token. + + NeMo's RollingBuffer drains with ``buffer = buffer[-context_size:]`` after + each flush. If ``context_size >= chunk_size`` the buffer never shrinks below + the ``len(buffer) >= chunk_size`` flush threshold, so every subsequent token + flushes on its own and triggers a full, serial ``self_check_output`` LLM + round-trip. That makes streaming latency and cost scale linearly with answer + length and is invisible without this check. + """ + streaming_config = getattr( + getattr(getattr(rails_config, "rails", None), "output", None), + "streaming", + None, + ) + if streaming_config is None or not getattr(streaming_config, "enabled", False): + return + + chunk_size = getattr(streaming_config, "chunk_size", 0) + context_size = getattr(streaming_config, "context_size", 0) + + if chunk_size > 0 and context_size >= chunk_size: + clamped = max(1, chunk_size // 4) + logger.error( + f"Invalid output-rails streaming config: context_size={context_size} " + f">= chunk_size={chunk_size}. This degrades to one guardrail LLM call " + f"per streamed token. Clamping context_size to {clamped}." + ) + streaming_config.context_size = clamped + else: + logger.debug( + f"Output-rails streaming buffer OK: chunk_size={chunk_size}, " + f"context_size={context_size}" + ) + async def check_input_async(self, user_message: str) -> GuardrailCheckResult: """ Check user input against guardrails (async version for streaming). @@ -365,6 +423,9 @@ async def stream_with_guardrails( logger.debug(f"Generator type: {type(bot_message_generator)}") chunk_count = 0 + stream_started_at = asyncio.get_running_loop().time() + last_chunk_at = stream_started_at + max_gap_seconds = 0.0 logger.info("Calling _rails.stream_async with generator parameter...") @@ -374,16 +435,33 @@ async def stream_with_guardrails( ): chunk_count += 1 - if chunk_count <= 10: + now = asyncio.get_running_loop().time() + gap = now - last_chunk_at + last_chunk_at = now + max_gap_seconds = max(max_gap_seconds, gap) + + # Log the head of the stream, then periodically. Logging only the + # first N chunks hides per-token guardrail stalls, which look like + # a total outage in the log while tokens are actually trickling. + if chunk_count <= 10 or chunk_count % 50 == 0: logger.debug( - f"[Chunk {chunk_count}] Validated and yielded: {repr(chunk)}" + f"[Chunk {chunk_count}] Validated and yielded: {repr(chunk)} " + f"(gap={gap:.2f}s, elapsed={now - stream_started_at:.2f}s)" ) yield chunk + total_elapsed = asyncio.get_running_loop().time() - stream_started_at logger.info( - f"NeMo streaming completed successfully - {chunk_count} chunks streamed" + f"NeMo streaming completed successfully - {chunk_count} chunks streamed " + f"in {total_elapsed:.2f}s (max inter-chunk gap {max_gap_seconds:.2f}s)" ) + if max_gap_seconds > 10.0: + logger.warning( + f"Output-rails streaming stalled for {max_gap_seconds:.2f}s between " + f"chunks. Check rails.output.streaming (context_size must be < " + f"chunk_size) — a per-token guardrail call causes this." + ) except Exception as e: logger.error(f"Error in stream_with_guardrails: {str(e)}") diff --git a/src/guardrails/optimized_guardrails_loader.py b/src/guardrails/optimized_guardrails_loader.py index aef76727..51a65906 100644 --- a/src/guardrails/optimized_guardrails_loader.py +++ b/src/guardrails/optimized_guardrails_loader.py @@ -6,7 +6,10 @@ from pathlib import Path from typing import Optional, Dict, Any, Tuple import json -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="optimized-guardrails-loader") class OptimizedGuardrailsLoader: diff --git a/src/guardrails/rails_config.yaml b/src/guardrails/rails_config.yaml index 42116e9a..a52ffd26 100644 --- a/src/guardrails/rails_config.yaml +++ b/src/guardrails/rails_config.yaml @@ -23,7 +23,11 @@ rails: streaming: enabled: True chunk_size: 200 - context_size: 300 + # context_size MUST be < chunk_size. NeMo's RollingBuffer drains with + # buffer[-context_size:]; if context_size >= chunk_size the buffer never + # shrinks below the flush threshold and every subsequent token triggers its + # own self_check_output LLM call. 50 is the NeMo default. + context_size: 50 stream_first: False prompts: diff --git a/src/guardrails/readme.md b/src/guardrails/readme.md index 7a69e931..17a97bb0 100644 --- a/src/guardrails/readme.md +++ b/src/guardrails/readme.md @@ -118,12 +118,12 @@ Cost: $0.000156 (7 tokens) #### 3. **GuardrailCheckResult** (Pydantic Model) ```python class GuardrailCheckResult(BaseModel): - allowed: bool # True if content passes - verdict: str # "yes" = blocked, "no" = allowed - content: str # Response message + allowed: bool # True if content passes + verdict: str # "yes" = blocked, "no" = allowed + content: str # Response message blocked_by_rail: Optional[str] # Exception type if blocked - reason: Optional[str] # Explanation - error: Optional[str] # Error message if failed + reason: Optional[str] # Explanation + error: Optional[str] # Error message if failed usage: Dict[str, Union[float, int]] # Cost tracking ``` @@ -136,8 +136,8 @@ When `enable_rails_exceptions: true` in config: "role": "exception", "content": { "type": "InputRailException", - "message": "I'm not able to respond to that" - } + "message": "I'm not able to respond to that", + }, } ``` @@ -167,11 +167,11 @@ result.usage = usage_info # Contains: total_cost, tokens, num_calls **Usage Dictionary Structure**: ```python { - "total_cost": 0.000245, # USD + "total_cost": 0.000245, # USD "total_prompt_tokens": 8, "total_completion_tokens": 2, "total_tokens": 10, - "num_calls": 1 + "num_calls": 1, } ``` @@ -181,10 +181,10 @@ result.usage = usage_info # Contains: total_cost, tokens, num_calls ```python costs_metric = { - "input_guardrails": {...}, # Step 1 - "prompt_refiner": {...}, # Step 2 - "response_generator": {...}, # Step 4 - "output_guardrails": {...} # Step 5 + "input_guardrails": {...}, # Step 1 + "prompt_refiner": {...}, # Step 2 + "response_generator": {...}, # Step 4 + "output_guardrails": {...}, # Step 5 } # Step 3 (retrieval) has no LLM cost @@ -197,7 +197,7 @@ costs_metric = { if not input_result.allowed: return OrchestrationResponse( inputGuardFailed=True, - content=input_result.content # Refusal message + content=input_result.content, # Refusal message ) # Saves costs: no refinement, retrieval, or generation ``` diff --git a/src/intent_data_enrichment/api_client.py b/src/intent_data_enrichment/api_client.py index 31ed96e2..bfcd240d 100644 --- a/src/intent_data_enrichment/api_client.py +++ b/src/intent_data_enrichment/api_client.py @@ -4,11 +4,14 @@ import httpx from typing import List, Optional from types import TracebackType -from loguru import logger +from src.loki_logger import LokiLogger from intent_data_enrichment.constants import EnrichmentConstants from intent_data_enrichment.models import ServiceData +# Initialize Loki logger +logger = LokiLogger(service_name="intent-enrichment-api-client") + class LLMAPIClient: """Client for calling LLM Orchestration Service endpoints.""" @@ -17,7 +20,7 @@ def __init__( self, api_base_url: str = EnrichmentConstants.DEFAULT_API_BASE_URL, environment: str = EnrichmentConstants.DEFAULT_ENVIRONMENT, - connection_id: str = EnrichmentConstants.DEFAULT_CONNECTION_ID, + connection_id: Optional[str] = EnrichmentConstants.DEFAULT_CONNECTION_ID, max_retries: int = EnrichmentConstants.MAX_RETRIES, retry_delay_base: int = EnrichmentConstants.RETRY_DELAY_BASE, timeout: int = EnrichmentConstants.REQUEST_TIMEOUT, diff --git a/src/intent_data_enrichment/constants.py b/src/intent_data_enrichment/constants.py index f506880a..c85d736c 100644 --- a/src/intent_data_enrichment/constants.py +++ b/src/intent_data_enrichment/constants.py @@ -10,7 +10,9 @@ class EnrichmentConstants: # API Configuration DEFAULT_API_BASE_URL = "http://llm-orchestration-service:8100" DEFAULT_ENVIRONMENT = "production" - DEFAULT_CONNECTION_ID = "gpt-4o-mini" + # None → orchestration service resolves the embedding model via the + # DB-fetched vault UUID (path: embeddings/connections/{provider}/{vault_uuid}). + DEFAULT_CONNECTION_ID = None # Retry Configuration MAX_RETRIES = 3 diff --git a/src/intent_data_enrichment/main_enrichment.py b/src/intent_data_enrichment/main_enrichment.py index b96e0d2f..9aff2959 100644 --- a/src/intent_data_enrichment/main_enrichment.py +++ b/src/intent_data_enrichment/main_enrichment.py @@ -14,13 +14,17 @@ import json import argparse import asyncio + from typing import List -from loguru import logger +from src.loki_logger import LokiLogger from intent_data_enrichment.models import ServiceData, EnrichedService, EnrichmentResult from intent_data_enrichment.api_client import LLMAPIClient from intent_data_enrichment.qdrant_manager import QdrantManager +# Initialize Loki logger +logger = LokiLogger(service_name="intent-enrichment-main") + # Import sparse encoder from tool_classifier (shared module) sys.path.insert(0, "/app/src") try: diff --git a/src/intent_data_enrichment/qdrant_manager.py b/src/intent_data_enrichment/qdrant_manager.py index d5593836..0bf0c744 100644 --- a/src/intent_data_enrichment/qdrant_manager.py +++ b/src/intent_data_enrichment/qdrant_manager.py @@ -2,7 +2,7 @@ import uuid from typing import Optional, List -from loguru import logger +from src.loki_logger import LokiLogger from qdrant_client import QdrantClient from qdrant_client.models import ( Distance, @@ -20,6 +20,9 @@ from intent_data_enrichment.constants import EnrichmentConstants from intent_data_enrichment.models import EnrichedService +# Initialize Loki logger +logger = LokiLogger(service_name="intent-qdrant-manager") + # Error messages _CLIENT_NOT_INITIALIZED = "Qdrant client not initialized" diff --git a/src/llm_orchestration_service.py b/src/llm_orchestration_service.py index 7601c603..40b4dd0e 100644 --- a/src/llm_orchestration_service.py +++ b/src/llm_orchestration_service.py @@ -4,12 +4,12 @@ import os import time import asyncio -from loguru import logger +import threading +from src.loki_logger import LokiLogger from langfuse import Langfuse, observe import dspy from datetime import datetime import json as json_module -import threading from llm_orchestrator_config.llm_manager import LLMManager from models.request_models import ( @@ -46,7 +46,11 @@ from src.vector_indexer.constants import ResponseGenerationConstants from src.utils.error_utils import generate_error_id, log_error_with_context from src.utils.stream_manager import stream_manager, StreamContext -from src.utils.cost_utils import calculate_total_costs, get_lm_usage_since +from src.utils.cost_utils import ( + calculate_total_costs, + get_lm_usage_since, + get_lm_usage_since_split, +) if TYPE_CHECKING: from src.llm_orchestrator_config.embedding_manager import EmbeddingManager @@ -60,6 +64,9 @@ from src.utils.language_detector import detect_language, get_language_name from src.utils.prompt_config_loader import PromptConfigurationLoader from src.utils.query_validator import validate_query_basic +from src.utils.sse_utils import extract_content_from_sse +from src.utils.conversation_history_store import should_save_history, save_history_round +from src.utils.conversation_history_helpers import get_conversation_history from src.guardrails import NeMoRailsAdapter, GuardrailCheckResult from src.contextual_retrieval import ContextualRetriever from src.contextual_retrieval.bm25_search import SmartBM25Search @@ -68,10 +75,29 @@ ContextualRetrievalFailureError, ) from src.llm_orchestrator_config.feature_flags import FeatureFlags -from src.tool_classifier import ToolClassifier +from src.tool_classifier import ToolClassifier, WorkflowType from src.tool_classifier.constants import SERVICE_STEP_PREFIXES from src.tool_classifier.workflows.service_workflow import ServiceWorkflowExecutor +# Initialize Loki logger for orchestration service +logger = LokiLogger(service_name="llm-orchestration-service") + +REFERENCES_SECTION_HEADER = "\n\n**References:**\n" + +# Set of content strings that must NOT be persisted in conversation history. +# Covers all multilingual error / OOS / guardrail-violation messages so that +# failed or blocked exchanges are never written to Redis. +_HISTORY_EXCLUDED_MESSAGES: frozenset[str] = frozenset( + { + *OUT_OF_SCOPE_MESSAGES.values(), + *TECHNICAL_ISSUE_MESSAGES.values(), + *INPUT_GUARDRAIL_VIOLATION_MESSAGES.values(), + *OUTPUT_GUARDRAIL_VIOLATION_MESSAGES.values(), + *QUERY_VALIDATION_FAILED_MESSAGES.values(), + STREAM_TOKEN_LIMIT_MESSAGE, + } +) + class LangfuseConfig: """Configuration for Langfuse integration.""" @@ -149,6 +175,11 @@ def __init__(self) -> None: # Workflow executors access it via self.orchestration_service.session_store. self.session_store: Any = None + # Redis-backed conversation history store. + # Set to None here; the FastAPI lifespan injects the live store after + # Redis initialises (app.state.orchestration_service.conversation_history_store = ...). + self.conversation_history_store: Any = None + # Shared BM25 search index pre-warmed at startup. # Populated by _prewarm_shared_bm25() which is called from the FastAPI # lifespan so it runs inside the async event loop. Until then it is None @@ -156,6 +187,9 @@ def __init__(self) -> None: # degradation path). self.shared_bm25_search: Optional[SmartBM25Search] = None + self._retriever_cache: Dict[tuple, ContextualRetriever] = {} + self._component_cache_lock = threading.Lock() + # Initialize shared guardrails adapters at startup (production and testing) self.shared_guardrails_adapters = ( self._initialize_shared_guardrails_at_startup() @@ -391,7 +425,7 @@ async def process_orchestration_request( if components["guardrails_adapter"]: start_time = time.time() input_blocked_response = await self.handle_input_guardrails( - components["guardrails_adapter"], request, {} + components["guardrails_adapter"], request, costs_metric ) time_metric["input_guardrails_check"] = time.time() - start_time @@ -428,7 +462,6 @@ async def process_orchestration_request( start_time = time.time() classification = await self.tool_classifier.classify( query=request.message, - conversation_history=request.conversationHistory, language=detected_language, request=request, ) @@ -487,25 +520,7 @@ async def process_orchestration_request( langfuse = self.langfuse_config.langfuse_client total_costs = calculate_total_costs(costs_metric) - total_input_tokens = sum( - c.get("total_prompt_tokens", 0) for c in costs_metric.values() - ) - total_output_tokens = sum( - c.get("total_completion_tokens", 0) for c in costs_metric.values() - ) - langfuse.update_current_generation( - model=components["llm_manager"] - .get_provider_info() - .get("model", "unknown"), - usage_details={ - "input": total_input_tokens, - "output": total_output_tokens, - "total": total_costs.get("total_tokens", 0), - }, - cost_details={ - "total": total_costs.get("total_cost", 0.0), - }, metadata={ "total_calls": total_costs.get("total_calls", 0), "cost_breakdown": costs_metric, @@ -515,6 +530,18 @@ async def process_orchestration_request( }, ) langfuse.flush() + + # Persist successful exchange to conversation history (non-streaming) + if should_save_history( + self.conversation_history_store, response, _HISTORY_EXCLUDED_MESSAGES + ): + await save_history_round( + self.conversation_history_store, + request.chatId, + request.message, + response.content, + ) + return response except Exception as e: @@ -542,7 +569,6 @@ async def process_orchestration_request( return self._create_error_response(request) - @observe(name="streaming_generation", as_type="generation", capture_output=False) async def stream_orchestration_response( self, request: OrchestrationRequest ) -> AsyncIterator[str]: @@ -583,6 +609,13 @@ async def stream_orchestration_response( costs_metric: Dict[str, Dict[str, Any]] = {} time_metric: Dict[str, float] = {} + # Capture DSPy history baseline before any LLM calls. + # Used at the end of the request to compute the total cost delta, + _lm = dspy.settings.lm + initial_history_length = ( + len(_lm.history) if _lm and hasattr(_lm, "history") else 0 + ) + # STEP 0: Detect language from user message (with timing) start_time = time.time() detected_language = detect_language(request.message) @@ -710,7 +743,6 @@ async def stream_orchestration_response( start_time = time.time() classification = await self.tool_classifier.classify( query=request.message, - conversation_history=request.conversationHistory, language=detected_language, request=request, ) @@ -723,9 +755,11 @@ async def stream_orchestration_response( # Route to appropriate workflow (streaming) # route_to_workflow returns AsyncIterator[str] when is_streaming=True - # Inject costs_metric into the classification context so the - # API Tool workflow can append its output guardrail costs. + # Inject costs_metric and pre-initialized components into the + # classification context so downstream workflows can reuse them + # without re-initializing (saves ~1.5s on fallback paths). classification.metadata["costs_metric"] = costs_metric + classification.metadata["components"] = components start_time = time.time() stream_result = await self.tool_classifier.route_to_workflow( classification=classification, @@ -735,17 +769,62 @@ async def stream_orchestration_response( ) time_metric["classifier.route"] = time.time() - start_time + # Accumulate content for history only on non-RAG workflows; + # RAG routes through _stream_rag_pipeline which has its own hook. + _save_classifier_history = ( + self.conversation_history_store is not None + and classification.workflow != WorkflowType.RAG + ) + _classifier_accumulated: list[str] = [] + # Tracks whether an excluded marker (OOS / guardrail violation / + # error) was observed at any point during the stream. When True + # the entire accumulated buffer is discarded so no partial content + # from before the blocked marker is ever written to Redis. + _history_blocked = False + async for sse_chunk in stream_result: yield sse_chunk + if _save_classifier_history and not _history_blocked: + extracted = extract_content_from_sse(sse_chunk) + if extracted is not None and extracted != "END": + if extracted in _HISTORY_EXCLUDED_MESSAGES: + # Excluded marker observed — discard any partial + # content accumulated before this point and stop + # accumulating for the rest of the stream. + _classifier_accumulated.clear() + _history_blocked = True + else: + _classifier_accumulated.append(extracted) # Successfully completed streaming through classifier logger.info( f"[{request.chatId}] [{stream_ctx.stream_id}] Tool classifier streaming completed" ) + # Persist conversation history (classifier streaming, non-RAG workflows) + if ( + _save_classifier_history + and not _history_blocked + and _classifier_accumulated + ): + await save_history_round( + self.conversation_history_store, + request.chatId, + request.message, + "".join(_classifier_accumulated), + ) + # Log costs and timings self.log_costs(costs_metric) log_step_timings(time_metric, request.chatId) + + # Budget update: use full DSPy history delta + _total_usage = get_lm_usage_since(initial_history_length) + self._update_connection_budget( + request.connection_id, + {"streaming_total": _total_usage}, + request.environment, + ) stream_ctx.mark_completed() return # Exit after successful classifier routing @@ -780,7 +859,15 @@ async def stream_orchestration_response( ): yield sse_chunk - # Pipeline completed successfully + # Pipeline completed successfully. + # Budget update: use full DSPy history delta (covers guardrails, + # refiner, and streaming generation across this request). + _total_usage = get_lm_usage_since(initial_history_length) + self._update_connection_budget( + request.connection_id, + {"streaming_total": _total_usage}, + request.environment, + ) return except Exception as e: @@ -796,9 +883,12 @@ async def stream_orchestration_response( self.log_costs(costs_metric) log_step_timings(time_metric, request.chatId) - # Update budget even on outer exception + # Budget update on outer exception using full DSPy history delta. + _total_usage = get_lm_usage_since(initial_history_length) self._update_connection_budget( - request.connection_id, costs_metric, request.environment + request.connection_id, + {"streaming_total": _total_usage}, + request.environment, ) if self.langfuse_config.langfuse_client: @@ -853,10 +943,16 @@ async def _stream_rag_pipeline( ) start_time = time.time() + conversation_history, conversation_summary = await get_conversation_history( + chat_id=request.chatId, + store=self.conversation_history_store, + fallback=request.conversationHistory, + ) refined_output, refiner_usage = self._refine_user_prompt( llm_manager=components["llm_manager"], original_message=request.message, - conversation_history=request.conversationHistory, + conversation_history=conversation_history, + conversation_summary=conversation_summary, ) time_metric["prompt_refiner"] = time.time() - start_time costs_metric["prompt_refiner"] = refiner_usage @@ -1056,7 +1152,7 @@ async def bot_response_generator() -> AsyncIterator[str]: # Send document references before END token doc_references = self._extract_document_references(relevant_chunks) if doc_references: - refs_text = "\n\n**References:**\n" + "\n".join( + refs_text = REFERENCES_SECTION_HEADER + "\n".join( f"{i + 1}. [{ref.document_url}]({ref.document_url})" for i, ref in enumerate(doc_references) ) @@ -1093,7 +1189,7 @@ async def bot_response_generator() -> AsyncIterator[str]: # Send document references before END token doc_references = self._extract_document_references(relevant_chunks) if doc_references: - refs_text = "\n\n**References:**\n" + "\n".join( + refs_text = REFERENCES_SECTION_HEADER + "\n".join( f"{i + 1}. [{ref.document_url}]({ref.document_url})" for i, ref in enumerate(doc_references) ) @@ -1101,9 +1197,15 @@ async def bot_response_generator() -> AsyncIterator[str]: yield self.format_sse(request.chatId, "END") - # Extract usage information after streaming completes - usage_info = get_lm_usage_since(history_length_before) + # Extract usage after streaming completes. Output-rail validation runs + # interleaved with generation in the same history window, so split the + # two - folding them together hid a 100x guardrail cost regression. + usage_info, guardrails_usage = get_lm_usage_since_split( + history_length_before + ) costs_metric["streaming_generation"] = usage_info + if guardrails_usage.get("num_calls", 0) > 0: + costs_metric["output_guardrails"] = guardrails_usage # Record timings time_metric["streaming_generation"] = time.time() - streaming_step_start @@ -1119,53 +1221,54 @@ async def bot_response_generator() -> AsyncIterator[str]: self.log_costs(costs_metric) log_step_timings(time_metric, request.chatId) - # Update budget - self._update_connection_budget( - request.connection_id, costs_metric, request.environment - ) - # Langfuse tracking if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - total_costs = calculate_total_costs(costs_metric) - langfuse.update_current_generation( - model=components["llm_manager"] - .get_provider_info() - .get("model", "unknown"), - usage_details={ - "input": usage_info.get("total_prompt_tokens", 0), - "output": usage_info.get("total_completion_tokens", 0), - "total": usage_info.get("total_tokens", 0), - }, - cost_details={"total": total_costs.get("total_cost", 0.0)}, - metadata={ - "streaming": True, - "streaming_duration_seconds": streaming_duration, - "chunks_streamed": chunk_count, - "cost_breakdown": costs_metric, - "chat_id": request.chatId, - "environment": request.environment, - "stream_id": stream_ctx.stream_id, - }, - ) - langfuse.flush() + metadata_payload = { + "streaming": True, + "streaming_duration_seconds": streaming_duration, + "chunks_streamed": chunk_count, + "cost_breakdown": costs_metric, + "chat_id": request.chatId, + "environment": request.environment, + "stream_id": stream_ctx.stream_id, + } - # Store inference data (for production and testing environments) - if request.environment in [ - PRODUCTION_DEPLOYMENT_ENVIRONMENT, - TEST_DEPLOYMENT_ENVIRONMENT, - ]: try: - await self._store_production_inference_data_async( - request=request, - refined_output=refined_output, - relevant_chunks=relevant_chunks, - accumulated_response="".join(accumulated_response), + langfuse.update_current_generation( + metadata=metadata_payload, ) - except Exception as storage_error: + langfuse.flush() + except Exception as langfuse_error: logger.error( - f"Storage failed for chat_id: {request.chatId}, environment: {request.environment} - {str(storage_error)}" + f"Langfuse streaming metadata update failed: {langfuse_error}", + exc_info=True, + ) + + # Store inference data (for production and testing environments) + # Set RAG data on request for unified storage method + setattr(request, "_rag_refined_questions", refined_output.refined_questions) # noqa: B010 + setattr(request, "_rag_ranked_chunks", relevant_chunks) # noqa: B010 + try: + await self.store_streaming_inference( + request=request, + final_answer="".join(accumulated_response), + ) + except Exception as storage_error: + logger.error( + f"Storage failed for chat_id: {request.chatId}, environment: {request.environment} - {str(storage_error)}" + ) + + # Persist conversation history (RAG streaming) + if self.conversation_history_store is not None: + _rag_bot_message = "".join(accumulated_response) + if _rag_bot_message not in _HISTORY_EXCLUDED_MESSAGES: + await save_history_round( + self.conversation_history_store, + request.chatId, + request.message, + _rag_bot_message, ) # Mark stream as completed successfully @@ -1205,11 +1308,6 @@ async def bot_response_generator() -> AsyncIterator[str]: self.log_costs(costs_metric) log_step_timings(time_metric, request.chatId) - # Update budget even on streaming error - self._update_connection_budget( - request.connection_id, costs_metric, request.environment - ) - def format_sse( self, chat_id: str, @@ -1248,9 +1346,17 @@ def _initialize_service_components( components: Dict[str, Any] = {} # Initialize LLM Manager - components["llm_manager"] = self._initialize_llm_manager( + llm_manager = self._initialize_llm_manager( environment=request.environment, connection_id=request.connection_id ) + components["llm_manager"] = llm_manager + + # Store resolved connection_id on request for downstream use (budget, inference storage) + if llm_manager.connection_id and not request.connection_id: + request.connection_id = llm_manager.connection_id + logger.debug( + f"Stored resolved vault_uuid on request: {llm_manager.connection_id}" + ) if request.environment in self.shared_guardrails_adapters: logger.info( @@ -1269,12 +1375,29 @@ def _initialize_service_components( request.environment, request.connection_id ) - # Initialize Contextual Retriever (replaces hybrid retriever) - components["contextual_retriever"] = self._safe_initialize_contextual_retriever( - request.environment, request.connection_id - ) - - # Initialize Response Generator + # Initialize Contextual Retriever (cached by environment + connection_id) + retriever_key = (request.environment, request.connection_id) + if retriever_key in self._retriever_cache: + components["contextual_retriever"] = self._retriever_cache[retriever_key] + logger.info(f"Using cached ContextualRetriever for key={retriever_key}") + else: + with self._component_cache_lock: + if retriever_key in self._retriever_cache: + components["contextual_retriever"] = self._retriever_cache[ + retriever_key + ] + else: + retriever = self._safe_initialize_contextual_retriever( + request.environment, request.connection_id + ) + if retriever is not None: + self._retriever_cache[retriever_key] = retriever + components["contextual_retriever"] = retriever + + # Initialize Response Generator (fresh per request - NOT cached) + # ResponseGeneratorAgent uses dspy.streamify() which has internal state + # that doesn't reset between calls, causing 0-token streaming on reuse. + # Only costs ~0.02s to create, so caching is not worth the risk. components["response_generator"] = self._safe_initialize_response_generator( components["llm_manager"] ) @@ -1397,10 +1520,16 @@ async def _execute_orchestration_pipeline( # Step 1: Refine user prompt start_time = time.time() + conversation_history, conversation_summary = await get_conversation_history( + chat_id=request.chatId, + store=self.conversation_history_store, + fallback=request.conversationHistory, + ) refined_output, refiner_usage = self._refine_user_prompt( llm_manager=components["llm_manager"], original_message=request.message, - conversation_history=request.conversationHistory, + conversation_history=conversation_history, + conversation_summary=conversation_summary, ) timing_key = f"{prefix}.prompt_refiner" if prefix else "prompt_refiner" time_metric[timing_key] = time.time() - start_time @@ -1443,6 +1572,31 @@ async def _execute_orchestration_pipeline( ) time_metric[timing_key] = time.time() - start_time + # Populate retrieval_context for eval mode (DeepEval metrics need this) + eval_mode = os.getenv("EVAL_MODE", "false").lower() == "true" + if eval_mode and isinstance(generated_response, OrchestrationResponse): + eval_chunks: List[Dict[str, Any]] = [] + for chunk in relevant_chunks: + meta = chunk.get("meta", {}) + meta_dict = meta if isinstance(meta, dict) else {} + eval_chunks.append( + { + "content": chunk.get("content", chunk.get("text", "")), + "metadata": { + "fused_score": meta_dict.get( + "fused_score", chunk.get("fused_score", 0) + ), + "bm25_score": meta_dict.get( + "bm25_score", chunk.get("bm25_score", 0) + ), + "semantic_score": meta_dict.get( + "semantic_score", chunk.get("semantic_score", 0) + ), + }, + } + ) + generated_response.retrieval_context = eval_chunks + # Step 4: Output Guardrails Check # Apply guardrails to all response types for consistent safety across all environments start_time = time.time() @@ -1464,12 +1618,29 @@ async def _execute_orchestration_pipeline( TEST_DEPLOYMENT_ENVIRONMENT, ] and isinstance(output_guardrails_response, OrchestrationResponse): try: - self._store_production_inference_data( - request=request, - refined_output=refined_output, - relevant_chunks=relevant_chunks, - final_response=output_guardrails_response, - ) + # Set RAG data on request for unified storage method + request._rag_refined_questions = refined_output.refined_questions # type: ignore[attr-defined] + request._rag_ranked_chunks = relevant_chunks # type: ignore[attr-defined] + + # Run async storage in a new event loop (sync context) + import threading + + def _store_async() -> None: + try: + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + loop.run_until_complete( + self.store_streaming_inference( + request=request, + final_answer=output_guardrails_response.content, + ) + ) + loop.close() + except Exception as e: + logger.error(f"Error in async storage thread: {str(e)}") + + storage_thread = threading.Thread(target=_store_async, daemon=True) + storage_thread.start() except Exception as storage_error: # Log storage error but don't fail the request logger.error( @@ -1742,129 +1913,63 @@ def _create_out_of_scope_response( content=localized_message, ) - def _store_production_inference_data( - self, - request: OrchestrationRequest, - refined_output: PromptRefinerOutput, - relevant_chunks: List[Dict[str, Union[str, float, Dict[str, Any]]]], - final_response: OrchestrationResponse, - ) -> None: + def _extract_content_from_sse(self, sse_chunk: str) -> Optional[str]: """ - Store production inference data to Resql endpoint for analytics. - - This method stores comprehensive inference data including: - - User question and refined questions - - Conversation history - - Retrieved chunks with rankings - - Embedding scores - - Final generated answer + Extract content from an SSE-formatted chunk. Args: - request: Original orchestration request - refined_output: Prompt refiner output with original and refined questions - relevant_chunks: Retrieved and ranked chunks - final_response: Final orchestration response with generated answer + sse_chunk: SSE-formatted string like 'data: {"chatId": ..., "payload": {"content": "..."}}\n\n' + + Returns: + The content string, or None if parsing fails """ try: - # Only store if the service was active and response was generated successfully - if not final_response.llmServiceActive: - logger.debug( - f"Skipping production data storage for chat_id: {request.chatId} " - f"- LLM service was not active" - ) - return - - # Extract embedding scores from chunks - embedding_scores = [] - for chunk in relevant_chunks: - score_value = chunk.get("fused_score", chunk.get("score", 0.0)) - try: - if isinstance(score_value, (int, float)): - embedding_scores.append(float(score_value)) - else: - embedding_scores.append(0.0) - except (ValueError, TypeError): - embedding_scores.append(0.0) - - # Convert conversation history to list of dicts - conversation_history_list = [ - {"role": item.authorRole, "content": item.message} - for item in (request.conversationHistory or []) - ] - - # Get the production store instance - production_store = get_production_store() - - # Store the inference result asynchronously without blocking - - def store_async() -> None: - """Run async storage in a new event loop in a separate thread.""" - try: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - result = loop.run_until_complete( - production_store.store_inference_result_async( - chat_id=request.chatId, - user_question=request.message, - refined_questions=refined_output.refined_questions, - conversation_history=conversation_history_list, - ranked_chunks=relevant_chunks, - embedding_scores=embedding_scores, - final_answer=final_response.content, - environment=request.environment, - ) - ) - loop.close() - - if result["success"]: - logger.info( - f"Successfully stored inference data for chat_id: {request.chatId}, environment: {request.environment}" - ) - else: - logger.warning( - f"Failed to store inference data for chat_id: {request.chatId}, environment: {request.environment} - " - f"Error: {result['error']}" - ) - except Exception as e: - logger.error(f"Error in async storage thread: {str(e)}") - - # Start storage in background thread (non-blocking) - storage_thread = threading.Thread(target=store_async, daemon=True) - storage_thread.start() - - except Exception as e: - # Log the error but don't fail the request - logger.error( - f"Error storing inference data for chat_id: {request.chatId}, environment: {request.environment} - {str(e)}" - ) + # SSE format: 'data: {...}\n\n' + if not sse_chunk.startswith("data: "): + return None + json_str = sse_chunk[6:].strip() # Remove 'data: ' prefix + if not json_str: + return None + parsed = json_module.loads(json_str) + return parsed.get("payload", {}).get("content") + except (json_module.JSONDecodeError, AttributeError, TypeError): + return None - async def _store_production_inference_data_async( + async def store_streaming_inference( self, request: OrchestrationRequest, - refined_output: PromptRefinerOutput, - relevant_chunks: List[Dict[str, Union[str, float, Dict[str, Any]]]], - accumulated_response: str, + final_answer: str, ) -> None: """ - Async version: Store production inference data to Resql endpoint for analytics. + Store streaming inference data. - This method stores comprehensive inference data including: - - User question and refined questions - - Conversation history - - Retrieved chunks with rankings - - Embedding scores - - Final generated answer (from streaming) + Checks for RAG data on request attributes (_rag_refined_questions, _rag_chunks). + If not present, uses empty arrays. This enables unified storage for all + streaming workflows (RAG, Service, Context, API Tool, etc.). Args: - request: Original orchestration request - refined_output: Prompt refiner output with original and refined questions - relevant_chunks: Retrieved and ranked chunks - accumulated_response: Complete streamed response + request: Orchestration request (may have RAG data as attributes) + final_answer: Complete streamed response """ + # Only store for production and testing environments + if request.environment not in [ + PRODUCTION_DEPLOYMENT_ENVIRONMENT, + TEST_DEPLOYMENT_ENVIRONMENT, + ]: + return + try: + # Get RAG data from request attributes if available (set by _stream_rag_pipeline) + refined_questions: List[str] = getattr( + request, "_rag_refined_questions", [] + ) + ranked_chunks: List[Dict[str, Any]] = getattr( + request, "_rag_ranked_chunks", [] + ) + # Extract embedding scores from chunks - embedding_scores = [] - for chunk in relevant_chunks: + embedding_scores: List[float] = [] + for chunk in ranked_chunks: score_value = chunk.get("fused_score", chunk.get("score", 0.0)) try: if isinstance(score_value, (int, float)): @@ -1887,31 +1992,33 @@ async def _store_production_inference_data_async( result = await production_store.store_inference_result_async( chat_id=request.chatId, user_question=request.message, - refined_questions=refined_output.refined_questions, + refined_questions=refined_questions, conversation_history=conversation_history_list, - ranked_chunks=relevant_chunks, + ranked_chunks=ranked_chunks, embedding_scores=embedding_scores, - final_answer=accumulated_response, + final_answer=final_answer, environment=request.environment, + vault_uuid=request.connection_id, ) if result["success"]: logger.info( - f"Successfully stored inference data (async) for chat_id: {request.chatId}, environment: {request.environment}" + f"Successfully stored streaming inference for chat_id: {request.chatId}, " + f"environment: {request.environment}, " + f"has_rag_data: {bool(refined_questions)}" ) else: logger.warning( - f"Failed to store inference data (async) for chat_id: {request.chatId}, environment: {request.environment} - " + f"Failed to store streaming inference for chat_id: {request.chatId} - " f"Error: {result['error']}" ) except Exception as e: # Log the error but don't fail the request logger.error( - f"Error storing inference data (async) for chat_id: {request.chatId}, environment: {request.environment} - {str(e)}" + f"Error storing streaming inference for chat_id: {request.chatId} - {str(e)}" ) - @observe(name="initialize_guardrails", as_type="span") def _initialize_guardrails( self, environment: str, connection_id: Optional[str] ) -> NeMoRailsAdapter: @@ -1969,7 +2076,7 @@ async def _check_input_guardrails_async( costs_metric["input_guardrails"] = result.usage if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( + langfuse.update_current_span( input=user_message, metadata={ "guardrail_type": "input", @@ -1978,14 +2085,6 @@ async def _check_input_guardrails_async( "blocked_reason": result.reason if not result.allowed else None, "error": result.error if result.error else None, }, - usage_details={ - "input": result.usage.get("total_prompt_tokens", 0), - "output": result.usage.get("total_completion_tokens", 0), - "total": result.usage.get("total_tokens", 0), - }, # type: ignore - cost_details={ - "total": result.usage.get("total_cost", 0.0), - }, ) logger.info( f"Input guardrails check completed: allowed={result.allowed}, " @@ -1998,7 +2097,7 @@ async def _check_input_guardrails_async( logger.error(f"Input guardrails check failed: {str(e)}") if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( + langfuse.update_current_span( metadata={ "error": str(e), "error_type": type(e).__name__, @@ -2041,7 +2140,7 @@ def _check_input_guardrails( costs_metric["input_guardrails"] = result.usage if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( + langfuse.update_current_span( input=user_message, metadata={ "guardrail_type": "input", @@ -2050,14 +2149,6 @@ def _check_input_guardrails( "blocked_reason": result.reason if not result.allowed else None, "error": result.error if result.error else None, }, - usage_details={ - "input": result.usage.get("total_prompt_tokens", 0), - "output": result.usage.get("total_completion_tokens", 0), - "total": result.usage.get("total_tokens", 0), - }, # type: ignore - cost_details={ - "total": result.usage.get("total_cost", 0.0), - }, ) logger.info( f"Input guardrails check completed: allowed={result.allowed}, " @@ -2070,7 +2161,7 @@ def _check_input_guardrails( logger.error(f"Input guardrails check failed: {str(e)}") if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( + langfuse.update_current_span( metadata={ "error": str(e), "error_type": type(e).__name__, @@ -2113,7 +2204,7 @@ async def _check_output_guardrails( costs_metric["output_guardrails"] = result.usage if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( + langfuse.update_current_span( input=assistant_message[:500], # Truncate for readability output=result.verdict, metadata={ @@ -2124,14 +2215,6 @@ async def _check_output_guardrails( "error": result.error if result.error else None, "response_length": len(assistant_message), }, - usage_details={ - "input": result.usage.get("total_prompt_tokens", 0), - "output": result.usage.get("total_completion_tokens", 0), - "total": result.usage.get("total_tokens", 0), - }, # type: ignore - cost_details={ - "total": result.usage.get("total_cost", 0.0), - }, ) logger.info( f"Output guardrails check completed: allowed={result.allowed}, " @@ -2144,7 +2227,7 @@ async def _check_output_guardrails( logger.error(f"Output guardrails check failed: {str(e)}") if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( + langfuse.update_current_span( metadata={ "error": str(e), "error_type": type(e).__name__, @@ -2234,38 +2317,15 @@ def _update_connection_budget( ) -> None: """ Update the budget for an LLM connection based on usage costs. - For production environment, fetches the connection ID asynchronously if not provided. Args: - connection_id: The LLM connection ID (optional) + connection_id: The vault_uuid identifying the LLM connection costs_metric: Dictionary of costs per component environment: The deployment environment (production/testing/development) """ try: budget_tracker = get_budget_tracker() - # For production environment, fetch connection ID if not provided - if environment == "production" and not connection_id: - logger.debug( - "Production environment detected, fetching connection ID..." - ) - try: - # Use synchronous fetch to avoid event loop issues - production_id = ( - budget_tracker.connection_fetcher.fetch_connection_id_sync( - "production" - ) - ) - if production_id: - connection_id = str(production_id) - logger.info(f"Using production connection_id: {connection_id}") - else: - logger.warning("Could not fetch production connection ID") - except Exception as fetch_error: - logger.error( - f"Error fetching production connection ID: {str(fetch_error)}" - ) - result = budget_tracker.update_budget_from_costs( connection_id, costs_metric ) @@ -2278,7 +2338,7 @@ def _update_connection_budget( ) else: logger.debug( - f"Budget updated successfully for connection_id={connection_id}" + f"Budget updated successfully for uuid={connection_id}" ) else: reason = result.get("reason", "unknown") @@ -2299,18 +2359,38 @@ def _initialize_llm_manager( """ Initialize LLM Manager with proper configuration. + For production environment, resolves vault_uuid from DB if not provided. + For testing environment, connection_id (vault_uuid) must be provided. + Args: - environment: Environment context (production/testing/development) - connection_id: Optional connection identifier + environment: Environment context (production/testing) + connection_id: Vault UUID for the connection (required for testing, auto-resolved for production) Returns: - LLMManager: Initialized LLM manager instance + LLMManager: Initialized LLM manager instance (with connection_id property) """ try: logger.info(f"Initializing LLM Manager for environment: {environment}") + resolved_connection_id = connection_id + + # Resolve vault_uuid for production if not provided + if environment == "production" and not connection_id: + from src.utils.connection_id_fetcher import get_connection_id_fetcher + + fetcher = get_connection_id_fetcher() + resolved_connection_id = fetcher.fetch_vault_uuid_sync("production") + if not resolved_connection_id: + raise ValueError( + "No production connection found in database. " + "Please create a production LLM connection first." + ) + logger.info( + f"Resolved production vault_uuid from DB: {resolved_connection_id}" + ) + llm_manager = LLMManager( - environment=environment, connection_id=connection_id + environment=environment, connection_id=resolved_connection_id ) llm_manager.ensure_global_config() @@ -2322,12 +2402,13 @@ def _initialize_llm_manager( logger.error(f"Failed to initialize LLM Manager: {str(e)}") raise - @observe(name="refine_user_prompt", as_type="chain") + @observe(name="refine_user_prompt", as_type="generation") def _refine_user_prompt( self, llm_manager: LLMManager, original_message: str, conversation_history: List[ConversationItem], + conversation_summary: Optional[str] = None, ) -> tuple[PromptRefinerOutput, Dict[str, Any]]: """ Refine user prompt using loaded LLM configuration and return usage info. @@ -2336,6 +2417,10 @@ def _refine_user_prompt( llm_manager: The LLM manager instance to use original_message: The original user message to refine conversation_history: Previous conversation context + conversation_summary: Optional summary of earlier conversation rounds + that were evicted from Redis. When provided it is prepended to the + DSPy history as a ``system`` turn so the refiner can use it for + context without re-summarising via an LLM call. Returns: Tuple of (PromptRefinerOutput, usage_dict): The refined prompt output and usage info @@ -2348,8 +2433,16 @@ def _refine_user_prompt( logger.info("Starting prompt refinement process") try: - # Convert conversation history to DSPy format + # Convert conversation history to DSPy format, optionally prepending + # a pre-computed summary of earlier (evicted) conversation rounds. history: List[Dict[str, str]] = [] + if conversation_summary: + history.append( + { + "role": "system", + "content": f"Summary of earlier conversation: {conversation_summary}", + } + ) for item in conversation_history: role = "assistant" if item.authorRole == "bot" else item.authorRole history.append({"role": role, "content": item.message}) @@ -2616,7 +2709,7 @@ def _extract_document_references( return references if references else None - @observe(name="generate_rag_response", as_type="generation") + @observe(name="generate_rag_response", as_type="span") def _generate_rag_response( self, llm_manager: LLMManager, @@ -2664,7 +2757,7 @@ def _generate_rag_response( llmServiceActive=False, questionOutOfLLMScope=False, inputGuardFailed=False, - content=localized_msg, + content=TECHNICAL_ISSUE_MESSAGE, ) try: @@ -2694,17 +2787,11 @@ def _generate_rag_response( costs_metric["response_generator"] = generator_usage if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( - model=llm_manager.get_provider_info().get("model", "unknown"), - usage_details={ - "input": generator_usage.get("total_prompt_tokens", 0), - "output": generator_usage.get("total_completion_tokens", 0), - "total": generator_usage.get("total_tokens", 0), - }, - cost_details={ - "total": generator_usage.get("total_cost", 0.0), - }, + langfuse.update_current_span( metadata={ + "model": llm_manager.get_provider_info().get( + "model", "unknown" + ), "num_calls": generator_usage.get("num_calls", 0), "question_out_of_scope": question_out_of_scope, "num_chunks_used": len(relevant_chunks) @@ -2753,7 +2840,7 @@ def _generate_rag_response( doc_references = self._extract_document_references(relevant_chunks) content_with_refs = answer if doc_references: - refs_text = "\n\n**References:**\n" + "\n".join( + refs_text = REFERENCES_SECTION_HEADER + "\n".join( f"{i + 1}. {ref.document_url}" for i, ref in enumerate(doc_references) ) @@ -2789,7 +2876,7 @@ def _generate_rag_response( ) if self.langfuse_config.langfuse_client: langfuse = self.langfuse_config.langfuse_client - langfuse.update_current_generation( + langfuse.update_current_span( metadata={ "error_id": error_id, "error_type": type(e).__name__, @@ -2821,7 +2908,7 @@ def _generate_rag_response( llmServiceActive=False, questionOutOfLLMScope=False, inputGuardFailed=False, - content=localized_msg, + content=TECHNICAL_ISSUE_MESSAGE, ) # ======================================================================== @@ -2955,9 +3042,15 @@ def _get_context_manager(self) -> "ContextGenerationManager": from src.llm_orchestrator_config.context_manager import ( ContextGenerationManager, ) + from src.utils.connection_id_fetcher import get_connection_id_fetcher - # Use existing LLM manager or create new one for context generation - llm_manager = LLMManager() + # Resolve production vault_uuid from DB before creating LLM manager + fetcher = get_connection_id_fetcher() + connection_id = fetcher.fetch_vault_uuid_sync("production") + + llm_manager = LLMManager( + environment="production", connection_id=connection_id + ) self._context_manager = ContextGenerationManager(llm_manager) logger.debug("Lazy initialized ContextGenerationManager for vector indexer") diff --git a/src/llm_orchestration_service_api.py b/src/llm_orchestration_service_api.py index 34279c7b..465905bf 100644 --- a/src/llm_orchestration_service_api.py +++ b/src/llm_orchestration_service_api.py @@ -1,5 +1,7 @@ """LLM Orchestration Service API - FastAPI application.""" +import asyncio +import os import logging from contextlib import asynccontextmanager from typing import Any, AsyncGenerator, Dict @@ -8,16 +10,18 @@ from fastapi.responses import StreamingResponse, JSONResponse from fastapi.exceptions import RequestValidationError from pydantic import ValidationError -from loguru import logger import uvicorn from llm_orchestration_service import LLMOrchestrationService +from llm_orchestrator_config.llm_manager import LLMManager from src.utils.redis_client import ( init_redis_client, close_redis_client, check_redis_health, ) from src.utils.api_tool_session_store import APIToolSessionStore +from src.utils.conversation_history_store import ConversationHistoryStore +from src.utils.conversation_summary_generator import create_incremental_summarizer from src.llm_orchestrator_config.llm_ochestrator_constants import ( STREAMING_ALLOWED_ENVS, STREAM_TIMEOUT_MESSAGE, @@ -34,7 +38,13 @@ ) from src.llm_orchestrator_config.stream_config import StreamConfig from src.llm_orchestrator_config.exceptions import StreamTimeoutError -from src.utils.stream_timeout import stream_timeout + +# NOTE: imported via the bare package path, not "src.llm_orchestrator_config". +# Both spellings resolve to separate module objects at runtime, so the class +# imported here must match the one the config loader raises or `except` misses. +from llm_orchestrator_config.exceptions import ConfigurationError +from src.utils.stream_timeout import stream_timeout, with_heartbeat +from src.utils.observation_utils import safe_observation_context from src.utils.error_utils import generate_error_id, log_error_with_context from src.utils.rate_limiter import RateLimiter from src.utils.prompt_config_loader import RefreshStatus @@ -48,7 +58,13 @@ ContextGenerationRequest, ContextGenerationResponse, EmbeddingErrorResponse, + DeepEvalTestOrchestrationResponse, ) +from src.utils.connection_id_fetcher import get_connection_id_fetcher +from src.loki_logger import LokiLogger + +# Initialize Loki logger for centralized logging +logger = LokiLogger(service_name="llm-orchestration-api") @asynccontextmanager @@ -91,23 +107,50 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: try: await init_redis_client() app.state.session_store = APIToolSessionStore() + + # Wire an incremental summarizer if the LLM manager singleton is available. + summarizer = None + try: + summarizer = create_incremental_summarizer(LLMManager()) + logger.info("Incremental conversation summarizer initialized") + except Exception as e: + logger.warning( + f"Could not create incremental summarizer, continuing without it: {e}" + ) + + app.state.conversation_history_store = ConversationHistoryStore( + summarizer=summarizer + ) logger.info("Redis session store initialized successfully") except Exception as e: logger.warning(f"Redis session store unavailable, continuing without it: {e}") app.state.session_store = None + app.state.conversation_history_store = None - # Expose session_store on the orchestration service so workflow executors - # (e.g. APIToolWorkflowExecutor) can reach it via self.orchestration_service. + # Expose session_store and conversation_history_store on the orchestration + # service so downstream components can reach them via self.orchestration_service. if ( hasattr(app.state, "orchestration_service") and app.state.orchestration_service is not None ): app.state.orchestration_service.session_store = app.state.session_store + app.state.orchestration_service.conversation_history_store = ( + app.state.conversation_history_store + ) yield # Shutdown logger.info("Shutting down LLM Orchestration Service API") + + # Await any in-flight incremental summary tasks to avoid lost work. + store = getattr(app.state, "conversation_history_store", None) + if store is not None and store._pending_tasks: + logger.info( + f"Waiting for {len(store._pending_tasks)} pending summary task(s) to complete..." + ) + await asyncio.gather(*store._pending_tasks, return_exceptions=True) + if ( hasattr(app.state, "orchestration_service") and app.state.orchestration_service is not None @@ -249,6 +292,22 @@ async def pydantic_validation_exception_handler( ) +@app.post("/cache/clear") +async def clear_connection_cache() -> dict[str, str]: + """Clear cached connection IDs and vault UUIDs.""" + try: + fetcher = get_connection_id_fetcher() + fetcher.clear_cache() + logger.info("Connection cache cleared via /cache/clear endpoint") + return {"status": "ok", "message": "Connection cache cleared"} + except Exception as e: + logger.error(f"Failed to clear connection cache: {e}") + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="Failed to clear connection cache", + ) from e + + @app.get("/health") async def health_check(request: Request) -> dict[str, str]: """Health check endpoint.""" @@ -495,16 +554,27 @@ async def stream_orchestrated_response( from datetime import datetime def create_sse_error_stream(chat_id: str, error_message: str) -> str: - """Create SSE format error response.""" + """Create an SSE error response, terminated by the END marker. + + The END frame is what the notification server translates into the + browser's ``stream_end``. Without it an error frame is indistinguishable + from ordinary content, so the client keeps waiting on a stream that has + already finished - which is how an upstream timeout turned into a chat + that hung indefinitely. Every caller is a terminal error path, so + closing the stream here is always correct. + """ from typing import Dict, Any - error_payload: Dict[str, Any] = { - "chatId": chat_id, - "payload": {"content": error_message}, - "timestamp": str(int(datetime.now().timestamp() * 1000)), - "sentTo": [], - } - return f"data: {json_module.dumps(error_payload)}\n\n" + def frame(content: str) -> str: + payload: Dict[str, Any] = { + "chatId": chat_id, + "payload": {"content": content}, + "timestamp": str(int(datetime.now().timestamp() * 1000)), + "sentTo": [], + } + return f"data: {json_module.dumps(payload)}\n\n" + + return frame(error_message) + frame("END") try: logger.info( @@ -512,11 +582,12 @@ def create_sse_error_stream(chat_id: str, error_message: str) -> str: f"chatId: {request.chatId}, " f"environment: {request.environment}, " f"message: {request.message[:100]}..." + f"connection_id: {request.connection_id}" ) # Streaming is only for allowed environments if request.environment not in STREAMING_ALLOWED_ENVS: - error_msg = f"Streaming is only available for production environment. Current environment: {request.environment}. Please use /orchestrate endpoint for non-streaming environments." + error_msg = f"Streaming is only available for production and testing environments. Current environment: {request.environment}. Please use /orchestrate endpoint for non-streaming environments." logger.warning(error_msg) async def env_error_stream() -> AsyncGenerator[str, None]: @@ -618,33 +689,51 @@ async def rate_limit_error_stream() -> AsyncGenerator[str, None]: # Wrap streaming response with timeout async def timeout_wrapped_stream() -> AsyncGenerator[str, None]: """Generator wrapper with timeout enforcement.""" - try: - async with stream_timeout(StreamConfig.MAX_STREAM_DURATION_SECONDS): - async for ( - chunk - ) in orchestration_service.stream_orchestration_response(request): - yield chunk - except StreamTimeoutError as timeout_exc: - # StreamTimeoutError already has error_id - log_error_with_context( - logger, - timeout_exc.error_id, - "streaming_timeout", - request.chatId, - timeout_exc, - ) - # Send timeout message to client - yield create_sse_error_stream(request.chatId, STREAM_TIMEOUT_MESSAGE) - except Exception as stream_error: - error_id = generate_error_id() - log_error_with_context( - logger, error_id, "streaming_error", request.chatId, stream_error - ) - # Send generic error message to client - yield create_sse_error_stream( - request.chatId, - "I apologize, but I encountered an issue while generating your response. Please try again.", - ) + with safe_observation_context( + as_type="generation", + name="streaming_generation", + input={"message": request.message[:500], "chat_id": request.chatId}, + ): + try: + async with stream_timeout(StreamConfig.MAX_STREAM_DURATION_SECONDS): + # Heartbeat frames keep proxies from closing a slow stream, + # and the idle budget fails fast on one that has truly + # stalled rather than waiting out the total-duration cap. + async for chunk in with_heartbeat( + orchestration_service.stream_orchestration_response( + request + ), + heartbeat_interval=StreamConfig.HEARTBEAT_INTERVAL_SECONDS, + idle_timeout=StreamConfig.IDLE_TIMEOUT_SECONDS, + ): + yield chunk + except StreamTimeoutError as timeout_exc: + # StreamTimeoutError already has error_id + log_error_with_context( + logger, + timeout_exc.error_id, + "streaming_timeout", + request.chatId, + timeout_exc, + ) + # Send timeout message to client + yield create_sse_error_stream( + request.chatId, STREAM_TIMEOUT_MESSAGE + ) + except Exception as stream_error: + error_id = generate_error_id() + log_error_with_context( + logger, + error_id, + "streaming_error", + request.chatId, + stream_error, + ) + # Send generic error message to client + yield create_sse_error_stream( + request.chatId, + "I apologize, but I encountered an issue while generating your response. Please try again.", + ) # Stream the response return StreamingResponse( @@ -745,6 +834,13 @@ async def generate_context_with_caching( return ContextGenerationResponse(**result) + except ConfigurationError as e: + # No usable LLM connection for this environment. This is an operator + # action, not a transient fault - 503 with the reason so callers (e.g. + # the vector indexer) can stop retrying and surface something useful. + error_id = generate_error_id() + log_error_with_context(logger, error_id, "context_generation_endpoint", None, e) + raise HTTPException(status_code=503, detail=str(e)) from e except Exception as e: error_id = generate_error_id() log_error_with_context(logger, error_id, "context_generation_endpoint", None, e) @@ -787,6 +883,84 @@ async def get_available_embedding_models( ) from e +@app.post("/orchestrate-eval") +async def orchestrate_llm_request_eval( + http_request: Request, + request: OrchestrationRequest, +) -> DeepEvalTestOrchestrationResponse: + """ + Process LLM orchestration request with additional testing data. + + This endpoint is only available when EVAL_MODE=true and returns + retrieval context and refined questions for DeepEval metrics evaluation. + + Args: + http_request: FastAPI Request object for accessing app state + request: OrchestrationRequest containing user message and context + + Returns: + OrchestrationResponse: Response with LLM output, status flags, and test data + + Raises: + HTTPException: For processing errors or if not in testing mode + """ + # Check if eval mode is enabled + eval_mode = os.getenv("EVAL_MODE", "false").lower() == "true" + if not eval_mode: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Eval endpoint not available in production mode", + ) + + try: + logger.info(f"Received EVAL orchestration request for chatId: {request.chatId}") + + if not hasattr(http_request.app.state, "orchestration_service"): + logger.error("Orchestration service not found in app state") + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="Service not initialized", + ) + + orchestration_service = http_request.app.state.orchestration_service + if orchestration_service is None: + logger.error("Orchestration service is None") + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="Service not initialized", + ) + + # Process the request (will include test data due to EVAL_MODE env var) + response = await orchestration_service.process_orchestration_request(request) + + # Convert to test response with additional fields + # Response may be OrchestrationResponse or TestOrchestrationResponse + chat_id = getattr(response, "chatId", request.chatId) + retrieval_ctx = getattr(response, "retrieval_context", None) + + test_response = DeepEvalTestOrchestrationResponse( + chatId=chat_id, + llmServiceActive=response.llmServiceActive, + questionOutOfLLMScope=response.questionOutOfLLMScope, + inputGuardFailed=response.inputGuardFailed, + content=response.content, + retrieval_context=retrieval_ctx, + expected_output=None, # Will be populated by test framework + ) + + logger.info(f"Successfully processed TEST request for chatId: {request.chatId}") + return test_response + + except HTTPException: + raise + except Exception as e: + logger.error(f"Unexpected error processing TEST request: {str(e)}") + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="Internal server error occurred", + ) from e + + @app.post("/prompt-config/refresh") def refresh_prompt_config(http_request: Request) -> Dict[str, Any]: """ diff --git a/src/llm_orchestrator_config/config/loader.py b/src/llm_orchestrator_config/config/loader.py index 2d9a3f60..7506fdae 100644 --- a/src/llm_orchestrator_config/config/loader.py +++ b/src/llm_orchestrator_config/config/loader.py @@ -7,7 +7,7 @@ import yaml from dotenv import load_dotenv -from loguru import logger +from src.loki_logger import LokiLogger from llm_orchestrator_config.config.schema import ( LLMConfiguration, @@ -27,6 +27,9 @@ # Constants DEFAULT_CONFIG_FILENAME = "llm_config.yaml" +# Initialize Loki logger +logger = LokiLogger(service_name="config-loader") + # Type alias for configuration values that can be processed ConfigValue = Union[str, Dict[str, Any], List[Any], int, float, bool, None] @@ -132,6 +135,11 @@ def load_config(self) -> LLMConfiguration: except yaml.YAMLError as e: raise ConfigurationError(f"Failed to parse YAML configuration: {e}") from e except Exception as e: + # Already a ConfigurationError with a specific, operator-facing + # message (e.g. "No production LLM connection configured") - keep it + # verbatim instead of nesting another prefix in front of it. + if isinstance(e, ConfigurationError): + raise raise ConfigurationError(f"Failed to load configuration: {e}") from e def _resolve_vault_secrets(self, config: Dict[str, Any]) -> Dict[str, Any]: @@ -196,12 +204,17 @@ def _resolve_provider_secrets( if "providers" not in config: return - # Validate environment-specific requirements - if self.environment in ["development", "test"]: - if not self.connection_id: - raise ConfigurationError( - f"connection_id is required for {self.environment} environment" - ) + # connection_id (vault_uuid) is required for all environments. A missing + # one means no connection row exists for this environment at all - the + # operator has not configured one yet, so say that rather than naming an + # internal field. + if not self.connection_id: + logger.error( + f"No {self.environment} LLM connection configured. " + f"Create a {self.environment} connection so its credentials are " + f"stored in Vault." + ) + raise ConfigurationError(f"No {self.environment} LLM connection configured") try: providers_to_update: Dict[str, Dict[str, Any]] = {} @@ -223,60 +236,31 @@ def _resolve_provider_secrets( ) continue - # For production: try to find any available model - # For dev/test: use connection_id to find specific model - if self.environment == "production": - # Find first available model for this provider - available_models = resolver.list_available_models( - provider_name, self.environment + # Use connection_id (vault_uuid) directly for all environments + secret = resolver.get_secret_for_model( + provider_name, + self.environment, + "", + self.connection_id, + ) + if secret: + model_name = secret.model + updated_config = self._merge_config_with_secrets( + provider_config, secret, model_name + ) + providers_to_update[provider_name] = updated_config + logger.info( + f"Configured {provider_name} with model {model_name} " + f"(vault_uuid: {self.connection_id})" ) - if available_models: - # Use the first available model in production - model_name = available_models[0] - secret = resolver.get_secret_for_model( - provider_name, self.environment, model_name - ) - if secret: - # Update provider config with secrets - updated_config = self._merge_config_with_secrets( - provider_config, secret, model_name - ) - providers_to_update[provider_name] = updated_config - logger.info( - f"Configured {provider_name} with model {model_name}" - ) - else: - logger.warning( - f"No secret found for {provider_name} model {model_name}" - ) - else: - logger.warning( - f"No available models found for provider {provider_name}" - ) else: - # For dev/test, try to find the specific connection_id - # Try each model to see if we can find the connection - for model_name in provider_config.get("models", {}): - secret = resolver.get_secret_for_model( - provider_name, - self.environment, - model_name, - self.connection_id, - ) - if secret: - # Update provider config with secrets - updated_config = self._merge_config_with_secrets( - provider_config, secret, model_name - ) - providers_to_update[provider_name] = updated_config - logger.info( - f"Configured {provider_name} with connection {self.connection_id}" - ) - break - else: - logger.warning( - f"No connection found for {provider_name} with connection_id {self.connection_id}" - ) + # Either nothing is stored for this provider under this + # connection, or what is stored failed validation - the + # resolver logs which one. + logger.warning( + f"No usable {provider_name} credentials for vault_uuid " + f"{self.connection_id} - skipping this provider" + ) except Exception as e: logger.error(f"Failed to process provider {provider_name}: {e}") @@ -287,22 +271,21 @@ def _resolve_provider_secrets( # Continue to next provider instead of failing completely continue - # Check if we have any providers configured + # Check if we have any providers configured. Reaching here means a + # connection row exists but none of its provider secrets could be + # loaded from Vault - either the cron has not written them yet or + # what it wrote does not match the expected schema (see the + # per-provider warnings above for which, and why). if not providers_to_update: - if self.environment == "production": - raise ConfigurationError( - "No providers available for production environment. " - "At least one provider must have production models configured." - ) - else: - raise ConfigurationError( - f"No providers available for {self.environment} environment" - + ( - f" with connection_id {self.connection_id}" - if self.connection_id - else "" - ) - ) + logger.error( + f"No usable LLM provider for the {self.environment} connection " + f"(vault_uuid: {self.connection_id}). The connection exists but " + f"none of its credentials could be loaded from Vault - see the " + f"per-provider messages above." + ) + raise ConfigurationError( + f"No usable LLM provider for the {self.environment} connection" + ) # Update the configuration with only available providers config["providers"] = providers_to_update @@ -645,9 +628,12 @@ def resolve_embedding_model( ) -> tuple[str, str]: """Resolve embedding model from vault based on environment and connection_id. + Uses the same vault_uuid as the LLM model resolver. The embedding secret + is stored at `secret/embeddings/connections/{platform}/{vault_uuid}`. + Args: - environment: Environment (production, development, test) - connection_id: Optional connection ID for dev/test environments + environment: Environment (production, testing) + connection_id: Vault UUID for the connection (same as LLM connection) Returns: Tuple of (provider_name, model_name) resolved from vault @@ -655,6 +641,17 @@ def resolve_embedding_model( Raises: ConfigurationError: If no embedding models are available """ + # Resolve vault_uuid for production if not provided + if not connection_id: + from src.utils.connection_id_fetcher import get_connection_id_fetcher + + fetcher = get_connection_id_fetcher() + connection_id = fetcher.fetch_vault_uuid_sync(environment) + if not connection_id: + raise ConfigurationError( + f"No {environment} connection found in database for embedding resolution" + ) + # Load raw config to get vault settings try: with open(self.config_path, "r", encoding="utf-8") as file: @@ -669,58 +666,33 @@ def resolve_embedding_model( resolver: SecretResolver = self._initialize_vault_resolver(config) # Get available providers from config - providers: List[str] = ["azure_openai", "aws_bedrock"] # Hardcoded for now - - if environment == "production": - # Find first available embedding model across all providers - for provider in providers: - try: - models: List[str] = resolver.list_available_embedding_models( - provider, environment - ) - embedding_models: List[str] = [ - m for m in models if self._is_embedding_model(m) - ] - if embedding_models: - logger.info( - f"Resolved production embedding model: {provider}/{embedding_models[0]}" - ) - return provider, embedding_models[0] - except Exception as e: - logger.debug( - f"Provider {provider} not available for embeddings: {e}" - ) - continue - - raise ConfigurationError("No embedding models available in production") - else: - # Use connection_id to find specific embedding model - if not connection_id: - raise ConfigurationError( - f"connection_id is required for {environment} environment" - ) + providers: List[str] = ["azure_openai", "aws_bedrock"] - for provider in providers: - try: - secret: Optional[Union[AzureOpenAISecret, AWSBedrockSecret]] = ( - resolver.get_embedding_secret_for_model( - provider, environment, "", connection_id - ) + # Use vault_uuid directly to find embedding secret (same UUID as LLM) + for provider in providers: + try: + secret: Optional[Union[AzureOpenAISecret, AWSBedrockSecret]] = ( + resolver.get_embedding_secret_for_model( + provider, environment, "", connection_id ) - if secret and self._is_embedding_model(secret.model): - logger.info( - f"Resolved {environment} embedding model: {provider}/{secret.model}" - ) - return provider, secret.model - except Exception as e: - logger.debug( - f"Provider {provider} not available with connection {connection_id}: {e}" + ) + if secret and self._is_embedding_model(secret.model): + logger.info( + f"Resolved embedding model: {provider}/{secret.model} " + f"(vault_uuid: {connection_id})" ) - continue + return provider, secret.model + except Exception as e: + logger.debug( + f"Provider {provider} not available for embeddings " + f"with vault_uuid {connection_id}: {e}" + ) + continue - raise ConfigurationError( - f"No embedding models available for {environment} with connection_id {connection_id}" - ) + raise ConfigurationError( + f"No embedding models available for {environment} " + f"with vault_uuid {connection_id}" + ) except yaml.YAMLError as e: raise ConfigurationError(f"Failed to parse YAML configuration: {e}") from e @@ -751,6 +723,17 @@ def get_embedding_provider_config( ConfigurationError: If configuration cannot be loaded or secrets not found """ try: + # Auto-resolve vault_uuid from DB if not provided + if not connection_id: + from src.utils.connection_id_fetcher import get_connection_id_fetcher + + fetcher = get_connection_id_fetcher() + connection_id = fetcher.fetch_vault_uuid_sync(environment) + if not connection_id: + raise ConfigurationError( + f"No {environment} connection found in database for embedding config" + ) + # Load raw config with open(self.config_path, "r", encoding="utf-8") as file: raw_config: Dict[str, Any] = yaml.safe_load(file) diff --git a/src/llm_orchestrator_config/context_manager.py b/src/llm_orchestrator_config/context_manager.py index 3bb6e5ac..09efb405 100644 --- a/src/llm_orchestrator_config/context_manager.py +++ b/src/llm_orchestrator_config/context_manager.py @@ -2,12 +2,14 @@ from typing import Any, Dict, Optional -from loguru import logger - +from src.loki_logger import LokiLogger from src.llm_orchestrator_config.llm_manager import LLMManager from src.models.request_models import ContextGenerationRequest from langfuse import observe +# Initialize Loki logger +logger = LokiLogger(service_name="context-manager") + class ContextGenerationManager: """Manager for context generation with Anthropic methodology.""" @@ -51,7 +53,7 @@ def generate_context_with_caching( ) # For now, call LLM directly (caching structure ready for future) - # TODO: Implement actual prompt caching when ready + # Implement actual prompt caching when ready response = self._call_llm_for_context( prompt=full_prompt, model=model_info["model"], diff --git a/src/llm_orchestrator_config/embedding_manager.py b/src/llm_orchestrator_config/embedding_manager.py index 6c9bf277..3fcca9c9 100644 --- a/src/llm_orchestrator_config/embedding_manager.py +++ b/src/llm_orchestrator_config/embedding_manager.py @@ -6,13 +6,15 @@ import dspy import numpy as np -from loguru import logger +from src.loki_logger import LokiLogger from pydantic import BaseModel - from .vault.vault_client import VaultAgentClient from .config.loader import ConfigurationLoader from .exceptions import ConfigurationError +# Initialize Loki logger +logger = LokiLogger(service_name="embedding-manager") + class EmbeddingFailure(BaseModel): """Model for tracking embedding failures.""" diff --git a/src/llm_orchestrator_config/feature_flags.py b/src/llm_orchestrator_config/feature_flags.py index e5a88f55..f0346a73 100644 --- a/src/llm_orchestrator_config/feature_flags.py +++ b/src/llm_orchestrator_config/feature_flags.py @@ -1,7 +1,10 @@ """Feature flags for tool classifier system.""" import os -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="feature-flags") class FeatureFlags: @@ -22,6 +25,7 @@ class FeatureFlags: - SERVICE_WORKFLOW_ENABLED: Enable Layer 1 service workflow (default: true) - API_TOOL_CALLING_WORKFLOW_ENABLED: Enable Layer 2 API tool calling workflow (default: true) - CONTEXT_WORKFLOW_ENABLED: Enable Layer 3 context workflow (default: true) + - MULTI_INTENT_ENABLED: Enable parallel multi-intent path in ATC (default: true) """ # Master switch for tool classifier @@ -43,6 +47,20 @@ class FeatureFlags: os.getenv("CONTEXT_WORKFLOW_ENABLED", "true").lower() == "true" ) + # Multi-intent (parallel multi-API) path + # When False: ambiguous-band ATC results go straight to the existing single-endpoint path + # When True: IntentDecomposer runs on ambiguous-band results; parallel endpoint searches + # are attempted when the query is detected as multi-intent + MULTI_INTENT_ENABLED = os.getenv("MULTI_INTENT_ENABLED", "true").lower() == "true" + + # ATC Response Cache — two-tier Redis cache for API Tool Calling responses + # When True: successful ATC API responses are cached and follow-up queries are + # served from cache where possible (L1 exact hit, L2 follow-up routing) + # When False: all cache reads and writes are skipped; normal ATC flow unchanged + ATC_RESPONSE_CACHE_ENABLED: bool = ( + os.getenv("ATC_RESPONSE_CACHE_ENABLED", "true").lower() == "true" + ) + # RAG and OOD workflows are always enabled (no flags) # RAG is the core fallback, OOD is the final safety net @@ -61,6 +79,10 @@ def log_configuration(cls) -> None: f" API_TOOL_CALLING_WORKFLOW_ENABLED: {cls.API_TOOL_CALLING_WORKFLOW_ENABLED}" ) logger.info(f" CONTEXT_WORKFLOW_ENABLED: {cls.CONTEXT_WORKFLOW_ENABLED}") + logger.info(f" MULTI_INTENT_ENABLED: {cls.MULTI_INTENT_ENABLED}") + logger.info( + f" ATC_RESPONSE_CACHE_ENABLED: {cls.ATC_RESPONSE_CACHE_ENABLED}" + ) logger.info(f" FALLBACK_TO_RAG_ON_ERROR: {cls.FALLBACK_TO_RAG_ON_ERROR}") else: logger.info(" (Classifier disabled - using RAG-only pipeline)") diff --git a/src/llm_orchestrator_config/llm_manager.py b/src/llm_orchestrator_config/llm_manager.py index 70df586e..4ebd21a1 100644 --- a/src/llm_orchestrator_config/llm_manager.py +++ b/src/llm_orchestrator_config/llm_manager.py @@ -77,6 +77,11 @@ def __init__( self._initialize_providers() LLMManager._initialized = True + @property + def connection_id(self) -> Optional[str]: + """Return the connection ID (vault_uuid) used by this manager.""" + return self._connection_id + def _load_configuration(self) -> None: """Load configuration from file. @@ -85,6 +90,10 @@ def _load_configuration(self) -> None: """ try: self._config = self._config_loader.load_config() + except ConfigurationError: + # The loader's message already identifies the environment and what + # the operator needs to do; re-wrapping only buries it. + raise except Exception as e: raise ConfigurationError(f"Failed to load LLM configuration: {e}") from e @@ -171,7 +180,7 @@ def configure_dspy(self, provider: Optional[LLMProvider] = None) -> None: provider: Optional specific provider to configure DSPY with. """ dspy_client = self.get_dspy_client(provider) - dspy.configure(lm=dspy_client) + dspy.configure(lm=dspy_client, track_usage=True) def ensure_global_config(self, provider: Optional[LLMProvider] = None) -> None: """Configure DSPy exactly once per process (thread-safe).""" @@ -181,7 +190,7 @@ def ensure_global_config(self, provider: Optional[LLMProvider] = None) -> None: # Re-check inside the lock to prevent race condition if not self._configured: dspy_client = self.get_dspy_client(provider) - dspy.configure(lm=dspy_client) + dspy.configure(lm=dspy_client, track_usage=True) self._configured = True @contextmanager diff --git a/src/llm_orchestrator_config/llm_ochestrator_constants.py b/src/llm_orchestrator_config/llm_ochestrator_constants.py index 789ef62a..82bfceb1 100644 --- a/src/llm_orchestrator_config/llm_ochestrator_constants.py +++ b/src/llm_orchestrator_config/llm_ochestrator_constants.py @@ -26,7 +26,9 @@ # Query validation messages - single generic message for all rejection types # (empty queries, special characters only, too short, repetitive characters) QUERY_VALIDATION_FAILED_MESSAGES = { - "et": "Palun esitage kehtiv küsimus või sõnum, et ma saaksin teid aidata." + "et": "Palun esitage kehtiv küsimus või sõnum, et ma saaksin teid aidata.", + "en": "Please provide a valid question or reply so I can assist you.", + "ru": "Пожалуйста, введите корректный вопрос или сообщение, чтобы я мог вам помочь.", } # Legacy constants for backward compatibility (English defaults) @@ -45,7 +47,7 @@ ] # Streaming configuration -STREAMING_ALLOWED_ENVS = {"production"} +STREAMING_ALLOWED_ENVS = {"production", "testing"} TEST_DEPLOYMENT_ENVIRONMENT = "testing" PRODUCTION_DEPLOYMENT_ENVIRONMENT = "production" diff --git a/src/llm_orchestrator_config/stream_config.py b/src/llm_orchestrator_config/stream_config.py index 84e5edd5..39b433d3 100644 --- a/src/llm_orchestrator_config/stream_config.py +++ b/src/llm_orchestrator_config/stream_config.py @@ -5,8 +5,13 @@ class StreamConfig: """Hardcoded configuration for streaming limits and timeouts.""" # Timeout Configuration - MAX_STREAM_DURATION_SECONDS: int = 300 # 5 minutes - IDLE_TIMEOUT_SECONDS: int = 60 # 1 minute idle timeout + MAX_STREAM_DURATION_SECONDS: int = 300 # 5 minutes, total wall clock + # Measured between chunks, not cumulatively: a long answer that keeps + # producing tokens is never cut off, however long it takes in total. + IDLE_TIMEOUT_SECONDS: int = 60 # 1 minute with no output at all + # How often to emit an SSE comment frame while the stream is quiet, so + # intermediate proxies see traffic and hold the connection open. + HEARTBEAT_INTERVAL_SECONDS: int = 15 # Size Limits MAX_MESSAGE_LENGTH: int = 10000 # Maximum characters in message diff --git a/src/llm_orchestrator_config/vault/models.py b/src/llm_orchestrator_config/vault/models.py index 60900522..bb8265e0 100644 --- a/src/llm_orchestrator_config/vault/models.py +++ b/src/llm_orchestrator_config/vault/models.py @@ -12,8 +12,9 @@ class BaseConnectionSecret(BaseModel): connection_id: str = Field(..., description="Unique connection identifier") model: str = Field(..., description="Model name (must match llm_config.yaml)") - environment: str = Field( - ..., description="Environment: production/development/test" + environment: Optional[str] = Field( + None, + description="Environment: production/development/test (legacy, not stored in new secrets)", ) tags: List[str] = Field(default_factory=list, description="Connection tags") diff --git a/src/llm_orchestrator_config/vault/secret_resolver.py b/src/llm_orchestrator_config/vault/secret_resolver.py index 5674c95a..2c32887e 100644 --- a/src/llm_orchestrator_config/vault/secret_resolver.py +++ b/src/llm_orchestrator_config/vault/secret_resolver.py @@ -3,8 +3,8 @@ import threading from datetime import datetime, timedelta from typing import Optional, Dict, Any, Union, List -from pydantic import BaseModel -from loguru import logger +from pydantic import BaseModel, ValidationError +from src.loki_logger import LokiLogger from llm_orchestrator_config.vault.vault_client import ( VaultAgentClient, @@ -17,6 +17,9 @@ ) from llm_orchestrator_config.vault.exceptions import VaultConnectionError +# Initialize Loki logger +logger = LokiLogger(service_name="secret-resolver") + class CachedSecret(BaseModel): """Cached secret with TTL information.""" @@ -119,55 +122,45 @@ def get_secret_for_model( except VaultConnectionError: logger.warning(f"Vault unavailable, trying fallback for {vault_path}") return self._get_fallback(vault_path) + except ValidationError as e: + # The secret exists but does not match the provider's schema. A + # last-known-good fallback would mask a stored-data problem that no + # retry can fix, so report it as a misconfiguration and stop here. + logger.error( + f"Malformed secret at {vault_path} for provider {provider} - the " + f"stored value does not match {get_secret_model(provider).__name__} " + f"and must be rewritten: {e}" + ) + return None except Exception as e: logger.error(f"Error resolving secret for {vault_path}: {e}") return self._get_fallback(vault_path) def list_available_models(self, provider: str, environment: str) -> list[str]: - """List available models for a provider and environment. + """List available models for a provider. Args: provider: Provider name (azure_openai, aws_bedrock) - environment: Environment (production, development, test) + environment: Environment (kept for API compatibility, not used in path) Returns: - List of available model names + List of available vault UUIDs under this provider """ - if environment == "production": - # For production: Check provider/production path for available models - production_path = f"llm/connections/{provider}/{environment}" - try: - models = self.vault_client.list_secrets(production_path) - if models: - logger.debug( - f"Found {len(models)} production models for {provider}: {models}" - ) - return models - else: - logger.debug(f"No production models found for {provider}") - return [] - - except Exception as e: - logger.debug(f"Provider {provider} not available in production: {e}") - return [] - else: - # For dev/test: Use existing logic with connection_id paths - base_path = f"llm/connections/{provider}/{environment}" - try: - models = self.vault_client.list_secrets(base_path) - if models: - logger.debug( - f"Found {len(models)} models for {provider}/{environment}" - ) - return models - else: - logger.debug(f"No models found for {provider}/{environment}") - return [] - - except Exception as e: - logger.error(f"Error listing models for {provider}/{environment}: {e}") + # List all UUIDs under provider (no environment in path) + base_path = f"llm/connections/{provider}" + try: + entries = self.vault_client.list_secrets(base_path) + if entries: + logger.debug(f"Found {len(entries)} entries for {provider}: {entries}") + return entries + else: + logger.debug(f"No entries found for {provider}") return [] + except Exception as e: + logger.error(f"Error listing models for {provider}: {e}") + return [] + def refresh_secret(self, vault_path: str) -> bool: """Manually refresh a specific secret. @@ -238,16 +231,17 @@ def _build_vault_path( ) -> str: """Build Vault path for a secret. - For production: llm/connections/{provider}/production/{model_name} - For dev/test: use connection_id if provided, otherwise model name + Uses UUID (connection_id) as the sole path terminal. + Environment is not part of the path — swaps are DB-only. + + Path format: llm/connections/{provider}/{uuid} """ - if environment == "production": - # Production uses provider/production/model_name path - return f"llm/connections/{provider}/{environment}/{model_name}" - else: - # Development/test can use connection_id or fall back to model name - model_identifier = connection_id if connection_id else model_name - return f"llm/connections/{provider}/{environment}/{model_identifier}" + if not connection_id: + raise ValueError( + f"connection_id (vault_uuid) is required to build vault path " + f"for provider={provider}, environment={environment}" + ) + return f"llm/connections/{provider}/{connection_id}" def _get_from_cache( self, vault_path: str @@ -373,6 +367,15 @@ def get_embedding_secret_for_model( f"Vault unavailable, trying fallback for embedding {vault_path}" ) return self._get_fallback(vault_path) + except ValidationError as e: + # See get_secret_for_model: a schema mismatch is a stored-data + # problem, not a transient one, so do not serve a stale fallback. + logger.error( + f"Malformed embedding secret at {vault_path} for provider {provider} - " + f"the stored value does not match " + f"{get_secret_model(provider).__name__} and must be rewritten: {e}" + ) + return None except Exception as e: logger.error(f"Error resolving embedding secret for {vault_path}: {e}") return self._get_fallback(vault_path) @@ -380,60 +383,32 @@ def get_embedding_secret_for_model( def list_available_embedding_models( self, provider: str, environment: str ) -> List[str]: - """List available embedding models for a provider and environment. + """List available embedding models for a provider. Args: provider: Provider name (azure_openai, aws_bedrock) - environment: Environment (production, development, test) + environment: Environment (kept for API compatibility, not used in path) Returns: - List of available embedding model names + List of available vault UUIDs under this provider """ - if environment == "production": - # For production: Check embeddings/connections/provider/production path - production_path: str = f"embeddings/connections/{provider}/{environment}" - try: - models_result: Optional[list[str]] = self.vault_client.list_secrets( - production_path - ) - if models_result: - logger.debug( - f"Found {len(models_result)} production embedding models for {provider}: {models_result}" - ) - return models_result - else: - logger.debug(f"No production embedding models found for {provider}") - return [] - - except Exception as e: + # List all UUIDs under provider (no environment in path) + base_path: str = f"embeddings/connections/{provider}" + try: + entries: Optional[list[str]] = self.vault_client.list_secrets(base_path) + if entries: logger.debug( - f"Provider {provider} embedding models not available in production: {e}" - ) - return [] - else: - # For dev/test: Use embeddings path with connection_id paths - base_path: str = f"embeddings/connections/{provider}/{environment}" - try: - models_result: Optional[list[str]] = self.vault_client.list_secrets( - base_path - ) - if models_result: - logger.debug( - f"Found {len(models_result)} embedding models for {provider}/{environment}" - ) - return models_result - else: - logger.debug( - f"No embedding models found for {provider}/{environment}" - ) - return [] - - except Exception as e: - logger.error( - f"Error listing embedding models for {provider}/{environment}: {e}" + f"Found {len(entries)} embedding entries for {provider}: {entries}" ) + return entries + else: + logger.debug(f"No embedding entries found for {provider}") return [] + except Exception as e: + logger.error(f"Error listing embedding models for {provider}: {e}") + return [] + def _build_embedding_vault_path( self, provider: str, @@ -443,23 +418,14 @@ def _build_embedding_vault_path( ) -> str: """Build Vault path for embedding secrets. - Args: - provider: Provider name (azure_openai, aws_bedrock) - environment: Environment (production, development, test) - model_name: Embedding model name - connection_id: Optional connection ID for dev/test environments + Uses UUID (connection_id) as the sole path terminal. + Environment is not part of the path — swaps are DB-only. - Returns: - Vault path for embedding secrets - - Examples: - Production: embeddings/connections/azure_openai/production/text-embedding-3-large - Dev/Test: embeddings/connections/azure_openai/development/dev-conn-123 + Path format: embeddings/connections/{provider}/{uuid} """ - if environment == "production": - # Production uses embeddings/connections/{provider}/production/{model_name} path - return f"embeddings/connections/{provider}/{environment}/{model_name}" - else: - # Development/test can use connection_id or fall back to model name - model_identifier: str = connection_id if connection_id else model_name - return f"embeddings/connections/{provider}/{environment}/{model_identifier}" + if not connection_id: + raise ValueError( + f"connection_id (vault_uuid) is required to build embedding vault path " + f"for provider={provider}, environment={environment}" + ) + return f"embeddings/connections/{provider}/{connection_id}" diff --git a/src/llm_orchestrator_config/vault/vault_client.py b/src/llm_orchestrator_config/vault/vault_client.py index 241f019e..51389421 100644 --- a/src/llm_orchestrator_config/vault/vault_client.py +++ b/src/llm_orchestrator_config/vault/vault_client.py @@ -4,7 +4,7 @@ import threading from pathlib import Path from typing import Optional, Dict, Any, cast -from loguru import logger +from src.loki_logger import LokiLogger import hvac from hvac.exceptions import InvalidPath, Forbidden @@ -14,6 +14,9 @@ VaultTokenError, ) +# Initialize Loki logger +logger = LokiLogger(service_name="vault-client") + # Global singleton instance _vault_client_instance: Optional["VaultAgentClient"] = None _vault_client_lock = threading.Lock() diff --git a/src/loki_logger.py b/src/loki_logger.py new file mode 100644 index 00000000..d780fcec --- /dev/null +++ b/src/loki_logger.py @@ -0,0 +1,184 @@ +#!/usr/bin/env python3 +""" +Loki Logger for RAG Module +Sends logs directly to Loki API for centralized logging + +[CANONICAL SOURCE] +This is the single source of truth for LokiLogger. +Two copies exist for environments where `src` is not a Python package: + - grafana-configs/loki_logger.py — mounted into CronManager container at runtime + - src/vector_indexer/loki_logger.py — used when running vector_indexer scripts locally + +If you change the logger logic here, apply the same change to both copies. +""" + +import json +import sys +import time +from datetime import datetime +from threading import Thread +from queue import Full, Queue + +import requests + + +class LokiLogger: + """Simple logger that sends logs directly to Loki API with async background thread""" + + _instances: dict[str, "LokiLogger"] = {} + + def __new__( + cls, loki_url: str = "http://loki:3100", service_name: str = "default" + ) -> "LokiLogger": + key = f"{loki_url}:{service_name}" + if key not in cls._instances: + cls._instances[key] = super().__new__(cls) + return cls._instances[key] + + def __init__( + self, loki_url: str = "http://loki:3100", service_name: str = "default" + ) -> None: + """ + Initialize LokiLogger + + Args: + loki_url: URL for Loki service (default: container URL in bykstack network) + service_name: Name of the service for labeling logs + """ + if hasattr(self, "_initialized"): + return + self._initialized = True + self.loki_url = loki_url + self.service_name = service_name + self.session = requests.Session() + # Set default timeout for all requests + self.timeout = 5 + + # Queue for async log processing (bounded to avoid unbounded memory growth under load) + self.log_queue: Queue[tuple[str, str]] = Queue(maxsize=10_000) + + # Start background worker thread + self.worker_thread = Thread(target=self._process_logs, daemon=True) + self.worker_thread.start() + + def _process_logs(self) -> None: + """Background worker that processes log queue""" + while True: + try: + # Get log entry from queue (blocking) + level, message = self.log_queue.get() + + # Send to Loki + self._send_to_loki_sync(level, message) + + # Mark task as done + self.log_queue.task_done() + except Exception: + # Silently ignore errors in background thread + pass + + def _send_to_loki_sync(self, level: str, message: str) -> None: + """Send log entry directly to Loki API (called from background thread)""" + try: + # Create timestamp in nanoseconds (Loki requirement) + timestamp_ns = str(int(time.time() * 1_000_000_000)) + + # Prepare labels for Loki + labels = { + "service": self.service_name, + "level": level, + } + + # Create log entry + log_entry = { + "level": level, + "message": message, + "service": self.service_name, + } + + # Prepare Loki payload + payload = { + "streams": [ + { + "stream": labels, + "values": [[timestamp_ns, json.dumps(log_entry)]], + } + ] + } + + # Send to Loki + self.session.post( + f"{self.loki_url}/loki/api/v1/push", + json=payload, + headers={"Content-Type": "application/json"}, + timeout=self.timeout, + ) + + except Exception: + # Silently ignore logging errors to not affect main application + pass + + def _log(self, level: str, message: str) -> None: + """Queue log entry for async processing (non-blocking)""" + # Print to console immediately for real-time feedback. Written to + # stderr (not stdout) so callers that capture a subprocess's stdout + # for its return value (e.g. decrypt_vault_secrets.py) never pick up + # log lines mixed in with the actual output. + timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + print(f"[{timestamp}] {level: <8} | {message}", file=sys.stderr) # noqa: T201 + + # Queue for async Loki sending (non-blocking, drops log if queue is full) + try: + self.log_queue.put_nowait((level, message)) + except Full: + # Queue full (Loki may be slow/unreachable) - drop log to avoid blocking + pass + + def info(self, message: str, **kwargs: object) -> None: + """Log info message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("INFO", message) + + def error(self, message: str, **kwargs: object) -> None: + """Log error message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("ERROR", message) + + def warning(self, message: str, **kwargs: object) -> None: + """Log warning message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("WARNING", message) + + def debug(self, message: str, **kwargs: object) -> None: + """Log debug message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("DEBUG", message) + + def success(self, message: str, **kwargs: object) -> None: + """Log success message (loguru compatibility). Extra kwargs ignored.""" + self._log("SUCCESS", message) + + def critical(self, message: str, **kwargs: object) -> None: + """Log critical message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("CRITICAL", message) + + def exception(self, message: str, **kwargs: object) -> None: + """Log exception message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("EXCEPTION", message) + + def add(self, *args: object, **kwargs: object) -> None: + """ + No-op method for loguru compatibility. + + LokiLogger sends logs to Loki/console only, not to files. + This method exists for backward compatibility with loguru code. + """ + pass # Silently ignore - logs go to Loki instead of files + + def remove(self, *args: object, **kwargs: object) -> None: + """No-op method for loguru compatibility.""" + pass # Silently ignore + + def bind(self, **kwargs: object) -> "LokiLogger": + """No-op method for loguru compatibility. Returns self for chaining.""" + return self # Allow method chaining + + def opt(self, **kwargs: object) -> "LokiLogger": + """No-op method for loguru compatibility. Returns self for chaining.""" + return self # Allow method chaining diff --git a/src/models/conversation_history_models.py b/src/models/conversation_history_models.py new file mode 100644 index 00000000..a91d8e97 --- /dev/null +++ b/src/models/conversation_history_models.py @@ -0,0 +1,45 @@ +"""Pydantic models for conversation history state.""" + +import time +from typing import Optional + +from pydantic import BaseModel, Field + + +class ConversationRound(BaseModel): + """A single user+bot exchange in a conversation. + + Stored as part of ``ConversationHistoryState`` in Redis, keyed by chat_id. + """ + + user_message: str = Field(..., description="The user's message text for this round") + bot_message: str = Field(..., description="The bot's response text for this round") + timestamp: float = Field( + default_factory=time.time, + description="Unix timestamp of when the round was recorded", + ) + + +class ConversationHistoryState(BaseModel): + """Persisted conversation history for a chat session. + + Keyed by chat_id in Redis with a sliding 30-minute TTL. + Retains up to the most recent 10 rounds; older rounds are trimmed. + An optional summary field holds a condensed representation of rounds + that have been evicted. When a summarizer is injected into the store, + it is automatically generated and persisted as rounds are evicted. + If no summarizer is provided, the summary must be managed by the caller. + """ + + chat_id: str = Field(..., description="Unique conversation identifier") + rounds: list[ConversationRound] = Field( + default_factory=list, + description="Ordered list of conversation rounds (newest last), capped at 10", + ) + summary: Optional[str] = Field( + default=None, + description=( + "Optional condensed summary of earlier conversation turns that have been " + "evicted from the rounds list. Generated and stored by the caller." + ), + ) diff --git a/src/models/request_models.py b/src/models/request_models.py index fd7ab795..fc93e125 100644 --- a/src/models/request_models.py +++ b/src/models/request_models.py @@ -6,7 +6,10 @@ from src.utils.input_sanitizer import InputSanitizer from src.llm_orchestrator_config.stream_config import StreamConfig -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="request-models") class ConversationItem(BaseModel): @@ -54,13 +57,20 @@ class OrchestrationRequest(BaseModel): ..., description="Previous conversation history" ) url: str = Field(..., description="Source URL context") - environment: Literal["production", "testing", "development"] = Field( - ..., description="Environment context" + environment: Literal["production", "testing"] = Field( + "production", description="Environment context (defaults to production)" ) connection_id: Optional[str] = Field( - None, description="Optional connection identifier" + None, description="Vault UUID for the connection (required for testing)" ) + @model_validator(mode="after") + def validate_connection_id_for_testing(self) -> "OrchestrationRequest": + """Ensure connection_id is provided when environment is testing.""" + if self.environment == "testing" and not self.connection_id: + raise ValueError("connection_id is required when environment is 'testing'") + return self + @field_validator("message") @classmethod def validate_and_sanitize_message(cls, v: str) -> str: @@ -93,7 +103,6 @@ def validate_conversation_history( cls, v: List[ConversationItem] ) -> List[ConversationItem]: """Validate conversation history limits.""" - from loguru import logger # Limit number of conversation history items max_history_items = 100 @@ -164,6 +173,10 @@ class OrchestrationResponse(BaseModel): default=None, description="Optional list of choice buttons for MCQ step responses", ) + retrieval_context: Optional[List[Dict[str, Any]]] = Field( + default=None, + description="Retrieved chunks with metadata, populated only in EVAL_MODE", + ) # New models for embedding and context generation @@ -255,7 +268,7 @@ class TestOrchestrationRequest(BaseModel): """Model for simplified test orchestration request.""" message: str = Field(..., description="User's message/query") - environment: Literal["production", "testing", "development"] = Field( + environment: Literal["production", "testing"] = Field( ..., description="Environment context" ) connectionId: Optional[int] = Field( @@ -288,3 +301,16 @@ class TestOrchestrationResponse(BaseModel): chunks: Optional[List[ChunkInfo]] = Field( default=None, description="Retrieved chunks with rank and content" ) + + +class DeepEvalTestOrchestrationResponse(BaseModel): + """Extended response model for testing with additional evaluation data.""" + + chatId: str + llmServiceActive: bool + questionOutOfLLMScope: bool + inputGuardFailed: bool + content: str + retrieval_context: Optional[List[Dict[str, Any]]] = None + refined_questions: Optional[List[str]] = None + expected_output: Optional[str] = None # For DeepEval diff --git a/src/models/session_models.py b/src/models/session_models.py index d7950860..d79355fd 100644 --- a/src/models/session_models.py +++ b/src/models/session_models.py @@ -1,7 +1,69 @@ """Pydantic models for API tool session state.""" -from typing import Any +from typing import Any, Literal from pydantic import BaseModel, Field +import time + + +class LastCallContext(BaseModel): + """Stores the result of the most recent successful ATC API call for a chat session. + + Used by the L2 cache tier to enable follow-up query handling: + param inheritance, response questions, and partial param updates. + + For multi-intent sessions, ``ATCCacheStore.set_l2()`` stores a + ``list[LastCallContext]`` keyed by ``atc:last:{chat_id}``. Single-intent + calls store a single-element list for consistency. The follow-up detector + always searches this list by ``api_name``. + """ + + api_name: str = Field(..., description="Endpoint name that was called") + endpoint: dict[str, Any] = Field( + ..., description="Full endpoint definition as stored in Qdrant payload" + ) + collected_params: dict[str, Any] = Field( + ..., description="Parameter values that were passed to the API call" + ) + raw_response: Any = Field( + ..., + description="Parsed API JSON response; Any to accept either list or dict", + ) + original_query: str = Field( + ..., description="User's first-turn query that triggered this API call" + ) + timestamp: float = Field( + default_factory=time.time, + description="Unix timestamp of when the call was made (for TTL/staleness checks)", + ) + + +class EndpointSessionState(BaseModel): + """Per-endpoint parameter collection state for parallel API tool execution. + + One instance per endpoint in a parallel-mode session. Tracks what has been + collected for that specific endpoint and whether all required params are ready + (``completed=True``). Note: ``completed`` signals param-readiness only — API + execution is gated on ``AgenticLoopStatus.COMPLETED`` at the loop level. + """ + + endpoint: dict[str, Any] = Field( + ..., + description="The API endpoint definition (name, url, method, params, etc.)", + ) + collected_params: dict[str, Any] = Field( + default_factory=dict, + description="Parameters collected for this endpoint so far", + ) + completed: bool = Field( + default=False, + description=( + "True once all required params for this endpoint have been collected. " + "This flag signals param-readiness only — it does NOT indicate that the " + "API has been called. Actual API execution is triggered by " + "AgenticLoopStatus.COMPLETED at the loop level, after all endpoints " + "reach this state." + ), + ) class APIToolSession(BaseModel): @@ -57,3 +119,29 @@ class APIToolSession(BaseModel): "full original intent, not just the last short follow-up message." ), ) + + # ── Parallel mode fields (all optional/defaulted; empty = single mode) ──── + + execution_mode: Literal["single", "parallel"] = Field( + default="single", + description=( + "'single' or 'parallel' — drives which code path manages this session. " + "Single-mode sessions leave the three fields below at their defaults." + ), + ) + parallel_endpoints: list[EndpointSessionState] = Field( + default_factory=list, + description=( + "Per-endpoint state list populated when execution_mode='parallel'. " + "Each entry tracks collected params and completion for one endpoint. " + "Always empty in single mode." + ), + ) + active_endpoint_index: int = Field( + default=0, + ge=0, + description=( + "Index into parallel_endpoints of the endpoint currently being collected. " + "Only relevant when execution_mode='parallel'; always 0 in single mode." + ), + ) diff --git a/src/optimization/metrics/generator_metrics.py b/src/optimization/metrics/generator_metrics.py index acb0a89f..ca0973a9 100644 --- a/src/optimization/metrics/generator_metrics.py +++ b/src/optimization/metrics/generator_metrics.py @@ -6,7 +6,10 @@ from typing import Any, Dict, List import dspy from dspy.evaluate import SemanticF1 -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="generator-metrics") class GeneratorMetric: diff --git a/src/optimization/metrics/guardrails_metrics.py b/src/optimization/metrics/guardrails_metrics.py index 97d4f523..953f4ce9 100644 --- a/src/optimization/metrics/guardrails_metrics.py +++ b/src/optimization/metrics/guardrails_metrics.py @@ -5,7 +5,10 @@ from typing import Any, Dict, List import dspy -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="guardrails-metrics") class GuardrailsMetric: diff --git a/src/optimization/metrics/refiner_metrics.py b/src/optimization/metrics/refiner_metrics.py index 8550d3a6..d29da021 100644 --- a/src/optimization/metrics/refiner_metrics.py +++ b/src/optimization/metrics/refiner_metrics.py @@ -5,7 +5,10 @@ from typing import Any, Dict, List import dspy -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="refiner-metrics") class RefinementJudge(dspy.Signature): diff --git a/src/optimization/optimization_scripts/check_paths.py b/src/optimization/optimization_scripts/check_paths.py index ff05e211..8dc33d2d 100644 --- a/src/optimization/optimization_scripts/check_paths.py +++ b/src/optimization/optimization_scripts/check_paths.py @@ -4,7 +4,10 @@ from pathlib import Path from typing import Dict -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="check-paths") def get_directory_structure() -> tuple[Path, Path]: diff --git a/src/optimization/optimization_scripts/diagnose_guardrails_loader.py b/src/optimization/optimization_scripts/diagnose_guardrails_loader.py index 28909caf..0dab5233 100644 --- a/src/optimization/optimization_scripts/diagnose_guardrails_loader.py +++ b/src/optimization/optimization_scripts/diagnose_guardrails_loader.py @@ -7,9 +7,12 @@ sys.path.append(str(Path(__file__).parent.parent.parent)) -from loguru import logger +from src.loki_logger import LokiLogger from src.guardrails.optimized_guardrails_loader import OptimizedGuardrailsLoader +# Initialize Loki logger +logger = LokiLogger(service_name="diagnose-guardrails-loader") + def main() -> None: """Run diagnostics.""" diff --git a/src/optimization/optimization_scripts/extract_guardrails_prompts.py b/src/optimization/optimization_scripts/extract_guardrails_prompts.py index 8c2654ae..9522d44b 100644 --- a/src/optimization/optimization_scripts/extract_guardrails_prompts.py +++ b/src/optimization/optimization_scripts/extract_guardrails_prompts.py @@ -8,12 +8,29 @@ import yaml from pathlib import Path from typing import Dict, Any, Optional, List, Tuple -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="extract-guardrails-prompts") # Constants FULL_TRACEBACK_MSG = "Full traceback:" FEW_SHOT_EXAMPLES_HEADER = "\nFew-shot Examples (from optimization):" +# Output-rails streaming buffer, written into every generated config. +# +# OUTPUT_STREAMING_CONTEXT_SIZE MUST stay below OUTPUT_STREAMING_CHUNK_SIZE. +# NeMo's RollingBuffer drains with ``buffer = buffer[-context_size:]`` after each +# flush; if context_size >= chunk_size the buffer never falls back below the +# flush threshold, so every subsequent token flushes on its own and triggers a +# full self_check_output LLM round-trip. That makes streaming latency and cost +# scale linearly with answer length. +# +# Keep in sync with rails.output.streaming in src/guardrails/rails_config.yaml. +OUTPUT_STREAMING_CHUNK_SIZE = 200 +OUTPUT_STREAMING_CONTEXT_SIZE = 50 +OUTPUT_STREAMING_STREAM_FIRST = False + # Type aliases for better readability JsonDict = Dict[str, Any] PromptDict = Dict[str, Any] @@ -359,9 +376,17 @@ def _ensure_required_config_structure(base_config: Dict[str, Any]) -> None: # Set required streaming parameters (override existing values to ensure consistency) output_streaming["enabled"] = True - output_streaming["chunk_size"] = 200 - output_streaming["context_size"] = 300 - output_streaming["stream_first"] = False + output_streaming["chunk_size"] = OUTPUT_STREAMING_CHUNK_SIZE + output_streaming["context_size"] = OUTPUT_STREAMING_CONTEXT_SIZE + output_streaming["stream_first"] = OUTPUT_STREAMING_STREAM_FIRST + + if OUTPUT_STREAMING_CONTEXT_SIZE >= OUTPUT_STREAMING_CHUNK_SIZE: + raise ValueError( + f"Invalid output-rails streaming constants: context_size=" + f"{OUTPUT_STREAMING_CONTEXT_SIZE} must be < chunk_size=" + f"{OUTPUT_STREAMING_CHUNK_SIZE}. Generating a config with this " + f"setting would cause one guardrail LLM call per streamed token." + ) logger.info("✓ Ensured required rails and streaming configuration structure") diff --git a/src/optimization/optimization_scripts/inspect_guardrails_optimization.py b/src/optimization/optimization_scripts/inspect_guardrails_optimization.py index f9632da7..9693aaee 100644 --- a/src/optimization/optimization_scripts/inspect_guardrails_optimization.py +++ b/src/optimization/optimization_scripts/inspect_guardrails_optimization.py @@ -4,7 +4,10 @@ import json from pathlib import Path -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="inspect-guardrails-optimization") def main() -> None: diff --git a/src/optimization/optimization_scripts/run_all_optimizations.py b/src/optimization/optimization_scripts/run_all_optimizations.py index 275148f1..02c041e1 100644 --- a/src/optimization/optimization_scripts/run_all_optimizations.py +++ b/src/optimization/optimization_scripts/run_all_optimizations.py @@ -13,13 +13,15 @@ sys.path.append(str(Path(__file__).parent.parent)) import dspy -from loguru import logger - +from src.loki_logger import LokiLogger from llm_orchestrator_config import LLMManager from optimizers.guardrails_optimizer import optimize_guardrails from optimizers.refiner_optimizer import optimize_refiner from optimizers.generator_optimizer import optimize_generator +# Initialize Loki logger +logger = LokiLogger(service_name="run-all-optimizations") + # Constants TRACEBACK_MSG = "Full traceback:" diff --git a/src/optimization/optimization_scripts/split_datasets.py b/src/optimization/optimization_scripts/split_datasets.py index 3316dffa..76649943 100644 --- a/src/optimization/optimization_scripts/split_datasets.py +++ b/src/optimization/optimization_scripts/split_datasets.py @@ -7,11 +7,13 @@ from typing import List, Dict, Any, Tuple import random import sys +from src.loki_logger import LokiLogger # Add src to path for imports sys.path.append(str(Path(__file__).parent.parent)) -from loguru import logger +# Initialize Loki logger +logger = LokiLogger(service_name="split-datasets") def load_dataset(filepath: Path) -> List[Dict[str, Any]]: diff --git a/src/optimization/optimized_module_loader.py b/src/optimization/optimized_module_loader.py index 19803415..f502e839 100644 --- a/src/optimization/optimized_module_loader.py +++ b/src/optimization/optimized_module_loader.py @@ -10,7 +10,10 @@ from datetime import datetime import threading import dspy -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="optimized-module-loader") class OptimizedModuleLoader: diff --git a/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251105_114631_config.yaml b/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251105_114631_config.yaml index 7565f994..4f992b33 100644 --- a/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251105_114631_config.yaml +++ b/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251105_114631_config.yaml @@ -42,7 +42,8 @@ rails: streaming: enabled: True chunk_size: 200 - context_size: 300 + # Must stay < chunk_size; see src/guardrails/rails_config.yaml + context_size: 50 stream_first: False prompts: diff --git a/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251112_205121_config.yaml b/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251112_205121_config.yaml index 7565f994..4f992b33 100644 --- a/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251112_205121_config.yaml +++ b/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251112_205121_config.yaml @@ -42,7 +42,8 @@ rails: streaming: enabled: True chunk_size: 200 - context_size: 300 + # Must stay < chunk_size; see src/guardrails/rails_config.yaml + context_size: 50 stream_first: False prompts: diff --git a/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251114_050437_config.yaml b/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251114_050437_config.yaml index 25e90011..d32d82c6 100644 --- a/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251114_050437_config.yaml +++ b/src/optimization/optimized_modules/guardrails/guardrails_optimized_20251114_050437_config.yaml @@ -39,7 +39,8 @@ rails: streaming: enabled: true chunk_size: 200 - context_size: 300 + # Must stay < chunk_size; see src/guardrails/rails_config.yaml + context_size: 50 stream_first: false prompts: - task: self_check_input diff --git a/src/optimization/optimizers/generator_optimizer.py b/src/optimization/optimizers/generator_optimizer.py index d141ba8a..54a65eb4 100644 --- a/src/optimization/optimizers/generator_optimizer.py +++ b/src/optimization/optimizers/generator_optimizer.py @@ -12,13 +12,15 @@ sys.path.append(str(Path(__file__).parent.parent.parent)) import dspy -from loguru import logger - +from src.loki_logger import LokiLogger from optimization.metrics.generator_metrics import ( GeneratorMetric, calculate_generator_stats, ) +# Initialize Loki logger +logger = LokiLogger(service_name="generator-optimizer") + class ResponseGeneratorSignature(dspy.Signature): """ diff --git a/src/optimization/optimizers/guardrails_optimizer.py b/src/optimization/optimizers/guardrails_optimizer.py index 02d9e9a1..4fa38d8e 100644 --- a/src/optimization/optimizers/guardrails_optimizer.py +++ b/src/optimization/optimizers/guardrails_optimizer.py @@ -13,13 +13,16 @@ sys.path.append(str(Path(__file__).parent.parent.parent)) import dspy -from loguru import logger +from src.loki_logger import LokiLogger from optimization.metrics.guardrails_metrics import ( safety_weighted_accuracy, calculate_guardrails_stats, ) +# Initialize Loki logger +logger = LokiLogger(service_name="guardrails-optimizer") + class GuardrailsChecker(dspy.Signature): """ diff --git a/src/optimization/optimizers/refiner_optimizer.py b/src/optimization/optimizers/refiner_optimizer.py index 526ab9dc..f545a3c3 100644 --- a/src/optimization/optimizers/refiner_optimizer.py +++ b/src/optimization/optimizers/refiner_optimizer.py @@ -12,13 +12,15 @@ sys.path.append(str(Path(__file__).parent.parent.parent)) import dspy -from loguru import logger - +from src.loki_logger import LokiLogger from optimization.metrics.refiner_metrics import ( RefinerMetric, calculate_refiner_stats, ) +# Initialize Loki logger +logger = LokiLogger(service_name="refiner-optimizer") + class PromptRefinerSignature(dspy.Signature): """ diff --git a/src/prompt_refine_manager/prompt_refiner.py b/src/prompt_refine_manager/prompt_refiner.py index 5cbe30ec..5b857992 100644 --- a/src/prompt_refine_manager/prompt_refiner.py +++ b/src/prompt_refine_manager/prompt_refiner.py @@ -2,7 +2,7 @@ from typing import Any, Sequence, Optional, Dict, Union, cast, List import contextlib -import logging +from src.loki_logger import LokiLogger import dspy from pydantic import BaseModel, Field @@ -10,7 +10,8 @@ from src.utils.cost_utils import get_lm_usage_since from src.optimization.optimized_module_loader import get_module_loader -LOGGER = logging.getLogger(__name__) +# Initialize Loki logger for prompt refiner +LOGGER = LokiLogger(service_name="prompt-refiner") class ConversationHistory(BaseModel): diff --git a/src/response_generator/response_generate.py b/src/response_generator/response_generate.py index 3dffbfb5..ca065aee 100644 --- a/src/response_generator/response_generate.py +++ b/src/response_generator/response_generate.py @@ -2,21 +2,36 @@ from typing import List, Dict, Any, Tuple, AsyncIterator, Optional import re import dspy -import logging +from src.loki_logger import LokiLogger import asyncio import dspy.streaming from dspy.streaming import StreamListener +from langfuse import observe +from src.utils.observation_utils import ( + safe_observation_context, + update_observation_safe, +) from src.llm_orchestrator_config.llm_ochestrator_constants import OUT_OF_SCOPE_MESSAGE from src.utils.cost_utils import get_lm_usage_since from src.optimization.optimized_module_loader import get_module_loader from src.vector_indexer.constants import ResponseGenerationConstants -# Configure logging -logging.basicConfig( - level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s" -) -logger = logging.getLogger(__name__) +# Initialize Loki logger for response generator +logger = LokiLogger(service_name="response-generator") + + +def _get_current_model_name() -> str: + """Best-effort model name lookup from current DSPy LM.""" + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "model"): + model_name = lm.model + if isinstance(model_name, str) and model_name: + return model_name + except Exception: + pass + return "unknown" class ResponseGenerator(dspy.Signature): @@ -237,90 +252,135 @@ async def stream_response( Yields: Token strings as they arrive from the LLM """ - if max_blocks is None: - max_blocks = ResponseGenerationConstants.DEFAULT_MAX_BLOCKS + with safe_observation_context( + as_type="generation", + name="response_generation_streaming", + input={"question": question[:500], "chunks_count": len(chunks)}, + ) as _generation: + if max_blocks is None: + max_blocks = ResponseGenerationConstants.DEFAULT_MAX_BLOCKS - logger.info( - f"Starting NATIVE DSPy streaming for question with {len(chunks)} chunks" - ) - - # Apply custom instructions while keeping the user question first, if provided - augmented_question = question - if self._custom_instructions_prefix: - augmented_question = f"{question}\n\n{self._custom_instructions_prefix}" - logger.debug( - f"Applied custom instructions after question for streaming ({len(self._custom_instructions_prefix)} chars)" + logger.info( + f"Starting NATIVE DSPy streaming for question with {len(chunks)} chunks" ) - output_stream = None - try: - # Build context - context_blocks, citation_labels, has_real_context = ( - build_context_and_citations(chunks, use_top_k=max_blocks) - ) + history_length_before = 0 + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + + # Apply custom instructions while keeping the user question first, if provided + augmented_question = question + if self._custom_instructions_prefix: + augmented_question = f"{question}\n\n{self._custom_instructions_prefix}" + logger.debug( + f"Applied custom instructions after question for streaming ({len(self._custom_instructions_prefix)} chars)" + ) - if not has_real_context: - logger.warning( - "No real context available for streaming, yielding nothing." + output_stream = None + streamed_tokens: List[str] = [] + final_answer_text: str = "" + try: + # Build context + context_blocks, citation_labels, has_real_context = ( + build_context_and_citations(chunks, use_top_k=max_blocks) ) - return - # Get the streamified predictor - stream_predictor = self._get_stream_predictor() + if not has_real_context: + logger.warning( + "No real context available for streaming, yielding nothing." + ) + return - # Call the streamified predictor with augmented question - logger.info("Calling streamified predictor with signature inputs...") - output_stream = stream_predictor( - question=augmented_question, - context_blocks=context_blocks, - citations=citation_labels, - ) + # Get the streamified predictor + stream_predictor = self._get_stream_predictor() - stream_started = False - try: - async for chunk in output_stream: - # The stream yields StreamResponse objects for tokens - # and a final Prediction object - if isinstance(chunk, dspy.streaming.StreamResponse): - if chunk.signature_field_name == "answer": - stream_started = True - yield chunk.chunk # Yield the token string - elif isinstance(chunk, dspy.Prediction): - # The final prediction object is yielded last - logger.info( - "Streaming complete, final Prediction object received." - ) - full_answer = getattr(chunk, "answer", "[No answer field]") - logger.debug(f"Full streamed answer: {full_answer}") - except GeneratorExit: - # Generator was closed early (e.g., by guardrails violation) - logger.info("Stream generator closed early - cleaning up") - # Properly close the stream + # Call the streamified predictor with augmented question + logger.info("Calling streamified predictor with signature inputs...") + output_stream = stream_predictor( + question=augmented_question, + context_blocks=context_blocks, + citations=citation_labels, + ) + + stream_started = False + try: + async for chunk in output_stream: + # The stream yields StreamResponse objects for tokens + # and a final Prediction object + if isinstance(chunk, dspy.streaming.StreamResponse): + if chunk.signature_field_name == "answer": + stream_started = True + streamed_tokens.append(chunk.chunk) + yield chunk.chunk # Yield the token string + elif isinstance(chunk, dspy.Prediction): + # The final prediction object is yielded last + logger.info( + "Streaming complete, final Prediction object received." + ) + full_answer = getattr(chunk, "answer", "[No answer field]") + if isinstance(full_answer, str): + final_answer_text = full_answer + logger.debug(f"Full streamed answer: {full_answer}") + except GeneratorExit: + # Generator was closed early (e.g., by guardrails violation) + logger.info("Stream generator closed early - cleaning up") + # Properly close the stream + if output_stream is not None: + try: + await output_stream.aclose() + except Exception as close_error: + logger.debug( + f"Error closing stream (expected): {close_error}" + ) + output_stream = None # Prevent double-close in finally block + raise + + if not stream_started: + logger.warning( + "Streaming call finished but no 'answer' tokens were received." + ) + + except Exception as e: + logger.error(f"Error during native DSPy streaming: {str(e)}") + logger.exception("Full traceback:") + raise + finally: + # Ensure cleanup even if exception occurs if output_stream is not None: try: await output_stream.aclose() - except Exception as close_error: - logger.debug(f"Error closing stream (expected): {close_error}") - output_stream = None # Prevent double-close in finally block - raise + except Exception as cleanup_error: + logger.debug(f"Error during cleanup (aclose): {cleanup_error}") + + usage_info = get_lm_usage_since(history_length_before) + stream_output = final_answer_text or "".join(streamed_tokens) + if not stream_output: + stream_output = "[stream completed with empty output]" - if not stream_started: + try: + _generation.update( + model=_get_current_model_name(), + usage_details={ + "input": usage_info.get("total_prompt_tokens", 0), + "output": usage_info.get("total_completion_tokens", 0), + "total": usage_info.get("total_tokens", 0), + }, + cost_details={ + "total": usage_info.get("total_cost", 0.0), + }, + metadata={ + "num_calls": usage_info.get("num_calls", 0), + "streaming": True, + }, + output=stream_output, + ) + except Exception as e: logger.warning( - "Streaming call finished but no 'answer' tokens were received." + f"Failed to update streaming generation usage in Langfuse: {e}" ) - except Exception as e: - logger.error(f"Error during native DSPy streaming: {str(e)}") - logger.exception("Full traceback:") - raise - finally: - # Ensure cleanup even if exception occurs - if output_stream is not None: - try: - await output_stream.aclose() - except Exception as cleanup_error: - logger.debug(f"Error during cleanup (aclose): {cleanup_error}") - + @observe(name="response_scope_check", as_type="generation") async def check_scope_quick( self, question: str, @@ -340,12 +400,34 @@ async def check_scope_quick( """ if max_blocks is None: max_blocks = ResponseGenerationConstants.DEFAULT_MAX_BLOCKS + + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception: + pass + try: context_blocks, _, has_real_context = build_context_and_citations( chunks, use_top_k=max_blocks ) if not has_real_context: + usage_info = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "question": question, + "chunks_count": len(chunks), + "max_blocks": max_blocks, + }, + output_data={"out_of_scope": True, "reason": "no_context"}, + metadata={ + "model": _get_current_model_name(), + "usage": usage_info, + }, + ) return True # Use DSPy to quickly check scope @@ -354,6 +436,19 @@ async def check_scope_quick( ) out_of_scope = getattr(result, "out_of_scope", False) + usage_info = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "question": question, + "chunks_count": len(chunks), + "max_blocks": max_blocks, + }, + output_data={"out_of_scope": bool(out_of_scope)}, + metadata={ + "model": _get_current_model_name(), + "usage": usage_info, + }, + ) logger.info( f"Quick scope check result: {'OUT OF SCOPE' if out_of_scope else 'IN SCOPE'}" ) @@ -362,6 +457,19 @@ async def check_scope_quick( except Exception as e: logger.error(f"Scope check error: {e}") + usage_info = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "question": question, + "chunks_count": len(chunks), + "max_blocks": max_blocks, + }, + output_data={"out_of_scope": False, "error": str(e)}, + metadata={ + "model": _get_current_model_name(), + "usage": usage_info, + }, + ) # On error, assume in-scope to allow generation to proceed return False @@ -393,6 +501,7 @@ def _validate_prediction(self, pred: dspy.Prediction) -> bool: logger.warning(f"Validation failed: {e}") return False + @observe(name="response_generation_non_streaming", as_type="generation") def forward( self, question: str, @@ -451,6 +560,24 @@ def forward( answer = OUT_OF_SCOPE_MESSAGE scope_flag = True + update_observation_safe( + input_data={ + "question": question, + "chunks_count": len(chunks), + "max_blocks": max_blocks, + "attempts": attempts, + "valid_prediction": False, + }, + output_data={ + "questionOutOfLLMScope": scope_flag, + "answer_preview": answer[:500], + }, + metadata={ + "model": _get_current_model_name(), + "usage": usage_info, + }, + ) + return { "answer": answer, "questionOutOfLLMScope": scope_flag, @@ -464,6 +591,24 @@ def forward( logger.warning("Flipping out-of-scope to True based on heuristics.") scope = True + update_observation_safe( + input_data={ + "question": question, + "chunks_count": len(chunks), + "max_blocks": max_blocks, + "attempts": attempts, + "valid_prediction": True, + }, + output_data={ + "questionOutOfLLMScope": scope, + "answer_preview": ans[:500], + }, + metadata={ + "model": _get_current_model_name(), + "usage": usage_info, + }, + ) + return { "answer": ans.strip(), "questionOutOfLLMScope": scope, diff --git a/src/tool_classifier/__init__.py b/src/tool_classifier/__init__.py index 62f24ef6..793cdd07 100644 --- a/src/tool_classifier/__init__.py +++ b/src/tool_classifier/__init__.py @@ -18,7 +18,10 @@ AgenticLoopResult, APICallResult, ClassificationResult, + MultiAPICallResult, ) +from tool_classifier.multi_api_caller import MultiAPICaller +from tool_classifier.multi_response_formatter import MultiResponseFormatterModule __all__ = [ "AgenticLoop", @@ -28,6 +31,9 @@ "APICallResult", "APIResponseFormatterModule", "ClassificationResult", + "MultiAPICaller", + "MultiAPICallResult", + "MultiResponseFormatterModule", "ToolClassifier", "WorkflowType", ] diff --git a/src/tool_classifier/agentic_loop.py b/src/tool_classifier/agentic_loop.py index 731a36f5..2f136ad2 100644 --- a/src/tool_classifier/agentic_loop.py +++ b/src/tool_classifier/agentic_loop.py @@ -1,10 +1,10 @@ """Standalone agentic loop for multi-turn parameter collection.""" import asyncio +import time from typing import Any, Dict, List, Optional -from loguru import logger - +from src.loki_logger import LokiLogger from src.utils.api_tool_session_store import APIToolSessionStore from tool_classifier.constants import ( CONTINUATION_QUESTION, @@ -12,10 +12,13 @@ CONTINUATION_QUESTION_RU, CONTINUATION_TURN, ) +from tool_classifier.continuation_utils import detect_continuation_response from tool_classifier.enums import AgenticLoopStatus from tool_classifier.models import AgenticLoopResult from tool_classifier.param_extractor import ParamExtractionModule +logger = LokiLogger(service_name="api-tool-calling") + _CONTINUATION_QUESTIONS: dict[str, str] = { "en": CONTINUATION_QUESTION, "et": CONTINUATION_QUESTION_ET, @@ -23,25 +26,6 @@ } -_YES_RESPONSES = frozenset( - { - "yes", - "y", - "jah", - "ja", - "да", - "ok", - "okay", - "sure", - "please", - "continue", - "jätka", - "продолжить", - "absolutely", - } -) - - class AgenticLoop: """Stateless multi-turn parameter collection loop. @@ -107,6 +91,7 @@ async def run_turn( continuation_turn: int = CONTINUATION_TURN, session_language: str = "en", continuation_language: Optional[str] = None, + seeded_params: Optional[Dict[str, Any]] = None, ) -> AgenticLoopResult: """Process one user turn of the parameter-collection loop. @@ -158,24 +143,30 @@ async def run_turn( """ updated_turn_count = turn_count + 1 + logger.info( + f"AgenticLoop: loop turn started | event_type=loop_turn_started chat_id={chat_id} turn_count={turn_count} max_turns={max_turns} awaiting_continuation={awaiting_continuation}" + ) + + # Seed inherited params from L2 follow-up detection — turn 0 only. + # seeded_params take lower priority than anything already in collected_params + # (i.e. values explicitly set by the session take precedence). + if turn_count == 0 and seeded_params: + collected_params = {**seeded_params, **collected_params} + # Step 0 — Continuation decision: user is responding to the yes/no prompt original_awaiting_continuation = awaiting_continuation if awaiting_continuation: wants_to_continue = self._detect_continuation_response(user_message) if wants_to_continue: logger.debug( - "AgenticLoop: user chose to continue on turn {} for chat_id={}", - turn_count, - chat_id, + f"AgenticLoop: user chose to continue on turn {turn_count} for chat_id={chat_id}" ) # Reset the flag so normal extraction takes over from here. awaiting_continuation = False else: logger.info( - "AgenticLoop: user chose to exit on turn {} for chat_id={}, " - "falling back to RAG", - turn_count, - chat_id, + f"AgenticLoop: user chose to exit on turn {turn_count} for chat_id={chat_id}, " + "falling back to RAG" ) return AgenticLoopResult( status=AgenticLoopStatus.MAX_TURNS_REACHED, @@ -187,9 +178,7 @@ async def run_turn( # Step 1 — Turn limit guard (no session save — caller deletes) if turn_count >= max_turns: logger.warning( - "AgenticLoop: max_turns={} reached for chat_id={}, abandoning", - max_turns, - chat_id, + f"AgenticLoop: max_turns={max_turns} reached for chat_id={chat_id}, abandoning" ) return AgenticLoopResult( status=AgenticLoopStatus.MAX_TURNS_REACHED, @@ -199,6 +188,7 @@ async def run_turn( ) # Step 2 — Extract params from the current user message + _t0 = time.time() try: extraction = await asyncio.to_thread( self._param_extractor, @@ -207,13 +197,16 @@ async def run_turn( conversation_history, collected_params, session_language, + turn_count, + ) + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.debug( + f"AgenticLoop: param extraction complete | event_type=param_extraction_complete chat_id={chat_id} turn_count={turn_count} extracted_count={len(extraction['extracted_params'])} duration_ms={_duration_ms}" ) except Exception as exc: + _duration_ms = round((time.time() - _t0) * 1000, 1) logger.error( - "AgenticLoop: param extraction failed on turn {} for chat_id={}: {}", - turn_count, - chat_id, - exc, + f"AgenticLoop: param extraction failed on turn {turn_count} for chat_id={chat_id}: {exc}" ) # If a continuation decision was already consumed this turn, persist the # updated flag so the next user message is not misread as another @@ -249,11 +242,13 @@ async def run_turn( } all_collected = required_param_names.issubset(merged_params.keys()) + logger.debug( + f"AgenticLoop: params merged | event_type=params_merged chat_id={chat_id} turn_count={turn_count} required_count={len(required_param_names)} collected_count={len(merged_params)} missing_count={len(required_param_names - merged_params.keys())}" + ) + if all_collected: - logger.debug( - "AgenticLoop: all required params collected on turn {} for chat_id={}", - turn_count, - chat_id, + logger.info( + f"AgenticLoop: loop completed | event_type=loop_completed chat_id={chat_id} turn_count={turn_count} status=completed collected_count={len(merged_params)} duration_ms={_duration_ms}" ) await self._save_session( chat_id, merged_params, updated_turn_count, awaiting_continuation=False @@ -267,18 +262,13 @@ async def run_turn( # Step 5 — Still missing params logger.debug( - "AgenticLoop: turn {} for chat_id={} — still missing: {}", - turn_count, - chat_id, - extraction["missing_required"], + f"AgenticLoop: loop needs input | event_type=loop_needs_input chat_id={chat_id} turn_count={turn_count} missing_params={extraction['missing_required']} status=needs_input" ) # At exactly the continuation threshold, ask whether to keep going. if updated_turn_count == continuation_turn: logger.info( - "AgenticLoop: continuation threshold reached on turn {} for chat_id={}", - turn_count, - chat_id, + f"AgenticLoop: continuation threshold reached | event_type=continuation_threshold_reached chat_id={chat_id} turn_count={turn_count} continuation_turn={continuation_turn} missing_count={len(extraction['missing_required'])}" ) effective_continuation_lang = continuation_language or session_language continuation_q = _CONTINUATION_QUESTIONS.get( @@ -317,6 +307,7 @@ async def stream_run_turn( continuation_turn: int = CONTINUATION_TURN, session_language: str = "en", continuation_language: Optional[str] = None, + seeded_params: Optional[Dict[str, Any]] = None, ) -> tuple[AgenticLoopResult, List[str]]: """Process one user turn like :meth:`run_turn` but stream clarifying_question tokens. @@ -332,23 +323,27 @@ async def stream_run_turn( """ updated_turn_count = turn_count + 1 + logger.info( + f"AgenticLoop: loop turn started | event_type=loop_turn_started chat_id={chat_id} turn_count={turn_count} max_turns={max_turns} awaiting_continuation={awaiting_continuation}" + ) + + # Seed inherited params from L2 follow-up detection — turn 0 only. + if turn_count == 0 and seeded_params: + collected_params = {**seeded_params, **collected_params} + # Step 0 — Continuation decision original_awaiting_continuation = awaiting_continuation if awaiting_continuation: wants_to_continue = self._detect_continuation_response(user_message) if wants_to_continue: - logger.debug( - "AgenticLoop: user chose to continue on turn {} for chat_id={}", - turn_count, - chat_id, + logger.info( + f"AgenticLoop: continuation user accepted | event_type=continuation_user_accepted chat_id={chat_id} turn_count={turn_count}" ) awaiting_continuation = False else: logger.info( - "AgenticLoop: user chose to exit on turn {} for chat_id={}, " - "falling back to RAG", - turn_count, - chat_id, + f"AgenticLoop: user chose to exit on turn {turn_count} for chat_id={chat_id}, " + "falling back to RAG" ) return ( AgenticLoopResult( @@ -363,9 +358,7 @@ async def stream_run_turn( # Step 1 — Turn limit guard if turn_count >= max_turns: logger.warning( - "AgenticLoop: max_turns={} reached for chat_id={}, abandoning", - max_turns, - chat_id, + f"AgenticLoop: max_turns={max_turns} reached for chat_id={chat_id}, abandoning" ) return ( AgenticLoopResult( @@ -378,6 +371,7 @@ async def stream_run_turn( ) # Step 2 — Stream-extract params from the current user message + _t0 = time.time() try: question_tokens, extraction = await self._param_extractor.stream_forward( user_message=user_message, @@ -385,13 +379,16 @@ async def stream_run_turn( conversation_history=conversation_history, already_collected=collected_params, session_language=session_language, + turn_count=turn_count, + ) + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.debug( + f"AgenticLoop: param extraction complete | event_type=param_extraction_complete chat_id={chat_id} turn_count={turn_count} extracted_count={len(extraction['extracted_params'])} duration_ms={_duration_ms}" ) except Exception as exc: + _duration_ms = round((time.time() - _t0) * 1000, 1) logger.error( - "AgenticLoop: stream param extraction failed on turn {} for chat_id={}: {}", - turn_count, - chat_id, - exc, + f"AgenticLoop: stream param extraction failed on turn {turn_count} for chat_id={chat_id}: {exc}" ) if awaiting_continuation != original_awaiting_continuation: await self._save_session( @@ -424,11 +421,13 @@ async def stream_run_turn( } all_collected = required_param_names.issubset(merged_params.keys()) + logger.debug( + f"AgenticLoop: params merged | event_type=params_merged chat_id={chat_id} turn_count={turn_count} required_count={len(required_param_names)} collected_count={len(merged_params)} missing_count={len(required_param_names - merged_params.keys())}" + ) + if all_collected: logger.debug( - "AgenticLoop: all required params collected on turn {} for chat_id={}", - turn_count, - chat_id, + f"AgenticLoop: all required params collected on turn {turn_count} for chat_id={chat_id}" ) await self._save_session( chat_id, merged_params, updated_turn_count, awaiting_continuation=False @@ -445,17 +444,12 @@ async def stream_run_turn( # Step 5 — Still missing params logger.debug( - "AgenticLoop: turn {} for chat_id={} — still missing: {}", - turn_count, - chat_id, - extraction["missing_required"], + f"AgenticLoop: turn {turn_count} for chat_id={chat_id} — still missing: {extraction['missing_required']}" ) if updated_turn_count == continuation_turn: logger.info( - "AgenticLoop: continuation threshold reached on turn {} for chat_id={}", - turn_count, - chat_id, + f"AgenticLoop: continuation threshold reached on turn {turn_count} for chat_id={chat_id}" ) effective_continuation_lang = continuation_language or session_language continuation_q = _CONTINUATION_QUESTIONS.get( @@ -508,8 +502,7 @@ async def _save_session( try: if self._session_store is None: logger.debug( - "AgenticLoop: session store unavailable — skipping save for chat_id={}", - chat_id, + f"AgenticLoop: session store unavailable — skipping save for chat_id={chat_id}" ) return await self._session_store.update( @@ -520,18 +513,13 @@ async def _save_session( ) except Exception as exc: logger.error( - "AgenticLoop: failed to save session for chat_id={}: {}", - chat_id, - exc, + f"AgenticLoop: failed to save session for chat_id={chat_id}: {exc}" ) def _detect_continuation_response(self, user_message: str) -> bool: """Detect whether the user's message indicates they want to continue. - Checks the normalised (lower-cased, stripped) message against a set of - known affirmative responses in Estonian, English, and Russian. - Any response that is not clearly affirmative is treated as a "no" so - the loop falls back to the RAG workflow. + Delegates to :func:`~continuation_utils.detect_continuation_response`. Args: user_message: The raw user message to inspect. @@ -539,5 +527,4 @@ def _detect_continuation_response(self, user_message: str) -> bool: Returns: True if the user wants to continue, False otherwise. """ - normalised = user_message.strip().lower() - return normalised in _YES_RESPONSES + return detect_continuation_response(user_message) diff --git a/src/tool_classifier/api_caller.py b/src/tool_classifier/api_caller.py index cbf38a4b..3ee4cccb 100644 --- a/src/tool_classifier/api_caller.py +++ b/src/tool_classifier/api_caller.py @@ -6,7 +6,8 @@ from typing import Any import httpx -from loguru import logger +from src.loki_logger import LokiLogger +from src.utils.error_utils import generate_error_id from llm_orchestrator_config.llm_ochestrator_constants import get_localized_message from tool_classifier.constants import ( @@ -23,7 +24,8 @@ SERVICE_UNAVAILABLE_MESSAGES, ) from tool_classifier.models import APICallResult -from src.utils.error_utils import generate_error_id, log_error_with_context + +logger = LokiLogger(service_name="api-tool-calling") @dataclass @@ -86,7 +88,9 @@ def can_execute(self, url: str) -> bool: if time.time() - breaker.last_failure_time >= self._cooldown_seconds: breaker.state = CB_STATE_HALF_OPEN breaker.probe_in_flight = True - logger.info(f"[CircuitBreaker] {url!r} → HALF_OPEN (probe allowed)") + logger.info( + f"CircuitBreaker: circuit half-open | event_type=circuit_breaker_half_open url={url}" + ) return True return False # HALF_OPEN: allow exactly one probe request through; gate subsequent @@ -101,8 +105,7 @@ def record_success(self, url: str) -> None: breaker = self._get_state(url) if breaker.state != CB_STATE_CLOSED: logger.info( - f"[CircuitBreaker] {url!r} → CLOSED " - f"(recovered after {breaker.failure_count} failure(s))" + f"CircuitBreaker: circuit recovered | event_type=circuit_breaker_recovered url={url} failure_count={breaker.failure_count}" ) breaker.state = CB_STATE_CLOSED breaker.failure_count = 0 @@ -120,8 +123,7 @@ def record_failure(self, url: str) -> None: if breaker.failure_count >= self._failure_threshold: if breaker.state != CB_STATE_OPEN: logger.warning( - f"[CircuitBreaker] {url!r} → OPEN " - f"after {breaker.failure_count} failure(s)" + f"CircuitBreaker: circuit opened | event_type=circuit_breaker_opened url={url} failure_count={breaker.failure_count} threshold={self._failure_threshold}" ) breaker.state = CB_STATE_OPEN @@ -192,7 +194,7 @@ async def call( if not self._circuit_breaker.can_execute(url): logger.warning( - f"[APICaller] Circuit breaker OPEN for {url!r} — rejecting call" + f"APICaller: api call rejected | event_type=api_call_rejected_circuit_open url={url} method={method_upper}" ) return APICallResult( success=False, @@ -202,6 +204,10 @@ async def call( ) effective_timeout = timeout if timeout is not None else self._default_timeout + logger.debug( + f"APICaller: api call started | event_type=api_call_started url={url} method={method_upper} timeout={effective_timeout}" + ) + _t0 = time.time() try: async with httpx.AsyncClient( timeout=effective_timeout, follow_redirects=True @@ -210,17 +216,19 @@ async def call( response = await client.post(url, json=params) else: response = await client.get(url, params=params) - return self._handle_response(response, url, language) + result = self._handle_response(response, url, language) + _duration_ms = round((time.time() - _t0) * 1000, 1) + if result.success: + logger.info( + f"APICaller: api call success | event_type=api_call_success url={url} method={method_upper} status_code={result.status_code} duration_ms={_duration_ms}" + ) + return result except httpx.TimeoutException as exc: - error_id = generate_error_id() - log_error_with_context( - logger, - error_id, - "api_call_timeout", - None, - exc, - {"url": url, "method": method_upper}, + _duration_ms = round((time.time() - _t0) * 1000, 1) + _error_id = generate_error_id() + logger.error( + f"APICaller: api call timeout | event_type=api_call_timeout url={url} method={method_upper} error_id={_error_id} duration_ms={_duration_ms} exc={exc!r}" ) self._circuit_breaker.record_failure(url) return APICallResult( @@ -231,14 +239,10 @@ async def call( ) except httpx.RequestError as exc: - error_id = generate_error_id() - log_error_with_context( - logger, - error_id, - "api_call_network_error", - None, - exc, - {"url": url, "method": method_upper}, + _duration_ms = round((time.time() - _t0) * 1000, 1) + _error_id = generate_error_id() + logger.error( + f"APICaller: api call network error | event_type=api_call_network_error url={url} method={method_upper} error_id={_error_id} duration_ms={_duration_ms} exc={exc!r}" ) self._circuit_breaker.record_failure(url) return APICallResult( @@ -271,8 +275,7 @@ def _handle_response( # Not a server fault — do NOT trip the circuit breaker. location = response.headers.get("location", "") logger.warning( - f"[APICaller] Unresolved redirect {status_code} from {url!r} " - f"→ {location!r}" + f"APICaller: api call redirect | event_type=api_call_redirect url={url} status_code={status_code} location={location!r}" ) base_msg = get_localized_message(REDIRECT_NOT_FOLLOWED_MESSAGES, language) error_msg = base_msg.format( @@ -293,7 +296,7 @@ def _handle_response( error_body = self._parse_response_body(response) raw_msg = error_body if isinstance(error_body, str) else str(error_body) logger.warning( - f"[APICaller] 4xx response {status_code} from {url!r}: {raw_msg[:200]}" + f"APICaller: api call client error | event_type=api_call_client_error url={url} status_code={status_code} error_preview={raw_msg[:200]!r}" ) return APICallResult( success=False, @@ -303,9 +306,8 @@ def _handle_response( ) # 5xx — server is misbehaving; trip the circuit breaker. - error_id = generate_error_id() logger.error( - f"[{error_id}] [APICaller] Server error {status_code} from {url!r}" + f"APICaller: api call server error | event_type=api_call_server_error url={url} status_code={status_code} error_id={generate_error_id()}" ) self._circuit_breaker.record_failure(url) return APICallResult( diff --git a/src/tool_classifier/api_response_formatter.py b/src/tool_classifier/api_response_formatter.py index 4262b41a..bab321c6 100644 --- a/src/tool_classifier/api_response_formatter.py +++ b/src/tool_classifier/api_response_formatter.py @@ -1,17 +1,38 @@ """API response formatter using DSPy — converts raw JSON API responses to natural language.""" import json +import re from typing import Any, AsyncIterator, Dict, List, Union import dspy import dspy.streaming from dspy.streaming import StreamListener -from loguru import logger - +from langfuse import observe +from src.utils.observation_utils import ( + safe_observation_context, + update_observation_safe, +) +from src.loki_logger import LokiLogger from llm_orchestrator_config.llm_ochestrator_constants import get_localized_message +from src.utils.cost_utils import get_lm_usage_since + +logger = LokiLogger(service_name="api-tool-calling") _MAX_ITEMS: int = 500 -_MAX_RESPONSE_BYTES: int = 50_000 +_MAX_RESPONSE_BYTES: int = 150_000 + + +def _get_current_model_name() -> str: + """Best-effort model name lookup from current DSPy LM.""" + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "model"): + model_name = lm.model + if isinstance(model_name, str) and model_name: + return model_name + except Exception: + pass + return "unknown" class APIResponseFormatterSignature(dspy.Signature): @@ -34,7 +55,7 @@ class APIResponseFormatterSignature(dspy.Signature): message that no results were found for the query. - If api_response contains an error field or error status, explain the issue to the user in a friendly, non-technical way. - - If the data contains more than 20 items, summarize the key highlights rather than + - If the data contains more than 50 items, summarize the key highlights rather than listing every item. Always mention the total count when summarizing. - Output must be clean text — no markdown headers (##), no code blocks (```), no raw JSON. The answer must be ready for direct display to the user. @@ -80,6 +101,13 @@ class APIResponseFormatterSignature(dspy.Signature): "When non-empty, follow these rules with highest priority." ) ) + query_params_context: str = dspy.InputField( + desc=( + "If non-empty, briefly acknowledge the time period or filter at the " + "start of the answer (e.g. 'For the period 2026-01-01 to 2026-12-31, ...'). " + "Empty string when no date or time filter was applied." + ) + ) formatted_answer: str = dspy.OutputField( desc=( @@ -95,6 +123,38 @@ class APIResponseFormatterSignature(dspy.Signature): _LANGUAGE_NAMES: Dict[str, str] = {"en": "English", "et": "Estonian", "ru": "Russian"} +# ISO-8601 datetime pattern with optional timezone (Z or ±HH:MM). +# Anchored at both start (^) and end ($) to avoid partial matches like "2026-01-01foo". +_DATE_VALUE_RE = re.compile( + r"^\d{4}-\d{2}-\d{2}(T\d{2}:\d{2}:\d{2}(Z|[+-]\d{2}:\d{2})?)?$" +) + + +def build_params_context(collected_params: Dict[str, Any]) -> str: + """Format date/datetime valued params into a readable context string. + + Detects values that look like ISO-8601 date strings (``YYYY-MM-DD`` or + ``YYYY-MM-DDTHH:MM:SS``) and formats them as human-readable key-value pairs. + Returns an empty string when no date-type values are found. + + Args: + collected_params: Dict of param names → values collected from the user. + + Returns: + A comma-separated string such as + ``"start date: 2026-01-01, end date: 2026-12-31"`` + or ``""`` when no date params are present. + """ + parts: List[str] = [] + for name, value in collected_params.items(): + if isinstance(value, str) and _DATE_VALUE_RE.match(value): + # Convert camelCase to spaced words, lower-case. + readable = re.sub(r"(?<=[a-z])(?=[A-Z])", " ", name).lower() + readable = readable.replace("_", " ") + parts.append(f"{readable}: {value}") + return ", ".join(parts) + + _FORMATTER_ERROR_MESSAGES: Dict[str, str] = { "et": "Vastuse kuvamine ebaõnnestus. Palun proovige uuesti.", "ru": "Не удалось отобразить ответ. Пожалуйста, попробуйте ещё раз.", @@ -118,12 +178,14 @@ def __init__(self, custom_instructions: str = "") -> None: self.formatter = dspy.Predict(APIResponseFormatterSignature) self._custom_instructions = custom_instructions + @observe(name="api_response_formatting_llm", as_type="generation") def forward( self, user_query: str, api_response: Union[str, Dict[str, Any], List[Any]], endpoint_description: str, detected_language: str = "en", + collected_params: Dict[str, Any] | None = None, ) -> str: """Convert a raw API response to a natural-language answer. @@ -134,15 +196,32 @@ def forward( detected_language: ISO language code from the agentic loop session ('en', 'et', 'ru'). Defaults to 'en'. This is the authoritative language for the answer — the LLM will not infer it from the data. + collected_params: Dict of param names → values collected from the user. + Used to build a date-range acknowledgment prefix when date params + are present. Defaults to None. Returns: A clean, natural-language answer ready for display to the user. """ + + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning( + f"Failed to get LM history length for response formatting: {e}" + ) + + collected_params = collected_params or {} + try: normalized = self._normalize_response(api_response) normalized = self._annotate_empty(normalized) normalized = self._truncate_if_needed(normalized) response_language = _LANGUAGE_NAMES.get(detected_language, "English") + params_context = build_params_context(collected_params) result = self.formatter( user_query=user_query, @@ -150,13 +229,49 @@ def forward( endpoint_description=endpoint_description, response_language=response_language, custom_instructions=self._custom_instructions, + query_params_context=params_context, ) - return result.formatted_answer # type: ignore[no-any-return] + formatted_answer = result.formatted_answer + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_query": user_query, + "api_response": api_response, + "endpoint_description": endpoint_description, + "response_language": response_language, + }, + output_data={ + "formatted_answer_preview": str(formatted_answer)[:500], + }, + metadata={ + "model": _get_current_model_name(), + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": False, + }, + ) + return formatted_answer # type: ignore[no-any-return] except Exception as e: logger.error( f"APIResponseFormatterModule.forward failed: {e}", exc_info=True ) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_query": user_query, + "endpoint_description": endpoint_description, + "detected_language": detected_language, + "api_response": api_response, + }, + output_data={"error": str(e)}, + metadata={ + "model": _get_current_model_name(), + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": False, + }, + ) safe_language = ( detected_language if detected_language in _FORMATTER_ERROR_MESSAGES @@ -190,6 +305,7 @@ async def stream_forward( api_response: Union[str, Dict[str, Any], List[Any]], endpoint_description: str, detected_language: str = "en", + collected_params: Dict[str, Any] | None = None, ) -> AsyncIterator[str]: """Stream formatted_answer tokens using DSPy native streaming. Yields individual token strings as they arrive from the LLM. @@ -205,83 +321,207 @@ async def stream_forward( api_response: Raw API response (dict, list, or string). endpoint_description: Short description of what the endpoint does. detected_language: ISO code ('en', 'et', 'ru'). Defaults to 'en'. + collected_params: Dict of param names → values collected from the user. + Used to build a date-range acknowledgment prefix. Defaults to None. Yields: Token strings from the LLM ``formatted_answer`` field. """ + collected_params = collected_params or {} safe_language = ( detected_language if detected_language in _FORMATTER_ERROR_MESSAGES else "en" ) output_stream = None - try: - normalized = self._normalize_response(api_response) - normalized = self._annotate_empty(normalized) - normalized = self._truncate_if_needed(normalized) - response_language = _LANGUAGE_NAMES.get(detected_language, "English") - - stream_predictor = self._get_stream_predictor() - output_stream = stream_predictor( - user_query=user_query, - api_response=normalized, - endpoint_description=endpoint_description, - response_language=response_language, - custom_instructions=self._custom_instructions, - ) - stream_started = False - token_count = 0 - async for chunk in output_stream: - if isinstance(chunk, dspy.streaming.StreamResponse): - if chunk.signature_field_name == "formatted_answer": - stream_started = True - token_count += 1 - yield chunk.chunk - elif isinstance(chunk, dspy.Prediction): - # dspy.streamify did not stream individual tokens — yield the - # full answer from the final Prediction as a single frame. - if not stream_started: - answer = getattr(chunk, "formatted_answer", None) - if answer: - logger.info( - "APIResponseFormatterModule.stream_forward: " - "no StreamResponse tokens — yielding full Prediction answer" - ) - stream_started = True - yield answer - - if stream_started and token_count > 0: - logger.debug( - f"APIResponseFormatterModule.stream_forward: streamed {token_count} tokens" - ) - - if not stream_started: - # Last-resort fallback: blocking forward() — covers cases where - # dspy.streamify yields neither StreamResponse nor Prediction. + with safe_observation_context( + as_type="generation", + name="api_response_formatting_streaming", + input={ + "user_query": user_query[:500], + "endpoint_description": endpoint_description, + "detected_language": detected_language, + }, + ) as generation: + output_stream = None + + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: logger.warning( - "APIResponseFormatterModule.stream_forward: " - "streamify produced no tokens and no Prediction — using blocking forward()" + "Failed to get LM history length for response formatting streaming: " + f"{e}" ) - result = self.forward( + try: + normalized = self._normalize_response(api_response) + normalized = self._annotate_empty(normalized) + normalized = self._truncate_if_needed(normalized) + response_language = _LANGUAGE_NAMES.get(detected_language, "English") + params_context = build_params_context(collected_params) + + stream_predictor = self._get_stream_predictor() + output_stream = stream_predictor( user_query=user_query, - api_response=api_response, + api_response=normalized, endpoint_description=endpoint_description, - detected_language=detected_language, + response_language=response_language, + custom_instructions=self._custom_instructions, + query_params_context=params_context, ) - yield result - except Exception as e: - logger.error( - f"APIResponseFormatterModule.stream_forward failed: {e}", exc_info=True - ) - yield get_localized_message(_FORMATTER_ERROR_MESSAGES, safe_language) - finally: - if output_stream is not None: + stream_started = False + token_count = 0 + accumulated: list[str] = [] + final_prediction: dspy.Prediction | None = None + async for chunk in output_stream: + if isinstance(chunk, dspy.streaming.StreamResponse): + if chunk.signature_field_name == "formatted_answer": + stream_started = True + token_count += 1 + accumulated.append(chunk.chunk) + yield chunk.chunk + elif isinstance(chunk, dspy.Prediction): + final_prediction = chunk + # dspy.streamify did not stream individual tokens — yield the + # full answer from the final Prediction as a single frame. + if not stream_started: + answer = getattr(chunk, "formatted_answer", None) + if answer: + logger.info( + "APIResponseFormatterModule.stream_forward: " + "no StreamResponse tokens — yielding full Prediction answer" + ) + stream_started = True + accumulated.append(answer) + yield answer + + assembled_answer = "".join(accumulated) + + if stream_started and token_count > 0: + logger.debug( + f"APIResponseFormatterModule.stream_forward: streamed {token_count} tokens" + ) + # DSPy streaming can drop the last few tokens before EOS. + # The final dspy.Prediction holds the authoritative complete answer. + # Yield any tail that wasn't delivered as StreamResponse chunks. + if final_prediction is not None: + full_answer = getattr( + final_prediction, "formatted_answer", None + ) + if full_answer: + streamed_text = assembled_answer + if full_answer.startswith(streamed_text) and len( + full_answer + ) > len(streamed_text): + tail = full_answer[len(streamed_text) :] + if tail.strip(): + logger.debug( + f"APIResponseFormatterModule.stream_forward: " + f"yielding {len(tail)} missing tail chars from Prediction" + ) + assembled_answer += tail + yield tail + elif streamed_text and not full_answer.startswith( + streamed_text + ): + logger.warning( + "APIResponseFormatterModule.stream_forward: " + "streamed output is not a prefix of final Prediction; " + "skipping tail reconciliation" + ) + + if not stream_started: + # Last-resort fallback: blocking forward() — covers cases where + # dspy.streamify yields neither StreamResponse nor Prediction. + logger.warning( + "APIResponseFormatterModule.stream_forward: " + "streamify produced no tokens and no Prediction — using blocking forward()" + ) + result = self.forward( + user_query=user_query, + api_response=api_response, + endpoint_description=endpoint_description, + detected_language=detected_language, + collected_params=collected_params, + ) + assembled_answer = result + yield result + + usage = get_lm_usage_since(history_length_before) try: - await output_stream.aclose() - except Exception as cleanup_error: - logger.debug(f"Error during stream cleanup: {cleanup_error}") + generation.update( + input={ + "user_query": user_query, + "api_response": api_response, + "endpoint_description": endpoint_description, + "detected_language": detected_language, + }, + output=assembled_answer, + metadata={ + "stream_started": stream_started, + "chunk_count": token_count, + "num_calls": usage.get("num_calls", 0), + "streaming": True, + }, + usage_details={ + "input": usage.get("total_prompt_tokens", 0), + "output": usage.get("total_completion_tokens", 0), + "total": usage.get("total_tokens", 0), + }, + cost_details={ + "total": usage.get("total_cost", 0.0), + }, + ) + except Exception as update_error: + logger.debug( + "Langfuse generation update skipped for response formatting " + f"streaming: {update_error}" + ) + + except Exception as e: + logger.error( + f"APIResponseFormatterModule.stream_forward failed: {e}", + exc_info=True, + ) + usage = get_lm_usage_since(history_length_before) + try: + generation.update( + input={ + "user_query": user_query, + "api_response": api_response, + "endpoint_description": endpoint_description, + "detected_language": detected_language, + }, + output={"error": str(e)}, + usage_details={ + "input": usage.get("total_prompt_tokens", 0), + "output": usage.get("total_completion_tokens", 0), + "total": usage.get("total_tokens", 0), + }, + cost_details={ + "total": usage.get("total_cost", 0.0), + }, + metadata={ + "num_calls": usage.get("num_calls", 0), + "streaming": True, + }, + ) + except Exception as update_error: + logger.debug( + "Langfuse error update skipped for response formatting " + f"streaming: {update_error}" + ) + yield get_localized_message(_FORMATTER_ERROR_MESSAGES, safe_language) + finally: + if output_stream is not None: + try: + await output_stream.aclose() + except Exception as cleanup_error: + logger.debug(f"Error during stream cleanup: {cleanup_error}") # ------------------------------------------------------------------ @@ -324,9 +564,14 @@ def _truncate_if_needed(api_response_str: str) -> str: encoded = api_response_str.encode("utf-8") if len(encoded) > _MAX_RESPONSE_BYTES: - truncated_str = encoded[:_MAX_RESPONSE_BYTES].decode( + suffix = "\n[NOTE: Response truncated due to size limit]" + suffix_bytes = len(suffix.encode("utf-8")) + truncated_str = encoded[: _MAX_RESPONSE_BYTES - suffix_bytes].decode( "utf-8", errors="ignore" ) - return truncated_str + "\n[NOTE: Response truncated due to size limit]" + logger.warning( + f"API response exceeded byte limit: {len(encoded)} bytes — truncating to {len(truncated_str.encode('utf-8'))} bytes" + ) + return truncated_str + suffix return api_response_str diff --git a/src/tool_classifier/api_semantic_searcher.py b/src/tool_classifier/api_semantic_searcher.py index 7261f3c5..efe3ee84 100644 --- a/src/tool_classifier/api_semantic_searcher.py +++ b/src/tool_classifier/api_semantic_searcher.py @@ -2,11 +2,13 @@ import asyncio import json +from dataclasses import dataclass from typing import Any, Dict, List, Optional, Protocol, cast import dspy import httpx -from loguru import logger +from src.utils.observation_utils import safe_observation_context +from src.loki_logger import LokiLogger from tool_classifier.constants import ( API_TOOL_COLLECTION, @@ -20,6 +22,23 @@ ) from tool_classifier.sparse_encoder import compute_sparse_vector from tool_classifier.sparse_encoder import SparseVector +from src.utils.cost_utils import get_lm_usage_since + + +def _get_current_model_name() -> str: + """Best-effort model name lookup from current DSPy LM.""" + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "model"): + model_name = lm.model + if isinstance(model_name, str) and model_name: + return model_name + except Exception: + pass + return "unknown" + + +logger = LokiLogger(service_name="api-tool-calling") class EmbeddingServiceProtocol(Protocol): @@ -48,6 +67,8 @@ def __init__( cosine_score: float, rrf_score: float, confidence: str, + llm_validated: bool = False, + multi_intent_hint: bool = False, ) -> None: self.endpoint_id = endpoint_id self.name = name @@ -60,6 +81,13 @@ def __init__( ) self.rrf_score = rrf_score # Hybrid RRF fusion score (used for ranking) self.confidence = confidence # "high", "medium", "none" + self.llm_validated = ( + llm_validated # True when disambiguator confirmed this match + ) + self.multi_intent_hint = ( + multi_intent_hint # True when disambiguator rejected all multi-candidates; + # this result is only valid when MULTI_INTENT_ENABLED=True + ) def to_dict(self) -> Dict[str, Any]: return { @@ -72,6 +100,7 @@ def to_dict(self) -> Dict[str, Any]: "cosine_score": round(self.cosine_score, 4), "rrf_score": round(self.rrf_score, 6), "confidence": self.confidence, + "llm_validated": self.llm_validated, } @@ -98,6 +127,16 @@ class EndpointDisambiguationSignature(dspy.Signature): ) +@dataclass +class DisambiguationResult: + """Wrapper for disambiguation result to satisfy DSPy's Langfuse callback.""" + + winner_id: Optional[str] + + def set_lm_usage(self, *args: object, **kwargs: object) -> None: + """No-op stub for DSPy's internal Langfuse callback compatibility.""" + + class EndpointDisambiguatorModule(dspy.Module): """DSPy Module for resolving ambiguous API endpoint candidates via LLM. @@ -114,7 +153,7 @@ def forward( self, user_query: str, candidates: List[Dict[str, Any]], - ) -> Optional[str]: + ) -> DisambiguationResult: """Pick the best matching endpoint_id from candidates, or return None. Args: @@ -123,7 +162,7 @@ def forward( description, and cosine_score. Returns: - The winning endpoint_id string, or None if no endpoint clearly fits. + DisambiguationResult with winner_id (endpoint_id string or None). """ candidates_payload = [ { @@ -143,14 +182,14 @@ def forward( ) winner = result.best_endpoint_id.strip() if winner.lower() == "none": - return None - return winner + return DisambiguationResult(winner_id=None) + return DisambiguationResult(winner_id=winner) except Exception as e: logger.error( f"EndpointDisambiguatorModule: Disambiguation failed: {e}", exc_info=True, ) - return None + return DisambiguationResult(winner_id=None) class APISemanticSearcher: @@ -253,8 +292,9 @@ async def search( ) query_embedding = precomputed_embedding else: - query_embedding = self._get_query_embedding( - query, environment, connection_id + # loop is not blocked — this allows parallel sub-query searches + query_embedding = await asyncio.to_thread( + self._get_query_embedding, query, environment, connection_id ) if query_embedding is None: logger.error("APISemanticSearcher: Failed to generate query embedding") @@ -391,8 +431,24 @@ async def search( ) winner_id = await self._disambiguate(query, medium_results) if winner_id is None: + if len(medium_results) > 1: + # Disambiguator rejected all candidates on a multi-candidate query. + # This happens when the query spans multiple independent intents and + # no single endpoint fully answers it. Return the top cosine-scored + # candidate WITHOUT llm_validated so the classifier's IntentDecomposer + # gate can run and detect the parallel path. + top = max(medium_results, key=lambda r: r.cosine_score) + logger.info( + f"APISemanticSearcher: disambiguator rejected all " + f"{len(medium_results)} candidates (possible multi-intent) — " + f"returning top candidate {top.name!r} " + f"(cosine={top.cosine_score:.4f}) for IntentDecomposer" + ) + top.multi_intent_hint = True + return [top] logger.info( - "APISemanticSearcher: disambiguator rejected all candidates — no API tool match" + "APISemanticSearcher: disambiguator rejected sole candidate — " + "no API tool match" ) return [] @@ -404,9 +460,10 @@ async def search( ) return [] + winner.llm_validated = True logger.info( f"APISemanticSearcher: disambiguated winner → {winner.name!r} " - f"(cosine={winner.cosine_score:.4f})" + f"(cosine={winner.cosine_score:.4f}, llm_validated=True)" ) return [winner] @@ -437,17 +494,49 @@ async def _disambiguate( f"APISemanticSearcher: disambiguating {len(candidates)} candidates " f"for query: {query!r}" ) - # Run the synchronous DSPy LLM call in a thread pool so it does not - # block the asyncio event loop while waiting for the LLM response. - # cast: asyncio.to_thread infers Prediction from DSPy; forward() returns Optional[str] - winner_id = cast( - Optional[str], - await asyncio.to_thread( - self._disambiguator, - user_query=query, - candidates=candidate_dicts, - ), - ) + + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning(f"Failed to get LM history length for disambiguation: {e}") + + with safe_observation_context( + name="api_endpoint_disambiguation_llm", + as_type="generation", + input={"user_query": query, "candidates_count": len(candidates)}, + ) as generation: + # Run the synchronous DSPy LLM call in a thread pool so it does not + # block the asyncio event loop while waiting for the LLM response. + disambiguation_result = cast( + DisambiguationResult, + await asyncio.to_thread( + self._disambiguator, + user_query=query, + candidates=candidate_dicts, + ), + ) + winner_id = disambiguation_result.winner_id + + # Update Langfuse observation with output and usage + try: + if generation is not None: + usage = get_lm_usage_since(history_length_before) + generation.update( + model=_get_current_model_name(), + output={"winner": winner_id}, + usage_details={ + "input": usage.get("total_prompt_tokens", 0), + "output": usage.get("total_completion_tokens", 0), + "total": usage.get("total_tokens", 0), + }, + cost_details={"total": usage.get("total_cost", 0.0)}, + ) + except Exception as e: + logger.debug(f"Langfuse generation update skipped: {e}") + if winner_id: logger.info( f"APISemanticSearcher: disambiguator picked endpoint_id={winner_id!r}" diff --git a/src/tool_classifier/classifier.py b/src/tool_classifier/classifier.py index 35c06571..9c28e68e 100644 --- a/src/tool_classifier/classifier.py +++ b/src/tool_classifier/classifier.py @@ -12,11 +12,12 @@ TYPE_CHECKING, ) import httpx -from loguru import logger +import asyncio + +from src.loki_logger import LokiLogger from llm_orchestrator_config.llm_manager import LLMManager from models.request_models import ( - ConversationItem, OrchestrationRequest, OrchestrationResponse, TestOrchestrationResponse, @@ -26,6 +27,7 @@ WorkflowType, WORKFLOW_DISPLAY_NAMES, WORKFLOW_LAYER_ORDER, + ExecutionMode, ) from tool_classifier.models import ClassificationResult from tool_classifier.constants import ( @@ -38,11 +40,17 @@ DENSE_MIN_THRESHOLD, DENSE_HIGH_CONFIDENCE_THRESHOLD, DENSE_SCORE_GAP_THRESHOLD, + API_TOOL_MIN_THRESHOLD, + API_TOOL_HIGH_CONFIDENCE_THRESHOLD, API_TOOL_INTENT_SWITCH_THRESHOLD, ) from tool_classifier.sparse_encoder import SparseVector, compute_sparse_vector -from tool_classifier.api_semantic_searcher import APISemanticSearcher +from tool_classifier.api_semantic_searcher import ( + APISemanticSearcher, + APIToolSearchResult, +) +from tool_classifier.intent_decomposer import IntentDecomposerModule from tool_classifier.workflows import ( APIToolWorkflowExecutor, @@ -52,6 +60,10 @@ OODWorkflowExecutor, ) from llm_orchestrator_config.feature_flags import FeatureFlags +from utils.atc_cache_store import ATCCacheStore + +# Initialize Loki logger +logger = LokiLogger(service_name="tool-classifier") if TYPE_CHECKING: from llm_orchestration_service import LLMOrchestrationService @@ -114,6 +126,9 @@ def __init__( self.context_workflow = ContextWorkflowExecutor( llm_manager=llm_manager, orchestration_service=orchestration_service, + conversation_history_store=getattr( + orchestration_service, "conversation_history_store", None + ), ) self.rag_workflow = RAGWorkflowExecutor( orchestration_service=orchestration_service, @@ -126,6 +141,9 @@ def __init__( qdrant_client=self._qdrant_client, ) + # Intent decomposer + self.intent_decomposer = IntentDecomposerModule() + logger.info( "Tool classifier initialized with hybrid search classification " f"(Qdrant: {self._qdrant_base_url})" @@ -142,7 +160,6 @@ async def aclose(self) -> None: async def classify( self, query: str, - conversation_history: List[ConversationItem], language: str, request: Optional[OrchestrationRequest] = None, ) -> ClassificationResult: @@ -160,7 +177,6 @@ async def classify( Args: query: User's query string - conversation_history: List of previous conversation messages language: Detected language code (e.g., 'en', 'et') request: Original orchestration request (needed for ATC search which requires environment and connection_id for embedding). @@ -210,6 +226,11 @@ async def classify( f"— abandoning old session" ) await session_store.delete(request.chatId) + if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED: + await ATCCacheStore().invalidate_l2(request.chatId) + logger.info( + f"[{request.chatId}] ATC cache: L2 invalidated on intent switch" + ) return new_api_match logger.info( @@ -802,25 +823,156 @@ async def _try_api_tool_classification( precomputed_embedding=precomputed_embedding, min_cosine_override=min_cosine_override, ) - if results: - matched = results[0] + if not results: + return None + + matched = results[0] + logger.info( + f"API tool match: {matched.name!r} " + f"(confidence={matched.confidence}, cosine={matched.cosine_score:.4f})" + ) + + # Suppress multi-intent hint results when the feature is disabled. + # The searcher returns a hint result (multi_intent_hint=True) when the + # disambiguator rejected all multi-candidates — it is only meaningful when + # MULTI_INTENT_ENABLED=True. + if matched.multi_intent_hint and not FeatureFlags.MULTI_INTENT_ENABLED: logger.info( - f"API tool match: {matched.name!r} " - f"(confidence={matched.confidence}, cosine={matched.cosine_score:.4f})" + f"ATC: {matched.name!r} is a multi-intent hint but " + f"MULTI_INTENT_ENABLED=False — falling through to CONTEXT/RAG" ) - return ClassificationResult( - workflow=WorkflowType.API_TOOL_CALLING, - confidence=matched.cosine_score, - metadata={"matched_endpoint": matched.to_dict()}, - reasoning=( - f"API tool match: {matched.name} " - f"(cosine={matched.cosine_score:.4f}, confidence={matched.confidence})" - ), + return None + + # ── Score-band gate: try multi-intent decomposition ─────────── + # A score between the min and high-confidence thresholds may indicate + # a diluted embedding caused by multiple intents in one query. + # Scores at or above the high-confidence threshold are normally clear + # single matches — UNLESS the disambiguator already ran and rejected + # all candidates (multi_intent_hint=True). In that case the score + # landed just above the threshold only because two close intents + # pushed the embedding upward together; IntentDecomposer must still run. + in_ambiguous_band = ( + API_TOOL_MIN_THRESHOLD + <= matched.cosine_score + < API_TOOL_HIGH_CONFIDENCE_THRESHOLD + ) + if ( + (in_ambiguous_band or matched.multi_intent_hint) + and FeatureFlags.MULTI_INTENT_ENABLED + and not matched.llm_validated + ): + reason = ( + "multi_intent_hint (disambiguator rejected all candidates)" + if matched.multi_intent_hint + else f"cosine={matched.cosine_score:.4f} in ambiguous band " + f"[{API_TOOL_MIN_THRESHOLD}, {API_TOOL_HIGH_CONFIDENCE_THRESHOLD})" ) + logger.info(f"ATC: {reason} — running IntentDecomposer") + decomposition = await self.intent_decomposer.decompose(query) + + if decomposition.mode == ExecutionMode.PARALLEL: + parallel_result = await self._try_parallel_api_tool_classification( + sub_queries=decomposition.sub_queries, + environment=environment, + connection_id=connection_id, + original_matched=matched, + ) + if parallel_result is not None: + return parallel_result + # Fewer than 2 endpoints matched — fall through to single path + + # ── Single-endpoint path (unchanged) ───────────────────────── + return ClassificationResult( + workflow=WorkflowType.API_TOOL_CALLING, + confidence=matched.cosine_score, + metadata={ + "matched_endpoint": matched.to_dict(), + "execution_mode": ExecutionMode.SINGLE, + }, + reasoning=( + f"API tool match: {matched.name} " + f"(cosine={matched.cosine_score:.4f}, confidence={matched.confidence})" + ), + ) + except Exception as e: logger.error(f"API tool classification failed: {e}", exc_info=True) return None + async def _try_parallel_api_tool_classification( + self, + sub_queries: list[str], + environment: str, + connection_id: Optional[str], + original_matched: APIToolSearchResult, + ) -> Optional[ClassificationResult]: + """Run parallel endpoint searches for each sub-query from IntentDecomposer. + + Searches all sub-queries concurrently. Returns a ClassificationResult with + execution_mode="parallel" if at least 2 distinct endpoints are matched, or + None to signal the caller to fall back to the single-endpoint path. + + Args: + sub_queries: Focused sub-queries from IntentDecomposer (2-3 items). + environment: LLM environment from the original request. + connection_id: Connection ID from the original request. + original_matched: The gate search result (used as fallback reference). + + Returns: + ClassificationResult with parallel metadata, or None if <2 matched. + """ + + async def search_one(sub_query: str) -> Optional[APIToolSearchResult]: + try: + sub_results = await self.api_tool_searcher.search( + query=sub_query, + environment=environment, + connection_id=connection_id, + ) + return sub_results[0] if sub_results else None + except Exception as exc: + logger.warning( + f"ATC parallel sub-search failed for {sub_query!r}: {exc}" + ) + return None + + sub_results = await asyncio.gather(*[search_one(q) for q in sub_queries]) + + # Deduplicate by endpoint name — keep first occurrence + seen: set[str] = set() + matched_endpoints: list[dict[str, Any]] = [] + for result in sub_results: + if result is None: + continue + name: str = result.name + if name not in seen: + seen.add(name) + matched_endpoints.append(result.to_dict()) + + if len(matched_endpoints) < 2: + logger.info( + f"ATC parallel: only {len(matched_endpoints)} distinct endpoint(s) matched " + f"— falling back to single path" + ) + return None + + logger.info( + f"ATC parallel: {len(matched_endpoints)} distinct endpoints matched: " + f"{[e.get('name') for e in matched_endpoints]}" + ) + return ClassificationResult( + workflow=WorkflowType.API_TOOL_CALLING, + confidence=original_matched.cosine_score, + metadata={ + "execution_mode": ExecutionMode.PARALLEL, + "matched_endpoints": matched_endpoints, + }, + reasoning=( + f"Multi-intent parallel match: " + f"{[e.get('name') for e in matched_endpoints]}" + ), + ) + async def _execute_with_fallback_async( self, workflow: BaseWorkflow, @@ -981,7 +1133,9 @@ async def _execute_with_fallback_streaming( f"(Layer {layer_number})" ) - result = await next_workflow.execute_streaming(request, {}, time_metric) + result = await next_workflow.execute_streaming( + request, context, time_metric + ) if result is not None: logger.info(f"[{chat_id}] {next_name} streaming started") @@ -999,7 +1153,7 @@ async def _execute_with_fallback_streaming( # Fallback to RAG on error logger.info(f"[{chat_id}] Falling back to RAG streaming due to error") streaming_result = await self.rag_workflow.execute_streaming( - request, {}, time_metric + request, context, time_metric ) if streaming_result is not None: async for chunk in streaming_result: diff --git a/src/tool_classifier/constants.py b/src/tool_classifier/constants.py index 14a744c7..23b17207 100644 --- a/src/tool_classifier/constants.py +++ b/src/tool_classifier/constants.py @@ -140,6 +140,15 @@ # Agentic Loop — Continuation Threshold # ============================================================================ +MULTI_API_MAX_TURNS = 9 +"""Hard cap on total turns for the multi-endpoint agentic loop (MultiEndpointAgenticLoop). + +The per-session limit is ``min(3 * num_endpoints, MULTI_API_MAX_TURNS)``. +``MULTI_API_MAX_ENDPOINTS`` independently caps ``num_endpoints`` at 3, so the +combined ceiling is ``min(3 * 3, 9) = 9`` turns — up to 3 turns per endpoint +at maximum endpoint count. +""" + CONTINUATION_TURN = 3 """1-based turn count (after increment) at which the loop asks the user whether to continue collecting parameters or fall back to the RAG workflow. @@ -154,6 +163,21 @@ run_turn #3 (turn 2→3): user doesn't answer properly → CONTINUATION CHECK """ +MULTI_INTENT_CONTINUATION_TURN = 4 +"""1-based turn count at which the multi-endpoint loop asks the user whether to +continue collecting parameters when **multiple** intents (endpoints) are active. +Fixed at 4 regardless of how many endpoints are active — if required params are +still missing at exactly this turn, the user is asked whether to keep going. +""" + +MULTI_INTENT_MAX_TURNS = 6 +"""Hard cap on total turns for the multi-endpoint agentic loop when multiple +intents are active. When the internal turn counter reaches or exceeds this value +(i.e., when `turn_count >= MULTI_INTENT_MAX_TURNS`, such as when attempting +`turn_count == 6`), the loop falls back to the RAG workflow regardless of how +many required params are still missing. +""" + CONTINUATION_QUESTION = ( "I still need a bit more information, but we've been at this for a while. " "Would you like to keep going and answer a few more questions " @@ -231,3 +255,46 @@ "en": "Your request could not be processed. Please check the provided information and try again.", } """Friendly message returned on 4xx client errors from external API calls.""" + + +# ============================================================================ +# Multi-Intent (Parallel Multi-API) Configuration +# ============================================================================ + +MULTI_API_MAX_ENDPOINTS = 3 +"""Maximum number of parallel endpoints allowed per multi-intent query. +Caps the number of sub-queries the IntentDecomposer may produce and the +number of concurrent API calls MultiAPICaller may execute.""" + +MULTI_API_BATCH_TIMEOUT = 30 +"""Total wall-clock timeout in seconds for all concurrent API calls in a batch. +Exceeds the per-call API_CALL_TIMEOUT (10 s) to allow parallel requests time to +complete without the batch being cancelled prematurely.""" + +MULTI_API_PARTIAL_FAILURE_MESSAGES = { + "et": "Mõned teenusekutsed ebaõnnestusid. Osad tulemused võivad puududa.", + "ru": "Некоторые запросы к сервисам завершились неудачно. Часть результатов может отсутствовать.", + "en": "Some service calls failed. Partial results may be missing.", +} +"""Friendly message returned when a batch of API calls has partial failures.""" + + +# ============================================================================ +# ATC Response Cache Configuration +# ============================================================================ + +ATC_CACHE_KEY_PREFIX = "atc:cache" +"""Redis key prefix for L1 exact-match response cache entries. +Full key format: atc:cache:{chat_id}:{api_name}:{param_hash}""" + +ATC_LAST_CALL_KEY_PREFIX = "atc:last" +"""Redis key prefix for L2 last-call context entries. +Full key format: atc:last:{chat_id}""" + +ATC_CACHE_DEFAULT_TTL_SECONDS = 1800 +"""Default L1 cache TTL in seconds (30 minutes). +Applied when an endpoint does not define a per-endpoint cache_ttl_seconds override.""" + +ATC_LAST_CALL_TTL_SECONDS = 1800 +"""L2 last-call context TTL in seconds (30 minutes). +Matches the session TTL so cached context never outlives the session.""" diff --git a/src/tool_classifier/context_analyzer.py b/src/tool_classifier/context_analyzer.py index da4eba1d..071ee36b 100644 --- a/src/tool_classifier/context_analyzer.py +++ b/src/tool_classifier/context_analyzer.py @@ -7,12 +7,32 @@ import dspy import dspy.streaming from dspy.streaming import StreamListener -from loguru import logger +from langfuse import observe +from src.utils.observation_utils import ( + safe_observation_context, + update_observation_safe, +) +from src.loki_logger import LokiLogger from pydantic import BaseModel, Field from src.utils.cost_utils import get_lm_usage_since from tool_classifier.greeting_constants import get_greeting_response +logger = LokiLogger(service_name="context-workflow") + + +def _get_current_model_name() -> str: + """Best-effort model name lookup from current DSPy LM.""" + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "model"): + model_name = lm.model + if isinstance(model_name, str) and model_name: + return model_name + except Exception: + pass + return "unknown" + class ContextAnalysisResult(BaseModel): """Result of context analysis.""" @@ -87,6 +107,42 @@ class ConversationSummarySignature(dspy.Signature): ) +class IncrementalSummarySignature(dspy.Signature): + """Merge newly evicted conversation rounds into an existing summary. + + Given an existing summary (which may be empty for the first eviction) and a + JSON-formatted list of conversation rounds that were just evicted from the + active history window, produce an updated summary that incorporates all new + information. + + Guidelines: + - Preserve all factual details from the existing summary (names, numbers, dates). + - Integrate only new, non-redundant information from the evicted rounds. + - Keep the summary concise — omit filler, focus on actionable/memorable facts. + - Respond in the SAME language as the conversation. + """ + + existing_summary: str = dspy.InputField( + desc=( + "Current conversation summary. May be an empty string if no summary " + "exists yet (first eviction)." + ) + ) + new_rounds: str = dspy.InputField( + desc=( + "JSON array of conversation rounds that were just evicted from the " + "active history window, each with user_message, bot_message, and timestamp." + ) + ) + updated_summary: str = dspy.OutputField( + desc=( + "Updated summary that merges the existing summary with the new rounds. " + "Preserve all factual details; drop redundant information. " + "Same language as the conversation." + ) + ) + + class SummaryAnalysisSignature(dspy.Signature): """Analyze if a user query can be answered from a conversation summary. @@ -281,6 +337,7 @@ def _merge_cost_dicts( "num_calls": cost1.get("num_calls", 0) + cost2.get("num_calls", 0), } + @observe(name="context_detection_phase1", as_type="generation") async def detect_context( self, query: str, @@ -365,24 +422,45 @@ async def detect_context( ) cost_dict = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "query": query, + "history_turns": total_turns, + }, + output_data={ + "is_greeting": result.is_greeting, + "greeting_type": result.greeting_type, + "can_answer_from_context": result.can_answer_from_context, + "has_context_snippet": bool(result.context_snippet), + }, + metadata={ + "model": _get_current_model_name(), + "usage": cost_dict, + }, + ) logger.info( f"Detection cost | Total: ${cost_dict.get('total_cost', 0):.6f} | " f"Tokens: {cost_dict.get('total_tokens', 0)}" ) return result, cost_dict + @observe(name="context_detection_with_summary_fallback", as_type="chain") async def detect_context_with_summary_fallback( self, query: str, conversation_history: List[Dict[str, Any]], + pre_computed_summary: Optional[str] = None, ) -> tuple[ContextDetectionResult, Dict[str, Any]]: """ Phase 1 with summary fallback: detect if query can be answered from history. Implements a 3-step flow: 1. Check the last 10 turns via detect_context(). - 2. If cannot answer AND total history > 10 turns: - - Generate a concise summary of the older turns (everything before the last 10). + 2. If cannot answer AND (total history > 10 turns OR a pre-computed summary + is available from Redis): + - Use *pre_computed_summary* directly when provided (skips the expensive + LLM summarisation call). + - Otherwise generate a concise summary of the older turns. - Check whether the query can be answered from that summary. 3. If still cannot answer, return can_answer=False (workflow falls back to RAG). @@ -396,6 +474,11 @@ async def detect_context_with_summary_fallback( Args: query: User query to classify conversation_history: Full conversation history + pre_computed_summary: Running conversation summary retrieved from Redis. + When provided the LLM summary-generation step is skipped and this + value is used directly. The summary path is also attempted even + when ``total_turns <= 10`` because evicted rounds may exist in Redis + beyond what is currently in *conversation_history*. Returns: Tuple of (ContextDetectionResult, cost_dict) @@ -411,19 +494,35 @@ async def detect_context_with_summary_fallback( if result.is_greeting or result.can_answer_from_context: return result, cost_dict - # Step 2 & 3: if history exceeds 10 turns, try summary-based detection - if total_turns > 10: - logger.info( - f"History has {total_turns} turns (> 10) | " - f"Cannot answer from recent 10 | Attempting summary-based detection" - ) - older_history = conversation_history[:-10] - logger.info(f"Summarizing {len(older_history)} older turns") - + # Step 2 & 3: try summary-based detection when history is long *or* Redis + # has a pre-computed summary (which may cover evicted rounds beyond what is + # currently in memory). + if total_turns > 10 or pre_computed_summary is not None: try: - summary, summary_cost = await self._generate_conversation_summary( - older_history - ) + if pre_computed_summary is not None: + # Redis path: reuse the incremental summary, skip LLM generation. + logger.info( + "Pre-computed summary available | " + "Skipping LLM summary generation, using Redis summary directly" + ) + summary = pre_computed_summary + summary_cost: Dict[str, Any] = { + "total_cost": 0.0, + "total_tokens": 0, + "num_calls": 0, + } + else: + # On-demand path: summarise older turns via LLM. + logger.info( + f"History has {total_turns} turns (> 10) | " + f"Cannot answer from recent 10 | Attempting summary-based detection" + ) + older_history = conversation_history[:-10] + logger.info(f"Summarizing {len(older_history)} older turns") + summary, summary_cost = await self._generate_conversation_summary( + older_history + ) + cost_dict = self._merge_cost_dicts(cost_dict, summary_cost) if summary: @@ -501,79 +600,136 @@ async def stream_context_response( Yields: Token strings as they arrive from the LLM (or simulated chunks) """ - logger.info(f"CONTEXT GENERATOR: Phase 2 streaming | Query: '{query[:100]}'") + with safe_observation_context( + as_type="generation", + name="context_response_streaming", + input={"query": query[:500]}, + ) as _generation: + logger.info( + f"CONTEXT GENERATOR: Phase 2 streaming | Query: '{query[:100]}'" + ) - self.llm_manager.ensure_global_config() - output_stream = None - stream_started = False - prediction_answer: Optional[str] = None - try: - with self.llm_manager.use_task_local(): - # Always create a fresh StreamListener + streamified predictor so that - # the listener's internal state is clean for this call. - answer_listener = StreamListener(signature_field_name="answer") - stream_predictor: Any = dspy.streamify( - dspy.Predict(ContextResponseGenerationSignature), - stream_listeners=[answer_listener], - ) - output_stream = stream_predictor( - context_snippet=context_snippet, - user_query=query, + self.llm_manager.ensure_global_config() + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning( + f"Failed to get LM history length for streaming generation: {e}" ) + output_stream = None + stream_started = False + prediction_answer: Optional[str] = None + assembled_answer = "" # Collect all yielded tokens here + try: + with self.llm_manager.use_task_local(): + # Always create a fresh StreamListener + streamified predictor so that + # the listener's internal state is clean for this call. + answer_listener = StreamListener(signature_field_name="answer") + stream_predictor: Any = dspy.streamify( + dspy.Predict(ContextResponseGenerationSignature), + stream_listeners=[answer_listener], + ) + output_stream = stream_predictor( + context_snippet=context_snippet, + user_query=query, + ) - async for chunk in output_stream: - if isinstance(chunk, dspy.streaming.StreamResponse): - if chunk.signature_field_name == "answer": - stream_started = True - yield chunk.chunk - elif isinstance(chunk, dspy.Prediction): - logger.info( - "Context response streaming complete (final Prediction received)" + async for chunk in output_stream: + if isinstance(chunk, dspy.streaming.StreamResponse): + if chunk.signature_field_name == "answer": + stream_started = True + assembled_answer += chunk.chunk + yield chunk.chunk + elif isinstance(chunk, dspy.Prediction): + logger.info( + "Context response streaming complete (final Prediction received)" + ) + if not stream_started: + # Tokens didn't stream — extract answer from the Prediction + # directly as first fallback before leaving the LM context. + prediction_answer = getattr(chunk, "answer", "") or "" + + except GeneratorExit: + raise + except Exception as e: + logger.error(f"Error during context response streaming: {e}") + raise + finally: + if output_stream is not None: + try: + await output_stream.aclose() + except Exception as cleanup_error: + logger.debug( + f"Error during context stream cleanup: {cleanup_error}" ) - if not stream_started: - # Tokens didn't stream — extract answer from the Prediction - # directly as first fallback before leaving the LM context. - prediction_answer = getattr(chunk, "answer", "") or "" - except GeneratorExit: - raise - except Exception as e: - logger.error(f"Error during context response streaming: {e}") - raise - finally: - if output_stream is not None: - try: - await output_stream.aclose() - except Exception as cleanup_error: - logger.debug( - f"Error during context stream cleanup: {cleanup_error}" - ) + stream_usage = get_lm_usage_since(history_length_before) - if stream_started: - return + # Determine final answer for Langfuse output + final_answer = assembled_answer + if not stream_started and prediction_answer: + final_answer = prediction_answer - # Fallback 1: answer was in the final Prediction but didn't stream as tokens - if prediction_answer: + try: + _generation.update( + model=_get_current_model_name(), + usage_details={ + "input": stream_usage.get("total_prompt_tokens", 0), + "output": stream_usage.get("total_completion_tokens", 0), + "total": stream_usage.get("total_tokens", 0), + }, + cost_details={ + "total": stream_usage.get("total_cost", 0.0), + }, + input={ + "query": query, + "context_snippet_preview": context_snippet[:500], + }, + output={ + "answer": final_answer, + "stream_started": stream_started, + "has_prediction_answer": bool(prediction_answer), + }, + metadata={ + "num_calls": stream_usage.get("num_calls", 0), + }, + ) + except Exception as e: + logger.warning( + f"Failed to update context streaming observation in Langfuse: {e}" + ) + + if stream_started: + return + + # Fallback 1: answer was in the final Prediction but didn't stream as tokens + if prediction_answer: + logger.warning( + "Stream tokens not received — yielding answer from final Prediction in chunks." + ) + for text_chunk in self._yield_in_chunks(prediction_answer): + yield text_chunk + return + + # Fallback 2: Prediction had no answer either — call generate_context_response logger.warning( - "Stream tokens not received — yielding answer from final Prediction in chunks." + "No answer from streamify — falling back to generate_context_response." ) - for text_chunk in self._yield_in_chunks(prediction_answer): - yield text_chunk - return - - # Fallback 2: Prediction had no answer either — call generate_context_response - logger.warning( - "No answer from streamify — falling back to generate_context_response." - ) - fallback_answer, _ = await self.generate_context_response( - query=query, context_snippet=context_snippet - ) - if fallback_answer: - for text_chunk in self._yield_in_chunks(fallback_answer): - yield text_chunk - else: - logger.error("All Phase 2 streaming fallbacks exhausted — empty response.") + fallback_answer, _ = await self.generate_context_response( + query=query, context_snippet=context_snippet + ) + if fallback_answer: + for text_chunk in self._yield_in_chunks(fallback_answer): + yield text_chunk + else: + logger.error( + "All Phase 2 streaming fallbacks exhausted — empty response." + ) + @observe(name="context_response_non_streaming", as_type="generation") async def generate_context_response( self, query: str, @@ -624,12 +780,27 @@ async def generate_context_response( logger.error(f"Context response generation failed: {e}", exc_info=True) cost_dict = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "query": query, + "context_snippet_preview": context_snippet[:500], + }, + output_data={ + "answer_preview": answer[:500], + "has_answer": bool(answer), + }, + metadata={ + "model": _get_current_model_name(), + "usage": cost_dict, + }, + ) logger.info( f"Generation cost | Total: ${cost_dict.get('total_cost', 0):.6f} | " f"Tokens: {cost_dict.get('total_tokens', 0)}" ) return answer, cost_dict + @observe(name="conversation_summary_generation", as_type="generation") async def _generate_conversation_summary( self, older_history: List[Dict[str, Any]], @@ -680,6 +851,19 @@ async def _generate_conversation_summary( summary = "" cost_dict = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "older_history_turns": len(older_history), + }, + output_data={ + "summary_preview": summary[:500], + "summary_length": len(summary), + }, + metadata={ + "model": _get_current_model_name(), + "usage": cost_dict, + }, + ) logger.info( f"Summary cost | Total: ${cost_dict.get('total_cost', 0):.6f} | " f"Tokens: {cost_dict.get('total_tokens', 0)}" @@ -687,6 +871,7 @@ async def _generate_conversation_summary( return summary, cost_dict + @observe(name="summary_answerability_analysis", as_type="generation") async def _analyze_from_summary( self, query: str, @@ -775,6 +960,21 @@ async def _analyze_from_summary( ) cost_dict = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "query": query, + "summary_preview": summary[:500], + }, + output_data={ + "can_answer_from_context": result.can_answer_from_context, + "answered_from_summary": result.answered_from_summary, + "has_answer": bool(result.answer), + }, + metadata={ + "model": _get_current_model_name(), + "usage": cost_dict, + }, + ) logger.info( f"Summary analysis cost | Total: ${cost_dict.get('total_cost', 0):.6f} | " f"Tokens: {cost_dict.get('total_tokens', 0)}" @@ -782,6 +982,7 @@ async def _analyze_from_summary( return result, cost_dict + @observe(name="context_analysis_orchestration", as_type="chain") async def analyze_context( self, query: str, @@ -886,6 +1087,7 @@ async def analyze_context( ) return result, cost_dict + @observe(name="context_analysis_recent_history", as_type="generation") async def _analyze_recent_history( self, query: str, @@ -1000,6 +1202,22 @@ async def _analyze_recent_history( # Calculate costs cost_dict = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "query": query, + "history_turns": len(conversation_history), + "language": language, + }, + output_data={ + "is_greeting": result.is_greeting, + "can_answer_from_context": result.can_answer_from_context, + "has_answer": bool(result.answer), + }, + metadata={ + "model": _get_current_model_name(), + "usage": cost_dict, + }, + ) logger.info( f"Cost tracking | Total cost: ${cost_dict.get('total_cost', 0):.6f} | " f"Tokens: {cost_dict.get('total_tokens', 0)} | " diff --git a/src/tool_classifier/continuation_utils.py b/src/tool_classifier/continuation_utils.py new file mode 100644 index 00000000..fd155f19 --- /dev/null +++ b/src/tool_classifier/continuation_utils.py @@ -0,0 +1,42 @@ +"""Shared utilities for agentic loop continuation detection.""" + +_YES_RESPONSES = frozenset( + { + "yes", + "y", + "jah", + "ja", + "да", + "ok", + "okay", + "sure", + "please", + "continue", + "jätka", + "продолжить", + "absolutely", + } +) +"""Normalised affirmative responses recognised across Estonian, English, and Russian. + +Any response that is not in this set is treated as a "no" so the loop falls +back to the RAG workflow. +""" + + +def detect_continuation_response(user_message: str) -> bool: + """Detect whether the user's message indicates they want to continue. + + Checks the normalised (lower-cased, stripped) message against a set of + known affirmative responses in Estonian, English, and Russian. + Any response that is not clearly affirmative is treated as a "no" so + the loop falls back to the RAG workflow. + + Args: + user_message: The raw user message to inspect. + + Returns: + True if the user wants to continue, False otherwise. + """ + normalised = user_message.strip().lower() + return normalised in _YES_RESPONSES diff --git a/src/tool_classifier/enums.py b/src/tool_classifier/enums.py index f734bc7d..6694f8a9 100644 --- a/src/tool_classifier/enums.py +++ b/src/tool_classifier/enums.py @@ -61,3 +61,15 @@ class AgenticLoopStatus(str, Enum): NEEDS_INPUT = "needs_input" MAX_TURNS_REACHED = "max_turns_reached" AWAITING_CONTINUATION_DECISION = "awaiting_continuation_decision" + + +class ExecutionMode(str, Enum): + """Execution mode for the API Tool Calling workflow. + + - SINGLE: One endpoint matched — existing single-API path, no changes. + - PARALLEL: Multiple independent endpoints matched via intent decomposition — + params collected together, APIs called concurrently. + """ + + SINGLE = "single" + PARALLEL = "parallel" diff --git a/src/tool_classifier/follow_up_detector.py b/src/tool_classifier/follow_up_detector.py new file mode 100644 index 00000000..cbd8a403 --- /dev/null +++ b/src/tool_classifier/follow_up_detector.py @@ -0,0 +1,228 @@ +"""Follow-up detection using DSPy — classifies whether a user query is a follow-up to a previous ATC API call.""" + +import json +import time +from typing import Any, Dict, List, TypedDict + +import dspy +from src.loki_logger import LokiLogger +from src.utils.error_utils import generate_error_id + +from .param_extractor import strip_format_hints + +logger = LokiLogger(service_name="api-tool-calling") + +_VALID_FOLLOW_UP_TYPES = {"param_update", "response_question", "new_intent"} + + +def _validate_updated_params( + updated_params: Dict[str, Any], + params_schema: List[Dict[str, Any]], +) -> Dict[str, Any]: + """Validate and filter updated_params against the schema. + + Removes unexpected keys not present in the schema and logs warnings for them. + This prevents untrusted LLM-generated parameters from leaking downstream. + + Args: + updated_params: Parameter dict from LLM (untrusted) + params_schema: List of parameter schema dicts with name, type, etc. + + Returns: + Filtered dict containing only schema-valid parameter keys + """ + # Build lookup: param_name -> param_schema + schema_lookup = { + p["name"]: p for p in params_schema if isinstance(p, dict) and "name" in p + } + + validated = {} + for key, value in updated_params.items(): + if key not in schema_lookup: + logger.warning( + f"FollowUpDetectorModule: unexpected param dropped | event_type=unexpected_param_dropped param_name={key!r}" + ) + else: + validated[key] = value + + return validated + + +class FollowUpDetectionResult(TypedDict): + """Return contract for FollowUpDetectorModule.forward().""" + + follow_up_type: str + updated_params: Dict[str, Any] + + +class FollowUpDetectionSignature(dspy.Signature): + """Classify whether a user query is a follow-up to a previous API call. + + Determine the relationship between the current user query and the previous query + that triggered an ATC (API Tool Call) API call. + + CLASSIFICATION RULES: + + param_update — The user is modifying or refining the SAME intent with different parameter values. + Examples: + - Previous: "Show public holidays in 2026" → Current: "What about 2025 instead?" → param_update + (same holiday lookup intent, only the year changed) + - Previous: "Convert EUR to USD" → Current: "What about GBP to JPY?" → param_update + (same currency conversion intent, only currencies changed) + - Previous: "Weather in Tallinn" → Current: "And in Tartu?" → param_update + (same weather intent, only city changed) + + response_question — The user is asking a question ABOUT the data/results already returned. + Examples: + - Previous: "List public holidays in 2026" → Response showed 12 holidays → Current: "Which of those falls on a Monday?" → response_question + (asking about the already-returned data, no new API call needed) + - Previous: "Show electricity prices" → Response showed prices → Current: "Why is the price so high in January?" → response_question + (analysing the returned results, not requesting new data) + - Previous: "Get train schedule from Tallinn to Tartu" → Response showed times → Current: "How long does the first journey take?" → response_question + (asking about details within the returned results) + + new_intent — The user is starting a completely different request, unrelated to the previous API call. + Examples: + - Previous: "Show public holidays in 2026" → Current: "How do I apply for a driving licence?" → new_intent + (completely different topic) + - Previous: "Convert EUR to USD" → Current: "What is the weather in Riga?" → new_intent + (unrelated request) + - Previous: "Get train schedules" → Current: "Tell me about Estonian history" → new_intent + (no connection to the previous API call) + + CRITICAL LANGUAGE RULE: + - Understand Estonian, Russian, and English queries equally + - Short follow-ups like "ja 2025?" (Estonian: "and 2025?") or "а в 2025?" (Russian: "and in 2025?") are param_update + - Language of the query does NOT affect classification — only intent relationship matters + + OUTPUT RULES: + - follow_up_type MUST be exactly one of: "param_update", "response_question", "new_intent" + - updated_params MUST be a valid JSON object (can be empty {}) + - For param_update: populate updated_params with the changed parameter values extracted from user_query + - For response_question or new_intent: return updated_params as empty {} + """ + + user_query: str = dspy.InputField( + desc="Current user query in Estonian, English, or Russian" + ) + previous_query: str = dspy.InputField( + desc="The user query that triggered the previous ATC API call" + ) + previous_params: str = dspy.InputField( + desc="JSON object of parameter values used in the previous API call: {param_name: value}" + ) + params_schema: str = dspy.InputField( + desc='JSON array of parameter schemas for the previous API: [{"name": str, "type": str, "required": bool, "description": str}]' + ) + + follow_up_type: str = dspy.OutputField( + desc='Classification result — MUST be exactly one of: "param_update", "response_question", "new_intent"' + ) + updated_params: str = dspy.OutputField( + desc="Valid JSON object of updated parameter values extracted from user_query. Non-empty only for param_update. Empty object {} for response_question and new_intent." + ) + + +class FollowUpDetectorModule(dspy.Module): + """DSPy Module for follow-up query classification.""" + + def __init__(self) -> None: + """Initialize follow-up detector module with Predict (direct prediction).""" + super().__init__() + self.detector = dspy.Predict(FollowUpDetectionSignature) + + def forward( + self, + user_query: str, + previous_query: str, + previous_params: Dict[str, Any], + params_schema: List[Dict[str, Any]], + ) -> FollowUpDetectionResult: + """ + Classify whether the user query is a follow-up to a previous API call. + + Args: + user_query: Current user query + previous_query: The query that triggered the previous ATC API call + previous_params: Parameter values used in the previous API call + params_schema: Parameter schemas for the previous API + + Returns: + FollowUpDetectionResult with follow_up_type and updated_params + """ + _safe_fallback: FollowUpDetectionResult = { + "follow_up_type": "new_intent", + "updated_params": {}, + } + + previous_params_json = json.dumps(previous_params, ensure_ascii=False) + sanitized_schema = [ + {**p, "description": strip_format_hints(p.get("description", ""))} + if isinstance(p, dict) + else p + for p in params_schema + ] + params_schema_json = json.dumps(sanitized_schema, ensure_ascii=False) + + result = None + _t0 = time.time() + try: + result = self.detector( + user_query=user_query, + previous_query=previous_query, + previous_params=previous_params_json, + params_schema=params_schema_json, + ) + _duration_ms = round((time.time() - _t0) * 1000, 1) + + # Parse and validate follow_up_type + follow_up_type = result.follow_up_type.strip().strip("'\"") + if follow_up_type not in _VALID_FOLLOW_UP_TYPES: + logger.warning( + f"FollowUpDetectorModule: invalid follow_up_type | event_type=follow_up_type_invalid received={follow_up_type!r}" + ) + follow_up_type = "new_intent" + + # Parse updated_params JSON only for param_update; force {} for other types + updated_params: Dict[str, Any] = {} + if follow_up_type == "param_update": + raw_updated_params = result.updated_params + try: + # Sanitize: strip whitespace and outer quotes (both single and double) + sanitized_params = raw_updated_params.strip().strip("'\"") + updated_params = json.loads(sanitized_params) + if not isinstance(updated_params, dict): + logger.warning( + f"FollowUpDetectorModule: updated_params not a dict | event_type=updated_params_not_dict received_type={type(updated_params).__name__}" + ) + updated_params = {} + except json.JSONDecodeError as exc: + _err_id = generate_error_id() + logger.error( + f"FollowUpDetectorModule: JSON parse error in updated_params | event_type=updated_params_json_error error_id={_err_id} error={exc}" + ) + updated_params = {} + + # Enforce OUTPUT RULES: updated_params must be {} unless follow_up_type is param_update. + # This prevents unintended params from leaking downstream even if the LLM returns them. + if follow_up_type != "param_update": + updated_params = {} + else: + # Validate against schema to drop unexpected/injected parameters + updated_params = _validate_updated_params(updated_params, params_schema) + + logger.debug( + f"FollowUpDetectorModule: detection complete | event_type=follow_up_detection_complete follow_up_type={follow_up_type} updated_params_count={len(updated_params)} duration_ms={_duration_ms}" + ) + + return FollowUpDetectionResult( + follow_up_type=follow_up_type, + updated_params=updated_params, + ) + + except Exception as exc: + _err_id = generate_error_id() + logger.error( + f"FollowUpDetectorModule: detection failed | event_type=follow_up_detection_failed error_id={_err_id} error={exc}" + ) + return _safe_fallback diff --git a/src/tool_classifier/intent_decomposer.py b/src/tool_classifier/intent_decomposer.py new file mode 100644 index 00000000..408b17d1 --- /dev/null +++ b/src/tool_classifier/intent_decomposer.py @@ -0,0 +1,249 @@ +"""Intent Decomposer — detects multi-intent queries and produces focused sub-queries.""" + +import asyncio +import json +from dataclasses import dataclass, field +from typing import Any, cast + +import dspy +from src.loki_logger import LokiLogger + +from src.utils.cost_utils import get_lm_usage_since +from src.utils.observation_utils import safe_observation_context +from tool_classifier.constants import MULTI_API_MAX_ENDPOINTS + +logger = LokiLogger(service_name="api-tool-calling") + + +def _get_current_model_name() -> str: + """Best-effort model name lookup from current DSPy LM.""" + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "model"): + model_name = lm.model + if isinstance(model_name, str) and model_name: + return model_name + except Exception: + pass + return "unknown" + + +@dataclass +class DecompositionResult: + """Result returned by IntentDecomposerModule.decompose(). + + Attributes: + mode: ``"single"`` if the query has one intent, ``"parallel"`` if + it contains multiple independent intents each requiring a + separate API call. + sub_queries: Focused sub-queries for each detected intent. + Empty list when ``mode="single"``. + Length is capped at ``MULTI_API_MAX_ENDPOINTS``. + """ + + mode: str + sub_queries: list[str] = field(default_factory=list) + + def set_lm_usage(self, *args: object, **kwargs: object) -> None: + """No-op stub for DSPy's internal Langfuse callback compatibility.""" + + +class IntentDecompositionSignature(dspy.Signature): + """Determine whether a user query contains a single intent or multiple + independent intents that each require a separate API call. + + Rules: + - Return mode="single" if the query has one clear intent + - Return mode="single" if you are uncertain — never force parallel + - Return mode="parallel" ONLY if the query CLEARLY asks for 2 or 3 + completely distinct, independently answerable things + - Each sub_query must be self-contained — answerable without the others + - Each sub_query must be phrased using domain-specific, descriptive + language that matches how the service would be described (e.g., + "address lookup and location search" not "find an address for me", + "vehicle tax calculation" not "calculate my tax"). This is critical + because sub_queries are used for semantic vector search. + - sub_queries must be a valid JSON list of strings, e.g. ["q1", "q2"] + - When mode="single", sub_queries MUST be an empty JSON list: [] + - Understands Estonian, English, and Russian queries + - Examples of multi-intent: "holidays in Estonia AND weather in Tallinn" + - Examples of single-intent: "renew my ID card", "book an appointment" + """ + + user_query: str = dspy.InputField( + desc="User's full natural language query in Estonian, English, or Russian" + ) + mode: str = dspy.OutputField(desc='Either "single" or "parallel"') + sub_queries: str = dspy.OutputField( + desc=( + 'JSON list of focused sub-queries when mode="parallel", ' + 'or empty list [] when mode="single". ' + "Each sub-query must be self-contained, independently answerable, " + "and phrased with domain-specific descriptive language suitable for " + "semantic search (e.g. 'address lookup and location search' rather " + "than 'find an address for me')." + ) + ) + + +class IntentDecomposerModule(dspy.Module): + """DSPy module that detects whether a user query contains multiple independent + intents and decomposes it into focused sub-queries for parallel API search. + + This module is invoked only when the gate search returns a medium-confidence + result (cosine score in the ambiguous band), indicating the query embedding + may be diluted by multiple intents. + + The ``decompose()`` coroutine wraps the synchronous DSPy forward call in a + thread pool so it does not block the asyncio event loop. + """ + + def __init__(self) -> None: + super().__init__() + self.predictor = dspy.Predict(IntentDecompositionSignature) + + def forward( + self, + user_query: str, + ) -> DecompositionResult: + """Run intent decomposition synchronously (called from a thread pool). + + Args: + user_query: The user's full natural language query. + + Returns: + DecompositionResult with mode and sub_queries. + Returns mode="single" with empty sub_queries on any LLM or + parse failure — conservative fallback to the existing single path. + """ + try: + prediction = self.predictor(user_query=user_query) + raw_mode = prediction.mode.strip().lower() + raw_sub_queries = prediction.sub_queries.strip() + + if raw_mode not in ("single", "parallel"): + logger.warning( + f"IntentDecomposer: unexpected mode={raw_mode!r} — " + f"falling back to single" + ) + return DecompositionResult(mode="single") + + if raw_mode == "single": + return DecompositionResult(mode="single") + + # mode = "parallel" — parse and validate sub_queries JSON + sub_queries = _parse_sub_queries(raw_sub_queries) + if len(sub_queries) < 2: + logger.warning( + f"IntentDecomposer: mode=parallel but only " + f"{len(sub_queries)} sub-query parsed — falling back to single" + ) + return DecompositionResult(mode="single") + + # Cap at MULTI_API_MAX_ENDPOINTS + if len(sub_queries) > MULTI_API_MAX_ENDPOINTS: + logger.info( + f"IntentDecomposer: {len(sub_queries)} sub-queries exceed cap " + f"({MULTI_API_MAX_ENDPOINTS}) — truncating" + ) + sub_queries = sub_queries[:MULTI_API_MAX_ENDPOINTS] + + logger.info( + f"IntentDecomposer: mode=parallel, " + f"{len(sub_queries)} sub-queries: {sub_queries}" + ) + return DecompositionResult(mode="parallel", sub_queries=sub_queries) + + except Exception as exc: + logger.error( + f"IntentDecomposer: decomposition failed: {exc} — " + f"falling back to single", + exc_info=True, + ) + return DecompositionResult(mode="single") + + async def decompose(self, user_query: str) -> DecompositionResult: + """Async wrapper — runs forward() in a thread pool. + + Keeps the asyncio event loop unblocked while the synchronous DSPy + LLM call executes. + + Args: + user_query: The user's full natural language query. + + Returns: + DecompositionResult with mode and sub_queries. + """ + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning( + f"Failed to get LM history length for intent decomposition: {e}" + ) + + with safe_observation_context( + name="intent_decomposition_llm", + as_type="generation", + input={"user_query": user_query}, + ) as generation: + result = cast( + DecompositionResult, await asyncio.to_thread(self, user_query) + ) + + # Update Langfuse observation with output and usage + try: + if generation is not None: + usage = get_lm_usage_since(history_length_before) + generation.update( + model=_get_current_model_name(), + output={"mode": result.mode, "sub_queries": result.sub_queries}, + usage_details={ + "input": usage.get("total_prompt_tokens", 0), + "output": usage.get("total_completion_tokens", 0), + "total": usage.get("total_tokens", 0), + }, + cost_details={"total": usage.get("total_cost", 0.0)}, + ) + except Exception as e: + logger.debug(f"Langfuse generation update skipped: {e}") + + return result + + +def _parse_sub_queries(raw: str) -> list[str]: + """Parse a JSON list of sub-query strings from LLM output. + + Handles common LLM output variations: markdown code fences, extra + whitespace, and non-list JSON values. Returns an empty list on any + parse failure. + + Args: + raw: Raw string output from the LLM (expected to be a JSON list). + + Returns: + List of non-empty sub-query strings, or [] on parse error. + """ + # Strip markdown code fences if present + cleaned = raw.strip() + if cleaned.startswith("```"): + lines = cleaned.splitlines() + # Drop opening fence (and optional language tag) and closing fence + inner = [line for line in lines[1:] if line.strip() != "```"] + cleaned = "\n".join(inner).strip() + + try: + parsed: Any = json.loads(cleaned) + except json.JSONDecodeError: + logger.warning(f"IntentDecomposer: could not parse sub_queries JSON: {raw!r}") + return [] + + if not isinstance(parsed, list): + logger.warning( + f"IntentDecomposer: sub_queries is not a list: {type(parsed).__name__}" + ) + return [] + + return [str(item).strip() for item in parsed if str(item).strip()] diff --git a/src/tool_classifier/intent_detector.py b/src/tool_classifier/intent_detector.py index a2abb74f..52f2af9c 100644 --- a/src/tool_classifier/intent_detector.py +++ b/src/tool_classifier/intent_detector.py @@ -4,7 +4,28 @@ from typing import Any, Dict, List, Optional import dspy -from loguru import logger +from langfuse import observe +from src.loki_logger import LokiLogger + +from src.utils.cost_utils import get_lm_usage_since +from src.utils.observation_utils import update_observation_safe + + +# Initialize Loki logger +logger = LokiLogger(service_name="intent-detector") + + +def _get_current_model_name() -> str: + """Best-effort model name lookup from current DSPy LM.""" + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "model"): + model_name = lm.model + if isinstance(model_name, str) and model_name: + return model_name + except Exception: + pass + return "unknown" class ServiceIntentDetector(dspy.Signature): @@ -46,6 +67,7 @@ def __init__(self) -> None: super().__init__() self.detector = dspy.Predict(ServiceIntentDetector) + @observe(name="service_intent_detection_llm", as_type="generation") def forward( self, user_query: str, @@ -77,6 +99,14 @@ def forward( services_json = json.dumps(services_formatted, ensure_ascii=False, indent=2) + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning(f"Failed to get LM history length for intent detection: {e}") + # Format conversation history if conversation_history: history_lines = [] @@ -111,12 +141,45 @@ def forward( intent_data.setdefault("entities", {}) intent_data.setdefault("reasoning", "") + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_query": user_query, + "services_count": len(services), + }, + output_data={ + "matched_service_id": intent_data.get("matched_service_id"), + "confidence": intent_data.get("confidence", 0.0), + "entities": intent_data.get("entities", {}), + }, + metadata={ + "model": _get_current_model_name(), + "usage": usage, + }, + ) + return intent_data except json.JSONDecodeError as e: logger.error(f"Failed to parse intent JSON: {e}") if result: logger.error(f"Raw response: {result.intent_result}") + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_query": user_query, + "services_count": len(services), + }, + output_data={ + "matched_service_id": None, + "confidence": 0.0, + "error": f"JSON parse error: {e}", + }, + metadata={ + "model": _get_current_model_name(), + "usage": usage, + }, + ) return { "matched_service_id": None, "confidence": 0.0, @@ -125,6 +188,22 @@ def forward( } except Exception as e: logger.error(f"Intent detection forward failed: {e}", exc_info=True) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_query": user_query, + "services_count": len(services), + }, + output_data={ + "matched_service_id": None, + "confidence": 0.0, + "error": f"Detection error: {e}", + }, + metadata={ + "model": _get_current_model_name(), + "usage": usage, + }, + ) return { "matched_service_id": None, "confidence": 0.0, diff --git a/src/tool_classifier/models.py b/src/tool_classifier/models.py index 2f830fd5..933386af 100644 --- a/src/tool_classifier/models.py +++ b/src/tool_classifier/models.py @@ -134,3 +134,43 @@ def is_client_error(self) -> bool: def is_server_error(self) -> bool: """True if the response was a 5xx server error.""" return 500 <= self.status_code < 600 + + +@dataclass +class MultiAPICallResult: + """ + Result returned by MultiAPICaller.call_all() after executing a batch of concurrent + HTTP requests. + + Preserves input order: ``results[i]`` corresponds to ``endpoints[i]``. + + Attributes: + results: One :class:`APICallResult` per endpoint, in input order. + endpoints: Corresponding endpoint metadata dicts (url, method, call_params, …). + """ + + results: list[APICallResult] + endpoints: list[Dict[str, Any]] + + @property + def all_succeeded(self) -> bool: + """True if every result in the batch has ``success=True``.""" + return all(r.success for r in self.results) + + @property + def successful_results(self) -> list[tuple[Dict[str, Any], APICallResult]]: + """Return ``(endpoint, result)`` pairs where ``result.success`` is True.""" + return [ + (ep, res) + for ep, res in zip(self.endpoints, self.results, strict=True) + if res.success + ] + + @property + def failed_results(self) -> list[tuple[Dict[str, Any], APICallResult]]: + """Return ``(endpoint, result)`` pairs where ``result.success`` is False.""" + return [ + (ep, res) + for ep, res in zip(self.endpoints, self.results, strict=True) + if not res.success + ] diff --git a/src/tool_classifier/multi_agentic_loop.py b/src/tool_classifier/multi_agentic_loop.py new file mode 100644 index 00000000..4ce8ae32 --- /dev/null +++ b/src/tool_classifier/multi_agentic_loop.py @@ -0,0 +1,1420 @@ +"""Multi-endpoint agentic loop for parallel parameter collection.""" + +import asyncio +import json +import time +from typing import Any, Dict, List, Optional, Tuple + +from src.loki_logger import LokiLogger +from src.utils.error_utils import generate_error_id + +from models.session_models import EndpointSessionState +from utils.api_tool_session_store import APIToolSessionStore +from tool_classifier.constants import ( + CONTINUATION_QUESTION, + CONTINUATION_QUESTION_ET, + CONTINUATION_QUESTION_RU, + MULTI_API_MAX_TURNS, + MULTI_INTENT_CONTINUATION_TURN, + MULTI_INTENT_MAX_TURNS, +) +from tool_classifier.continuation_utils import detect_continuation_response +from tool_classifier.enums import AgenticLoopStatus +from tool_classifier.models import AgenticLoopResult +from tool_classifier.param_extractor import ParamExtractionModule, strip_format_hints + +logger = LokiLogger(service_name="api-tool-calling") + +_CONTINUATION_QUESTIONS: dict[str, str] = { + "en": CONTINUATION_QUESTION, + "et": CONTINUATION_QUESTION_ET, + "ru": CONTINUATION_QUESTION_RU, +} + + +class MultiEndpointAgenticLoop: + """Multi-turn parameter collection loop for parallel endpoint execution. + + Merges the parameter schemas of all active endpoints, deduplicates shared + params, asks the user **once** per turn, and distributes extracted values + back to the per-endpoint :class:`~models.session_models.EndpointSessionState` + objects. + + The loop is stateless between HTTP requests — all mutable state is held in + the :class:`~models.session_models.EndpointSessionState` list passed in on + each call. The session store is used exclusively to persist updated state + between turns. + + Turn limit + ---------- + * Multi-intent (>1 endpoint): ``max_turns = MULTI_INTENT_MAX_TURNS`` (6) + * Single-intent: ``max_turns = min(3 * num_endpoints, MULTI_API_MAX_TURNS)`` + + Continuation threshold + ---------------------- + * Multi-intent (>1 endpoint): ``continuation_turn = MULTI_INTENT_CONTINUATION_TURN`` (4) + * Single-intent: ``continuation_turn = num_endpoints + 1`` + + Fixed turn-4 continuation for multi-intent gives the user enough turns to + supply all intent parameters before being asked whether to keep going. + All 6 turns (``updated_turn_count`` 1–6) execute normally; the RAG fallback + is triggered on the **7th** call, when the pre-increment ``turn_count`` + reaches ``MULTI_INTENT_MAX_TURNS`` (6) and the guard ``turn_count >= + max_turns`` fires before any extraction is attempted. + + Typical usage:: + + loop = MultiEndpointAgenticLoop( + session_store=app.state.session_store, + param_extractor=ParamExtractionModule(), + ) + + result = await loop.run_turn( + chat_id=request.chatId, + user_message=request.message, + conversation_history=request.conversationHistory, + endpoint_states=session.parallel_endpoints, + turn_count=session.turn_count, + ) + + if result.status == AgenticLoopStatus.COMPLETED: + # All endpoints have their params — proceed to API calls + ... + elif result.status == AgenticLoopStatus.NEEDS_INPUT: + # Session already saved inside run_turn — return question to user + ... + else: # MAX_TURNS_REACHED + # Delete session and fall back gracefully + ... + """ + + def __init__( + self, + session_store: Optional[APIToolSessionStore], + param_extractor: ParamExtractionModule, + ) -> None: + """Initialise the loop with injected session store and param extractor. + + Args: + session_store: Redis-backed store used to persist loop state between + HTTP requests. Injected to allow easy mocking in tests. May be + ``None`` in environments where session persistence is unavailable. + param_extractor: DSPy module that extracts parameter values from a + user message. Injected to allow easy mocking in tests. + """ + self._session_store: Optional[APIToolSessionStore] = session_store + self._param_extractor = param_extractor + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + + async def run_turn( + self, + chat_id: str, + user_message: str, + conversation_history: List[Dict[str, Any]], + endpoint_states: List[EndpointSessionState], + turn_count: int, + awaiting_continuation: bool = False, + session_language: str = "en", + continuation_language: Optional[str] = None, + ) -> AgenticLoopResult: + """Process one user turn of the multi-endpoint parameter-collection loop. + + Steps: + 0. Continuation decision — if ``awaiting_continuation`` is True, detect + whether the user said yes (keep going) or no (fall back to RAG). + 1. Turn limit guard — return MAX_TURNS_REACHED when exhausted. + 2. Build merged schema from all incomplete endpoints (deduplicating by name). + 3. Compute merged already_collected (union of per-endpoint collected_params). + 4. Call :class:`~param_extractor.ParamExtractionModule` with the merged schema. + 5. Distribute extracted params back to owning endpoints; mark endpoints + completed when all their required params are present. + 6. If all endpoints complete → save session, return COMPLETED. + 7. At the continuation threshold → save session, return + AWAITING_CONTINUATION_DECISION. + 8. Otherwise → save session, return NEEDS_INPUT with the clarifying question. + + Args: + chat_id: Unique conversation identifier used as the Redis session key. + user_message: The user's latest message for this turn. + conversation_history: Recent conversation turns as a list of + ``{"authorRole": str, "message": str}`` dicts. + endpoint_states: Per-endpoint state objects. Mutated in-place when + params are distributed. + turn_count: The current turn index (0-based before this call). + awaiting_continuation: True when the previous turn returned + AWAITING_CONTINUATION_DECISION and we are now processing the + user's yes/no reply. + session_language: Language code for clarifying questions (``"en"``, + ``"et"``, or ``"ru"``). + continuation_language: Override language for the continuation + yes/no question. Falls back to ``session_language`` when None. + + Returns: + :class:`~models.AgenticLoopResult` with updated status, + collected_params (merged across all endpoints), and turn_count. + """ + if not endpoint_states: + logger.debug( + f"MultiEndpointAgenticLoop: no endpoints provided for chat_id={chat_id} — returning COMPLETED" + ) + return AgenticLoopResult( + status=AgenticLoopStatus.COMPLETED, + collected_params={}, + clarifying_question="", + turn_count=turn_count, + ) + + num_endpoints = len(endpoint_states) + max_turns, continuation_turn = self._compute_turn_limits(num_endpoints) + updated_turn_count = turn_count + 1 + + logger.info( + f"MultiEndpointAgenticLoop: multi loop turn started | event_type=multi_loop_turn_started" + f" chat_id={chat_id} turn_count={turn_count} endpoint_count={num_endpoints}" + f" incomplete_count={sum(1 for s in endpoint_states if not s.completed)}" + ) + + # Step 0 — Continuation decision + if awaiting_continuation: + wants_to_continue = detect_continuation_response(user_message) + if wants_to_continue: + logger.info( + f"MultiEndpointAgenticLoop: continuation user accepted | event_type=continuation_user_accepted" + f" chat_id={chat_id} turn_count={turn_count}" + ) + awaiting_continuation = False + else: + logger.info( + f"MultiEndpointAgenticLoop: continuation user declined | event_type=continuation_user_declined" + f" chat_id={chat_id} turn_count={turn_count} status=max_turns_reached" + ) + return AgenticLoopResult( + status=AgenticLoopStatus.MAX_TURNS_REACHED, + collected_params=self._merged_collected(endpoint_states), + clarifying_question="", + turn_count=updated_turn_count, + ) + + # Step 1 — Turn limit guard (no session save — caller deletes) + if turn_count >= max_turns: + logger.warning( + f"MultiEndpointAgenticLoop: turn limit reached | event_type=turn_limit_reached" + f" chat_id={chat_id} turn_count={turn_count} max_turns={max_turns}" + ) + return AgenticLoopResult( + status=AgenticLoopStatus.MAX_TURNS_REACHED, + collected_params=self._merged_collected(endpoint_states), + clarifying_question="", + turn_count=updated_turn_count, + ) + + # Step 2 — Build merged schema + param owner map + namespace map + merged_schema, param_owners, namespace_map = self._build_merged_schema( + endpoint_states + ) + + logger.debug( + f"MultiEndpointAgenticLoop: schema merged | event_type=schema_merged" + f" chat_id={chat_id} turn_count={turn_count}" + f" total_param_count={sum(len(s.endpoint.get('params', [])) for s in endpoint_states if not s.completed)}" + f" deduplicated_count={len(merged_schema)} namespaced_count={len(namespace_map)}" + ) + + # Step 3 — Namespaced already_collected for the LLM extractor + merged_already_collected = self._build_namespaced_already_collected( + endpoint_states, namespace_map + ) + + # Step 4 — Extract params from the current user message + intent_groups = self._build_intent_groups( + endpoint_states, merged_already_collected, namespace_map + ) + _t0 = time.time() + try: + extraction = await asyncio.to_thread( + self._param_extractor, + user_message, + merged_schema, + conversation_history, + merged_already_collected, + session_language, + turn_count, + intent_groups, + ) + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.debug( + f"MultiEndpointAgenticLoop: param extraction complete | event_type=param_extraction_complete" + f" chat_id={chat_id} turn_count={turn_count}" + f" extracted_count={len(extraction['extracted_params'])} duration_ms={_duration_ms}" + ) + except Exception as exc: + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.error( + f"MultiEndpointAgenticLoop: param extraction failed | event_type=param_extraction_failed" + f" chat_id={chat_id} turn_count={turn_count}" + f" error_id={generate_error_id()} duration_ms={_duration_ms} exc={exc}" + ) + await self._save_session( + chat_id, + endpoint_states, + updated_turn_count, + awaiting_continuation=awaiting_continuation, + ) + return AgenticLoopResult( + status=AgenticLoopStatus.NEEDS_INPUT, + collected_params=self._merged_collected(endpoint_states), + clarifying_question="", + turn_count=updated_turn_count, + ) + + # Step 5 — Distribute extracted params to owning endpoints + previously_completed_indices: set[int] = { + i for i, s in enumerate(endpoint_states) if s.completed + } + self._distribute_params( + extraction["extracted_params"], endpoint_states, param_owners, namespace_map + ) + + # Step 5b — If the user gave a single value that was duplicated across + # multiple endpoints' conflicting params, clear the duplicates so the loop + # re-asks for each endpoint's params separately on the next turn. + enforced = self._enforce_sequential_conflicting_params( + endpoint_states, namespace_map + ) + + # Step 5c — If multiple endpoints with non-conflicting param names were all + # completed by the same single-turn user response (e.g. one date range applied + # to both electricity prices and parliament stats), keep only the first newly- + # completed endpoint and clear the rest so they are asked for separately. + enforced = ( + self._enforce_sequential_parallel_completion( + endpoint_states, previously_completed_indices + ) + or enforced + ) + + # Step 6 — Check global completion + all_done = all(state.completed for state in endpoint_states) + merged_after = self._merged_collected(endpoint_states) + + if all_done: + logger.info( + f"MultiEndpointAgenticLoop: all endpoints completed | event_type=multi_loop_all_endpoints_completed" + f" chat_id={chat_id} turn_count={turn_count} endpoint_count={num_endpoints}" + f" status=completed duration_ms={_duration_ms}" + ) + await self._save_session( + chat_id, + endpoint_states, + updated_turn_count, + awaiting_continuation=False, + ) + return AgenticLoopResult( + status=AgenticLoopStatus.COMPLETED, + collected_params=merged_after, + clarifying_question="", + turn_count=updated_turn_count, + ) + + # Step 7 — Still missing params + logger.debug( + f"MultiEndpointAgenticLoop: loop needs input | event_type=loop_needs_input" + f" chat_id={chat_id} turn_count={turn_count}" + f" missing_params={extraction['missing_required']} status=needs_input" + ) + + # At exactly the continuation threshold, ask whether to keep going. + if updated_turn_count == continuation_turn: + logger.info( + f"MultiEndpointAgenticLoop: continuation threshold reached | event_type=continuation_threshold_reached" + f" chat_id={chat_id} turn_count={turn_count} continuation_turn={continuation_turn}" + f" missing_count={len(extraction['missing_required'])}" + ) + effective_lang = continuation_language or session_language + continuation_q = _CONTINUATION_QUESTIONS.get( + effective_lang, CONTINUATION_QUESTION + ) + await self._save_session( + chat_id, endpoint_states, updated_turn_count, awaiting_continuation=True + ) + return AgenticLoopResult( + status=AgenticLoopStatus.AWAITING_CONTINUATION_DECISION, + collected_params=merged_after, + clarifying_question=continuation_q, + turn_count=updated_turn_count, + ) + + await self._save_session( + chat_id, endpoint_states, updated_turn_count, awaiting_continuation=False + ) + + # If enforcement cleared duplicate params from a later endpoint and the LLM's + # question is stale (it said "none" believing all params were satisfied), we + # must regenerate a fresh question that covers the remaining missing params. + clarifying_q = extraction["clarifying_question"] + if enforced and clarifying_q.strip().lower() == "none": + clarifying_q = await self._regenerate_question_after_enforcement( + endpoint_states, merged_schema, namespace_map, session_language + ) + # If regeneration failed/returned empty, use a fallback question to prevent dead-end + if not clarifying_q.strip(): + clarifying_q = self._build_fallback_question( + endpoint_states, namespace_map, session_language + ) + + return AgenticLoopResult( + status=AgenticLoopStatus.NEEDS_INPUT, + collected_params=merged_after, + clarifying_question=clarifying_q, + turn_count=updated_turn_count, + ) + + async def stream_run_turn( + self, + chat_id: str, + user_message: str, + conversation_history: List[Dict[str, Any]], + endpoint_states: List[EndpointSessionState], + turn_count: int, + awaiting_continuation: bool = False, + session_language: str = "en", + continuation_language: Optional[str] = None, + ) -> Tuple[AgenticLoopResult, List[str]]: + """Process one user turn like :meth:`run_turn` but stream clarifying question tokens. + + Delegates extraction to + :meth:`~param_extractor.ParamExtractionModule.stream_forward` so + ``clarifying_question`` tokens are captured as they arrive from the LLM. + All session management is identical to :meth:`run_turn`. + + Returns: + Tuple of ``(AgenticLoopResult, question_tokens)``. + ``question_tokens`` is the list of streamed token strings for the + clarifying question, or an empty list when no question is needed. + """ + if not endpoint_states: + logger.debug( + f"MultiEndpointAgenticLoop: no endpoints provided for chat_id={chat_id} — returning COMPLETED" + ) + return ( + AgenticLoopResult( + status=AgenticLoopStatus.COMPLETED, + collected_params={}, + clarifying_question="", + turn_count=turn_count, + ), + [], + ) + + num_endpoints = len(endpoint_states) + max_turns, continuation_turn = self._compute_turn_limits(num_endpoints) + updated_turn_count = turn_count + 1 + + logger.info( + f"MultiEndpointAgenticLoop: multi loop turn started | event_type=multi_loop_turn_started" + f" chat_id={chat_id} turn_count={turn_count} endpoint_count={num_endpoints}" + f" incomplete_count={sum(1 for s in endpoint_states if not s.completed)}" + ) + + # Step 0 — Continuation decision + if awaiting_continuation: + wants_to_continue = detect_continuation_response(user_message) + if wants_to_continue: + logger.info( + f"MultiEndpointAgenticLoop: continuation user accepted | event_type=continuation_user_accepted" + f" chat_id={chat_id} turn_count={turn_count}" + ) + awaiting_continuation = False + else: + logger.info( + f"MultiEndpointAgenticLoop: continuation user declined | event_type=continuation_user_declined" + f" chat_id={chat_id} turn_count={turn_count} status=max_turns_reached" + ) + return ( + AgenticLoopResult( + status=AgenticLoopStatus.MAX_TURNS_REACHED, + collected_params=self._merged_collected(endpoint_states), + clarifying_question="", + turn_count=updated_turn_count, + ), + [], + ) + + # Step 1 — Turn limit guard + if turn_count >= max_turns: + logger.warning( + f"MultiEndpointAgenticLoop: turn limit reached | event_type=turn_limit_reached" + f" chat_id={chat_id} turn_count={turn_count} max_turns={max_turns}" + ) + return ( + AgenticLoopResult( + status=AgenticLoopStatus.MAX_TURNS_REACHED, + collected_params=self._merged_collected(endpoint_states), + clarifying_question="", + turn_count=updated_turn_count, + ), + [], + ) + + # Step 2 — Build merged schema + param owner map + namespace map + merged_schema, param_owners, namespace_map = self._build_merged_schema( + endpoint_states + ) + + # Step 3 — Namespaced already_collected for the LLM extractor + merged_already_collected = self._build_namespaced_already_collected( + endpoint_states, namespace_map + ) + + # Step 4 — Stream-extract params from the current user message + intent_groups = self._build_intent_groups( + endpoint_states, merged_already_collected, namespace_map + ) + _t0 = time.time() + try: + question_tokens, extraction = await self._param_extractor.stream_forward( + user_message=user_message, + params_schema=merged_schema, + conversation_history=conversation_history, + already_collected=merged_already_collected, + session_language=session_language, + turn_count=turn_count, + intent_groups=intent_groups, + ) + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.debug( + f"MultiEndpointAgenticLoop: param extraction complete | event_type=param_extraction_complete" + f" chat_id={chat_id} turn_count={turn_count}" + f" extracted_count={len(extraction['extracted_params'])} duration_ms={_duration_ms}" + ) + except Exception as exc: + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.error( + f"MultiEndpointAgenticLoop: param extraction failed | event_type=param_extraction_failed" + f" chat_id={chat_id} turn_count={turn_count}" + f" error_id={generate_error_id()} duration_ms={_duration_ms} exc={exc}" + ) + await self._save_session( + chat_id, + endpoint_states, + updated_turn_count, + awaiting_continuation=awaiting_continuation, + ) + return ( + AgenticLoopResult( + status=AgenticLoopStatus.NEEDS_INPUT, + collected_params=self._merged_collected(endpoint_states), + clarifying_question="", + turn_count=updated_turn_count, + ), + [], + ) + + # Step 5 — Distribute extracted params to owning endpoints + previously_completed_indices_stream: set[int] = { + i for i, s in enumerate(endpoint_states) if s.completed + } + self._distribute_params( + extraction["extracted_params"], endpoint_states, param_owners, namespace_map + ) + + # Step 5b — If the user gave a single value that was duplicated across + # multiple endpoints' conflicting params, clear the duplicates so the loop + # re-asks for each endpoint's params separately on the next turn. + enforced = self._enforce_sequential_conflicting_params( + endpoint_states, namespace_map + ) + + # Step 5c — If multiple endpoints with non-conflicting param names were all + # completed by the same single-turn user response, keep only the first newly- + # completed endpoint and clear the rest so they are asked for separately. + enforced = ( + self._enforce_sequential_parallel_completion( + endpoint_states, previously_completed_indices_stream + ) + or enforced + ) + + # Step 6 — Check global completion + all_done = all(state.completed for state in endpoint_states) + merged_after = self._merged_collected(endpoint_states) + + if all_done: + logger.info( + f"MultiEndpointAgenticLoop: all endpoints completed | event_type=multi_loop_all_endpoints_completed" + f" chat_id={chat_id} turn_count={turn_count} endpoint_count={num_endpoints}" + f" status=completed duration_ms={_duration_ms}" + ) + await self._save_session( + chat_id, + endpoint_states, + updated_turn_count, + awaiting_continuation=False, + ) + return ( + AgenticLoopResult( + status=AgenticLoopStatus.COMPLETED, + collected_params=merged_after, + clarifying_question="", + turn_count=updated_turn_count, + ), + [], + ) + + # Step 7 — Still missing params + logger.debug( + f"MultiEndpointAgenticLoop: loop needs input | event_type=loop_needs_input" + f" chat_id={chat_id} turn_count={turn_count}" + f" missing_params={extraction['missing_required']} status=needs_input" + ) + + if updated_turn_count == continuation_turn: + logger.info( + f"MultiEndpointAgenticLoop: continuation threshold reached | event_type=continuation_threshold_reached" + f" chat_id={chat_id} turn_count={turn_count} continuation_turn={continuation_turn}" + f" missing_count={len(extraction['missing_required'])}" + ) + effective_lang = continuation_language or session_language + continuation_q = _CONTINUATION_QUESTIONS.get( + effective_lang, CONTINUATION_QUESTION + ) + await self._save_session( + chat_id, endpoint_states, updated_turn_count, awaiting_continuation=True + ) + words = continuation_q.split(" ") + continuation_tokens = [ + w + " " if i < len(words) - 1 else w for i, w in enumerate(words) + ] + return ( + AgenticLoopResult( + status=AgenticLoopStatus.AWAITING_CONTINUATION_DECISION, + collected_params=merged_after, + clarifying_question=continuation_q, + turn_count=updated_turn_count, + ), + continuation_tokens, + ) + + await self._save_session( + chat_id, endpoint_states, updated_turn_count, awaiting_continuation=False + ) + + # If enforcement cleared duplicate params from a later endpoint and the LLM's + # question is stale (it said "none" believing all params were satisfied), we + # must regenerate a fresh question that covers the remaining missing params. + # The stale question_tokens are also discarded so the workflow falls back to + # streaming the regenerated question as a single chunk. + clarifying_q = extraction["clarifying_question"] + final_question_tokens = question_tokens + if enforced and clarifying_q.strip().lower() == "none": + clarifying_q = await self._regenerate_question_after_enforcement( + endpoint_states, merged_schema, namespace_map, session_language + ) + final_question_tokens = [] + # If regeneration failed/returned empty, use a fallback question to prevent dead-end + if not clarifying_q.strip(): + clarifying_q = self._build_fallback_question( + endpoint_states, namespace_map, session_language + ) + + return ( + AgenticLoopResult( + status=AgenticLoopStatus.NEEDS_INPUT, + collected_params=merged_after, + clarifying_question=clarifying_q, + turn_count=updated_turn_count, + ), + final_question_tokens, + ) + + # ------------------------------------------------------------------ + # Private helpers + # ------------------------------------------------------------------ + + async def _regenerate_question_after_enforcement( + self, + endpoint_states: List[EndpointSessionState], + merged_schema: List[Dict[str, Any]], + namespace_map: Dict[str, Tuple[int, str]], + session_language: str, + ) -> str: + """Generate a fresh clarifying question after sequential param enforcement. + + Called when :meth:`_enforce_sequential_conflicting_params` cleared + duplicate values from a later endpoint, making it incomplete again after + the LLM had already reported all params as collected (clarifying_question + == "none"). The stale "none" question is discarded and a new question is + produced from the updated endpoint states. + + An empty user message is used so the extractor performs no new extraction + — it only observes what is still missing in the updated ``already_collected`` + and generates the appropriate clarifying question for those params. + + Args: + endpoint_states: Per-endpoint states after enforcement (some may have + had params cleared and ``completed`` reset to ``False``). + merged_schema: The merged parameter schema built earlier this turn. + namespace_map: Namespacing map from :meth:`_build_merged_schema`. + session_language: Language code for the clarifying question. + + Returns: + A natural-language question for the still-missing params, or an empty + string if regeneration fails or no params remain missing. + """ + updated_already_collected = self._build_namespaced_already_collected( + endpoint_states, namespace_map + ) + updated_intent_groups = self._build_intent_groups( + endpoint_states, updated_already_collected, namespace_map + ) + if not updated_intent_groups: + return "" + try: + regen = await asyncio.to_thread( + self._param_extractor, + "", # empty — no new params to extract, just generate the question + merged_schema, + [], + updated_already_collected, + session_language, + 0, # turn_count: no acknowledgment needed for regenerated questions + updated_intent_groups, + ) + question = regen["clarifying_question"] + if question.strip().lower() == "none": + return "" + return question + except Exception as exc: + logger.warning( + f"MultiEndpointAgenticLoop: failed to regenerate clarifying question after enforcement | exc={exc}" + ) + return "" + + def _build_fallback_question( + self, + endpoint_states: List[EndpointSessionState], + namespace_map: Dict[str, Tuple[int, str]], + session_language: str, + ) -> str: + """Build a generic fallback question when LLM regeneration fails. + + Constructs a safe, non-empty question from the still-missing required + params when ``_regenerate_question_after_enforcement()`` returns empty + due to failed regeneration or empty intent groups. This prevents the + workflow from dead-ending with an empty clarifying_question when + endpoints are still incomplete. + + The question lists all missing required param descriptions from incomplete + endpoints, providing clear guidance to the user on what information is + still needed. + + Args: + endpoint_states: Per-endpoint state objects. + namespace_map: Mapping of ``namespaced_name → (ep_idx, original_name)`` + for conflicting params. + session_language: Language code for the question. + + Returns: + A generic fallback question listing missing params, or a generic + continuation prompt if no missing params are found. + """ + conflicting_original_names: set[str] = { + orig_name for (_, orig_name) in namespace_map.values() + } + + # Collect all missing required param descriptions from incomplete endpoints + missing_descriptions: List[str] = [] + seen_descriptions: set[str] = set() + + for state in endpoint_states: + if state.completed: + continue + params_schema: List[Dict[str, Any]] = state.endpoint.get("params", []) + for param in params_schema: + if not isinstance(param, dict): + continue + if not param.get("required", False): + continue + name: str = param.get("name", "") + if not name: + continue + + # Check if already collected (accounting for conflicting params) + if name in conflicting_original_names: + if name in state.collected_params: + continue + else: + # Non-conflicting: check union of all collected_params + found = False + for s in endpoint_states: + if name in s.collected_params: + found = True + break + if found: + continue + + # Add description if not already seen (deduplication) + desc = strip_format_hints(str(param.get("description", name))) + if desc and desc not in seen_descriptions: + missing_descriptions.append(desc) + seen_descriptions.add(desc) + + # Build the fallback question + if missing_descriptions: + items_str = ", ".join(missing_descriptions) + # Localized fallback prompts + if session_language == "et": + return f"Palun sisestage järgmine teave: {items_str}" + elif session_language == "ru": + return f"Пожалуйста, предоставьте следующую информацию: {items_str}" + else: # Default to English + return f"Please provide the following: {items_str}" + else: + # No specific missing params found — use generic continuation prompt + if session_language == "et": + return CONTINUATION_QUESTION_ET + elif session_language == "ru": + return CONTINUATION_QUESTION_RU + else: # Default to English + return CONTINUATION_QUESTION + + def _compute_turn_limits(self, num_endpoints: int) -> Tuple[int, int]: + """Compute max_turns and continuation_turn based on endpoint count. + + Args: + num_endpoints: Number of endpoints in the current session. + + Returns: + Tuple of (max_turns, continuation_turn). + """ + is_multi_intent = num_endpoints > 1 + max_turns = ( + MULTI_INTENT_MAX_TURNS + if is_multi_intent + else min(3 * num_endpoints, MULTI_API_MAX_TURNS) + ) + continuation_turn = ( + MULTI_INTENT_CONTINUATION_TURN if is_multi_intent else num_endpoints + 1 + ) + return max_turns, continuation_turn + + def _build_merged_schema( + self, + endpoint_states: List[EndpointSessionState], + ) -> Tuple[List[Dict[str, Any]], Dict[str, List[int]], Dict[str, Tuple[int, str]]]: + """Build a merged parameter schema from all incomplete endpoints. + + Uses a three-pass approach: + + * **Pass 0** — Counts how many *incomplete* endpoints define each param name. + This detects conflicts among endpoints that are actively requesting params. + Completed endpoints do not participate in conflict detection. Names that + appear in more than one incomplete endpoint are "conflicting" and will + be namespaced in Pass 1 so each intent can supply an independent value. + + * **Pass 1** — Iterates over *incomplete* endpoints only to build + ``merged_schema``. *Non-conflicting* params are deduplicated by name; + the first occurrence's definition (type, description) is used. A type + conflict is logged as a warning. A param is ``required=True`` in the + merged schema if it is required in **any** owning endpoint. + *Conflicting* params are emitted as ``{name}__{endpoint_idx}`` so all + intents receive their own entry in the schema simultaneously. Duplicate + values across endpoints are handled after distribution by + ``_enforce_sequential_conflicting_params``. + + * **Pass 2** — Iterates over *completed* endpoints and adds them to + ``param_owners`` for any *non-namespaced* param name already present + in the merged schema. This ensures that if a shared non-conflicting + param is re-extracted or corrected in a later turn the updated value + is distributed back to completed endpoint states. Completed endpoints + do not affect conflict detection or namespacing. + + Args: + endpoint_states: All per-endpoint states. + + Returns: + A 3-tuple of: + - ``merged_schema``: Deduplicated/namespaced list of param dicts. + - ``param_owners``: Mapping of param key (possibly namespaced) → + list of endpoint indices that own it. + - ``namespace_map``: Mapping of ``namespaced_name → + (ep_idx, original_name)`` for all conflicting params. Empty dict + when no conflicts exist. + """ + # Pass 0: detect conflicting param names across incomplete endpoints only. + # Conflicts are determined solely among endpoints that need params, not including + # completed endpoints. This ensures no unnecessary namespacing when only one + # endpoint is incomplete. + name_to_endpoints: Dict[str, set[int]] = {} + for idx, state in enumerate(endpoint_states): + if state.completed: + continue + for param in state.endpoint.get("params", []): + if not isinstance(param, dict): + continue + name: str = param.get("name", "") + if name: + name_to_endpoints.setdefault(name, set()).add(idx) + conflicting_names: set[str] = { + n for n, eps in name_to_endpoints.items() if len(eps) > 1 + } + + merged_schema: List[Dict[str, Any]] = [] + param_owners: Dict[str, List[int]] = {} + seen_types: Dict[str, str] = {} # key → type string of first occurrence + namespace_map: Dict[str, Tuple[int, str]] = {} + + # Pass 1: build merged_schema from incomplete endpoints only. + for idx, state in enumerate(endpoint_states): + if state.completed: + continue + params_schema: List[Dict[str, Any]] = state.endpoint.get("params", []) + for param in params_schema: + if not isinstance(param, dict): + continue + name = param.get("name", "") + if not name: + continue + param_type: str = str(param.get("type") or "string") + + if name in conflicting_names: + # Conflicting — emit as a namespaced key owned by this endpoint only. + ns_name: str = f"{name}__{idx}" + ns_param = dict(param) + ns_param["name"] = ns_name + merged_schema.append(ns_param) + param_owners[ns_name] = [idx] + seen_types[ns_name] = param_type + namespace_map[ns_name] = (idx, name) + elif name not in param_owners: + # Non-conflicting first occurrence — add to merged schema. + merged_schema.append(dict(param)) + param_owners[name] = [idx] + seen_types[name] = param_type + else: + # Non-conflicting duplicate name — update owner list and required flag. + param_owners[name].append(idx) + + # Warn on type conflicts + if param_type != seen_types[name]: + logger.warning( + f"MultiEndpointAgenticLoop: param '{name}' has conflicting types " + f"across endpoints ({seen_types[name]} vs {param_type}). Using first occurrence's type." + ) + + # Promote to required if required in any owner + if param.get("required", False): + for merged_param in merged_schema: + if merged_param.get("name") == name: + merged_param["required"] = True + break + + # Pass 2: register completed endpoints as owners for non-namespaced params + # that appear in the merged schema. This ensures that if a shared + # non-conflicting param is re-extracted or corrected in a later turn the + # updated value is also written back to already-completed endpoint states. + # + # Also apply the same "required in any owner" promotion used in Pass 1. + for idx, state in enumerate(endpoint_states): + if not state.completed: + continue + params_schema = state.endpoint.get("params", []) + for param in params_schema: + if not isinstance(param, dict): + continue + name = param.get("name", "") + # Only register for non-namespaced (non-conflicting) params. + if name and name in param_owners: + param_owners[name].append(idx) + if param.get("required", False): + for merged_param in merged_schema: + if merged_param.get("name") == name: + merged_param["required"] = True + break + + return merged_schema, param_owners, namespace_map + + def _build_namespaced_already_collected( + self, + endpoint_states: List[EndpointSessionState], + namespace_map: Dict[str, Tuple[int, str]], + ) -> Dict[str, Any]: + """Build the ``already_collected`` dict for the LLM, namespacing conflicting params. + + For non-conflicting params: merged union across all endpoints (original + names), excluding any names that are conflicting. + For conflicting params: namespaced keys (e.g. ``startDate__0``, + ``startDate__1``), read from each endpoint's own ``collected_params``. + + Args: + endpoint_states: All per-endpoint states. + namespace_map: Mapping of ``namespaced_name → (ep_idx, original_name)`` + produced by ``_build_merged_schema()``. Empty dict when no conflicts. + + Returns: + Dict of already-collected param values, ready to pass to the LLM + extractor. + """ + if not namespace_map: + # Fast path — no conflicts, use simple merge. + merged: Dict[str, Any] = {} + for state in endpoint_states: + merged.update(state.collected_params) + return merged + + conflicting_original_names: set[str] = { + orig_name for (_, orig_name) in namespace_map.values() + } + + # Non-conflicting: merged union, excluding conflicting original names. + result: Dict[str, Any] = { + k: v + for state in endpoint_states + for k, v in state.collected_params.items() + if k not in conflicting_original_names + } + + # Conflicting: namespaced keys from each endpoint's own collected_params. + result.update( + { + ns_name: endpoint_states[ep_idx].collected_params[original_name] + for ns_name, (ep_idx, original_name) in namespace_map.items() + if original_name in endpoint_states[ep_idx].collected_params + } + ) + + return result + + def _build_intent_groups( + self, + endpoint_states: List[EndpointSessionState], + already_collected: Dict[str, Any], + namespace_map: Dict[str, Tuple[int, str]], + ) -> List[Dict[str, Any]]: + """Build intent groups for multi-intent clarifying question separation. + + For each incomplete endpoint, collects the required params not yet present + in ``already_collected``, strips format hints from their descriptions, and + returns a list of groups suitable for the ``intent_groups`` field of + :class:`~tool_classifier.param_extractor.ParamExtractionSignature`. + + **Conflicting params** (those in ``namespace_map``) are checked against + each endpoint's own ``collected_params`` rather than the merged + ``already_collected``, and deduplication across groups is suppressed so + that both intents independently list their date-range params. + + **Non-conflicting params** preserve the existing cross-endpoint + deduplication via ``seen_param_names``. + + Returns an empty list when no groups have missing required params. + Returns a single-element list when only one endpoint has outstanding + params — the endpoint name is still included so the LLM can phrase the + question with proper context (e.g. "For the parliament member participation + stats, what is the start and end date?"). + + Args: + endpoint_states: All per-endpoint states. + already_collected: Namespaced already-collected dict (from + ``_build_namespaced_already_collected()``) used for non-conflicting + param lookup. + namespace_map: Mapping of ``namespaced_name → (ep_idx, original_name)`` + produced by ``_build_merged_schema()``. + + Returns: + List of ``{"intent": str, "missing_param_descriptions": [str, ...]}`` + dicts, or an empty list if no intent has missing required params. + """ + conflicting_original_names: set[str] = { + orig_name for (_, orig_name) in namespace_map.values() + } + + groups: List[Dict[str, Any]] = [] + seen_param_names: set[str] = set() + for state in endpoint_states: + if state.completed: + continue + params_schema: List[Dict[str, Any]] = state.endpoint.get("params", []) + missing_descriptions: List[str] = [] + for param in params_schema: + if not isinstance(param, dict): + continue + if not param.get("required", False): + continue + name: str = param.get("name", "") + if not name: + continue + + if name in conflicting_original_names: + # Conflicting param: check per-endpoint collected_params. + # Do NOT deduplicate across groups — each intent lists its own entry. + if name in state.collected_params: + continue + else: + # Non-conflicting: use merged already_collected + seen_param_names. + if name in already_collected: + continue + if name in seen_param_names: + continue + seen_param_names.add(name) + + desc = strip_format_hints(str(param.get("description", name))) + missing_descriptions.append(desc) + if missing_descriptions: + intent_name: str = state.endpoint.get("name") or state.endpoint.get( + "description", "" + ) + groups.append( + { + "intent": intent_name, + "missing_param_descriptions": missing_descriptions, + } + ) + if not groups: + return [] + return groups + + def _distribute_params( + self, + extracted_params: Dict[str, Any], + endpoint_states: List[EndpointSessionState], + param_owners: Dict[str, List[int]], + namespace_map: Dict[str, Tuple[int, str]], + ) -> None: + """Distribute extracted param values to their owning endpoint states. + + For **namespaced** keys (those in ``namespace_map``): the value is + reverse-translated and written to the specific endpoint's + ``collected_params`` under the original param name. + + For **non-namespaced** keys: the existing owner-broadcast logic writes + the value into every owning endpoint's ``collected_params``. + + Marks an endpoint as ``completed=True`` once all its required params are + present. For completed endpoints a write only occurs when the incoming + value differs from the stored one; the override is logged at WARNING + level. + + Args: + extracted_params: Dict of newly extracted param values. + endpoint_states: Per-endpoint state objects (mutated in-place). + param_owners: Mapping of param key → list of owning endpoint indices. + namespace_map: Mapping of ``namespaced_name → (ep_idx, original_name)`` + produced by ``_build_merged_schema()``. + """ + for param_name, value in extracted_params.items(): + if param_name in namespace_map: + # Namespaced key — write to the specific endpoint under the original name. + ep_idx, original_name = namespace_map[param_name] + state = endpoint_states[ep_idx] + if state.completed: + existing = state.collected_params.get(original_name) + if value == existing: + continue + logger.warning( + f"MultiEndpointAgenticLoop: overwriting completed endpoint param" + f" | event_type=endpoint_param_overwrite" + f" endpoint_name={state.endpoint.get('name', '')}" + f" param_name={original_name} old_value={repr(existing)} new_value={repr(value)}" + ) + state.collected_params[original_name] = value + else: + # Non-namespaced key — broadcast to all owning endpoints. + owner_indices = param_owners.get(param_name, []) + for idx in owner_indices: + state = endpoint_states[idx] + if state.completed: + existing = state.collected_params.get(param_name) + if value == existing: + continue + logger.warning( + f"MultiEndpointAgenticLoop: overwriting completed endpoint param" + f" | event_type=endpoint_param_overwrite" + f" endpoint_name={state.endpoint.get('name', '')}" + f" param_name={param_name} old_value={repr(existing)} new_value={repr(value)}" + ) + state.collected_params[param_name] = value + + # Check completion for each non-completed endpoint + for state in endpoint_states: + if state.completed: + continue + params_schema: List[Dict[str, Any]] = state.endpoint.get("params", []) + required_names = { + p["name"] + for p in params_schema + if isinstance(p, dict) and p.get("required", False) + } + if required_names.issubset(state.collected_params.keys()): + state.completed = True + logger.debug( + f"MultiEndpointAgenticLoop: endpoint completed | event_type=endpoint_completed" + f" endpoint_name={state.endpoint.get('name', '')}" + ) + + def _enforce_sequential_parallel_completion( + self, + endpoint_states: List[EndpointSessionState], + previously_completed_indices: set[int], + ) -> bool: + """Clear params from endpoints newly completed with duplicate values in this turn. + + When endpoints have *different* param names (non-conflicting), the existing + :meth:`_enforce_sequential_conflicting_params` does not apply. Instead, + if a single user response satisfies **multiple** endpoints at once with the + **same underlying values** (e.g. one date range fills both ``start``/``end`` + on the electricity prices endpoint *and* ``startDate``/``endDate`` on the + parliament stats endpoint with identical dates), all but the *first* + (lowest-index) such endpoint have their required params cleared so the loop + can ask for each endpoint's params separately. + + If the newly-completed endpoints have **disjoint** required-param value sets + (e.g. the user provided April dates for electricity and January dates for + parliament in a single reply), both are kept as completed because the user + intentionally supplied separate values for each intent. + + Endpoints that were already completed *before* this turn are untouched. + + Args: + endpoint_states: Per-endpoint state objects (mutated in-place). + previously_completed_indices: Set of endpoint indices that were already + completed at the start of this turn (before distribution). + + Returns: + ``True`` if any endpoint had its params cleared; ``False`` otherwise. + """ + newly_completed = [ + i + for i, state in enumerate(endpoint_states) + if state.completed and i not in previously_completed_indices + ] + + if len(newly_completed) <= 1: + return False + + first_ep_idx = newly_completed[0] + first_state = endpoint_states[first_ep_idx] + first_params_schema: List[Dict[str, Any]] = first_state.endpoint.get( + "params", [] + ) + first_required_names: set[str] = { + p["name"] + for p in first_params_schema + if isinstance(p, dict) and p.get("required", False) and p.get("name") + } + # Normalize values to hashable strings to handle unhashable types (lists, dicts). + first_values_normalized: set[str] = { + self._normalize_value_to_hashable(first_state.collected_params[name]) + for name in first_required_names + if name in first_state.collected_params + } + + cleared_any = False + for ep_idx in newly_completed[1:]: + state = endpoint_states[ep_idx] + params_schema: List[Dict[str, Any]] = state.endpoint.get("params", []) + required_names: set[str] = { + p["name"] + for p in params_schema + if isinstance(p, dict) and p.get("required", False) and p.get("name") + } + # Normalize values to hashable strings to handle unhashable types (lists, dicts). + later_values_normalized: set[str] = { + self._normalize_value_to_hashable(state.collected_params[name]) + for name in required_names + if name in state.collected_params + } + + # If both endpoints have non-empty value sets and they are completely + # disjoint, the user deliberately provided different values for each + # intent in one reply — keep this endpoint completed. + if ( + first_values_normalized + and later_values_normalized + and later_values_normalized.isdisjoint(first_values_normalized) + ): + logger.debug( + f"MultiEndpointAgenticLoop: parallel endpoints completed with distinct values — keeping both" + f" | event_type=enforcement_parallel_completion" + f" endpoint_name={state.endpoint.get('name', '')}" + f" first_endpoint_name={first_state.endpoint.get('name', '')}" + ) + continue + + # Value sets are exactly equal — the LLM likely applied the same answer + # to multiple endpoints. Clear this endpoint so the loop re-asks. + if later_values_normalized == first_values_normalized: + logger.debug( + f"MultiEndpointAgenticLoop: clearing duplicate params from endpoint" + f" | event_type=enforcement_sequential_conflict" + f" endpoint_name={state.endpoint.get('name', '')}" + ) + for param in params_schema: + if isinstance(param, dict) and param.get("required", False): + name: str = param.get("name", "") + if name: + state.collected_params.pop(name, None) + state.completed = False + cleared_any = True + + return cleared_any + + def _enforce_sequential_conflicting_params( + self, + endpoint_states: List[EndpointSessionState], + namespace_map: Dict[str, Tuple[int, str]], + ) -> bool: + """Clear duplicate conflicting param values from later endpoints. + + After distribution, if multiple endpoints received **identical** values + for all their shared conflicting params it means the user gave a single + value (e.g. one date range) that the LLM duplicated across all namespaced + slots. In that case, keep the values only on the **first** such endpoint + and clear the rest, so the loop re-asks the user for each remaining + endpoint's params separately on the next turn. + + When the user provides **distinct** values for each endpoint's conflicting + params (e.g. two different date ranges) the values will differ, and this + method leaves all endpoints untouched. + + Args: + endpoint_states: Per-endpoint state objects (mutated in-place). + namespace_map: Mapping of ``namespaced_name → (ep_idx, original_name)`` + produced by ``_build_merged_schema()``. Empty when there are no + conflicting params. + + Returns: + ``True`` if any later endpoint had its duplicate params cleared; + ``False`` otherwise. The caller uses this flag to detect when the + LLM's already-generated clarifying question is stale (it said "none" + because it thought all params were satisfied) and needs regenerating. + """ + if not namespace_map: + return False + + # Build a map of endpoint_idx → set of original conflicting param names. + ep_conflicting: Dict[int, set[str]] = {} + for ep_idx, original_name in namespace_map.values(): + ep_conflicting.setdefault(ep_idx, set()).add(original_name) + + if len(ep_conflicting) < 2: + return False + + # Determine the shared conflicting param names across all participating endpoints. + shared_names: set[str] = set.intersection(*ep_conflicting.values()) + if not shared_names: + return False + + ep_indices = sorted(ep_conflicting.keys()) + first_ep_idx = ep_indices[0] + first_state = endpoint_states[first_ep_idx] + + # Collect the first endpoint's values for the shared params (only those + # that were actually filled this turn). Normalize to hashable strings to + # handle unhashable types (lists, dicts). + first_values: Dict[str, str] = { + name: self._normalize_value_to_hashable(first_state.collected_params[name]) + for name in shared_names + if name in first_state.collected_params + } + if not first_values: + # First endpoint has no values yet — nothing to compare. + return False + + # Clear later endpoints that received the same values as the first. + cleared_any = False + for ep_idx in ep_indices[1:]: + state = endpoint_states[ep_idx] + later_values: Dict[str, str] = { + name: self._normalize_value_to_hashable(state.collected_params[name]) + for name in shared_names + if name in state.collected_params + } + if not later_values: + continue + if later_values == first_values: + logger.debug( + f"MultiEndpointAgenticLoop: clearing duplicate conflicting params from endpoint" + f" | event_type=enforcement_sequential_conflict" + f" endpoint_name={state.endpoint.get('name', '')}" + ) + for name in shared_names: + state.collected_params.pop(name, None) + state.completed = False + cleared_any = True + return cleared_any + + @staticmethod + def _normalize_value_to_hashable(value: Any) -> str: # noqa: ANN401 + """Convert a value to a hashable string for safe comparison. + + Uses JSON serialization with `repr` fallback for non-JSON-serializable + types (lists, dicts, custom objects, etc.). Ensures that unhashable + values like lists and dicts can be safely compared in sets or as dict keys. + + Args: + value: Any value to normalize. + + Returns: + A string that uniquely represents the value. + """ + try: + # JSON serialization with sort_keys ensures consistent repr across identical + # value structures, using str as fallback for non-JSON-serializable objects. + return json.dumps(value, sort_keys=True, default=str) + except (TypeError, ValueError): + # Fallback to repr for types that can't be JSON-serialized. + return repr(value) + + @staticmethod + def _merged_collected( + endpoint_states: List[EndpointSessionState], + ) -> Dict[str, Any]: + """Return the union of all per-endpoint collected_params. + + Later endpoints overwrite earlier ones on key conflicts, consistent with + how the single-endpoint loop merges extraction results. + """ + merged: Dict[str, Any] = {} + for state in endpoint_states: + merged.update(state.collected_params) + return merged + + async def _save_session( + self, + chat_id: str, + endpoint_states: List[EndpointSessionState], + turn_count: int, + awaiting_continuation: bool = False, + ) -> None: + """Persist updated loop state to the Redis session store. + + Only updates fields owned by the loop (parallel_endpoints, turn_count, + awaiting_continuation). Workflow-owned fields are preserved. + A missing or unavailable session is logged but never raises. + + Args: + chat_id: Unique conversation identifier. + endpoint_states: Updated per-endpoint states to persist. + turn_count: Updated turn counter. + awaiting_continuation: Updated continuation flag. + """ + try: + if self._session_store is None: + logger.debug( + f"MultiEndpointAgenticLoop: session store unavailable — skipping save for chat_id={chat_id}" + ) + return + await self._session_store.update( + chat_id, + parallel_endpoints=endpoint_states, + turn_count=turn_count, + awaiting_continuation=awaiting_continuation, + ) + except Exception as exc: + logger.error( + f"MultiEndpointAgenticLoop: failed to save session | event_type=session_save_failed" + f" chat_id={chat_id} turn_count={turn_count} error_id={generate_error_id()} exc={exc}" + ) diff --git a/src/tool_classifier/multi_api_caller.py b/src/tool_classifier/multi_api_caller.py new file mode 100644 index 00000000..6674f334 --- /dev/null +++ b/src/tool_classifier/multi_api_caller.py @@ -0,0 +1,187 @@ +"""MultiAPICaller — concurrent batch execution of multiple API endpoints.""" + +import asyncio +import time +from typing import Any + +from src.loki_logger import LokiLogger +from src.utils.error_utils import generate_error_id + +from tool_classifier.api_caller import APICaller +from tool_classifier.constants import ( + MULTI_API_BATCH_TIMEOUT, + MULTI_API_PARTIAL_FAILURE_MESSAGES, +) +from tool_classifier.models import APICallResult, MultiAPICallResult +from llm_orchestrator_config.llm_ochestrator_constants import get_localized_message + +logger = LokiLogger(service_name="api-tool-calling") + + +class MultiAPICaller: + """ + Executes multiple API endpoint calls concurrently using :func:`asyncio.gather`. + + Reuses the caller's :class:`~tool_classifier.api_caller.APICaller` instance so + that per-URL circuit breaker state is shared across single and batch invocations. + + A batch-level timeout caps total wall-clock time. When the timeout fires, any + still-pending tasks are cancelled and their slots are filled with a failure result, + ensuring the caller always receives a fully populated :class:`MultiAPICallResult`. + + Args: + api_caller: Shared :class:`~tool_classifier.api_caller.APICaller` instance. + Circuit breaker state is inherited from this instance. + batch_timeout: Maximum seconds to wait for the entire batch. Defaults to + :data:`~tool_classifier.constants.MULTI_API_BATCH_TIMEOUT`. + """ + + def __init__( + self, + api_caller: APICaller, + batch_timeout: int = MULTI_API_BATCH_TIMEOUT, + ) -> None: + self._api_caller = api_caller + self._batch_timeout = batch_timeout + + async def call_all( + self, + endpoints: list[dict[str, Any]], + language: str = "et", + ) -> MultiAPICallResult: + """Execute all *endpoints* concurrently and return a consolidated result. + + Each endpoint dict must contain at minimum: + - ``"url"`` (str): Full URL for the request. + - ``"method"`` (str): HTTP method — ``"GET"`` or ``"POST"``. + - ``"call_params"`` (dict): Query / body parameters to send. + This key is intentionally distinct from the ``"params"`` key on + endpoint schema dicts (which holds a *list* of parameter + descriptors used during parameter collection), to prevent the + schema list from being forwarded to the HTTP call. + + Results are returned in the same order as *endpoints*. Failed tasks + (exceptions or batch timeout) produce ``APICallResult(success=False, …)``. + + Args: + endpoints: List of endpoint descriptor dicts. + language: BCP-47 language code for user-facing error messages. + + Returns: + :class:`~tool_classifier.models.MultiAPICallResult` with one result per + endpoint. + """ + if not endpoints: + return MultiAPICallResult(results=[], endpoints=[]) + + logger.info( + f"MultiAPICaller: batch started | event_type=multi_api_batch_started" + f" endpoint_count={len(endpoints)} batch_timeout={self._batch_timeout}" + ) + _t0 = time.time() + + tasks: list[asyncio.Task[APICallResult]] = [ + asyncio.create_task( + self._api_caller.call( + url=ep["url"], + method=ep["method"], + params=ep.get("call_params", {}), + language=language, + ), + name=f"multi_api_{i}", + ) + for i, ep in enumerate(endpoints) + ] + + results: list[APICallResult] + try: + raw = await asyncio.wait_for( + asyncio.gather(*tasks, return_exceptions=True), + timeout=self._batch_timeout, + ) + results = [self._coerce_result(item, language) for item in raw] + + except asyncio.TimeoutError: + pending = [t for t in tasks if not t.done()] + logger.warning( + f"MultiAPICaller: batch timeout | event_type=multi_api_batch_timeout" + f" pending_count={len(pending)} batch_timeout={self._batch_timeout}" + ) + for task in pending: + task.cancel() + # Await cancelled tasks so their coroutines can complete cleanup + # (e.g. close HTTP connections) before we return. + if pending: + await asyncio.gather(*pending, return_exceptions=True) + + results = [] + for task in tasks: + if task.done() and not task.cancelled(): + exc = task.exception() + outcome: APICallResult | BaseException = ( + exc if exc is not None else task.result() + ) + results.append(self._coerce_result(outcome, language)) + else: + results.append( + APICallResult( + success=False, + status_code=0, + response_data="", + error=get_localized_message( + MULTI_API_PARTIAL_FAILURE_MESSAGES, language + ), + ) + ) + + had_failure = any(not r.success for r in results) + if had_failure: + logger.warning( + f"MultiAPICaller: partial failure | event_type=multi_api_partial_failure" + f" failed_count={sum(not r.success for r in results)} total_count={len(results)}" + ) + + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.info( + f"MultiAPICaller: batch completed | event_type=multi_api_batch_completed" + f" endpoint_count={len(endpoints)}" + f" success_count={sum(r.success for r in results)}" + f" failed_count={sum(not r.success for r in results)}" + f" duration_ms={_duration_ms}" + ) + + return MultiAPICallResult(results=results, endpoints=endpoints) + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + def _coerce_result( + self, + item: APICallResult | BaseException | None, + language: str, + ) -> APICallResult: + """Convert a raw gather() output item to an :class:`APICallResult`. + + :func:`asyncio.gather` with ``return_exceptions=True`` returns either the + coroutine's return value or the exception that was raised. We identify + exceptions by checking against :class:`BaseException` rather than the + positive type, which avoids issues with Python's dual import paths + (``tool_classifier.models`` vs ``src.tool_classifier.models``). + """ + if not isinstance(item, BaseException): + # Assume any non-exception value is already an APICallResult. + return item # type: ignore[return-value] + error_id = generate_error_id() + exc_type = type(item).__name__ + exc_msg = str(item) + logger.error( + f"MultiAPICaller: task exception | event_type=multi_api_task_exception" + f" error_id={error_id} exc_type={exc_type} exc_msg={exc_msg!r}" + ) + return APICallResult( + success=False, + status_code=0, + response_data="", + error=get_localized_message(MULTI_API_PARTIAL_FAILURE_MESSAGES, language), + ) diff --git a/src/tool_classifier/multi_response_formatter.py b/src/tool_classifier/multi_response_formatter.py new file mode 100644 index 00000000..a0edf0a9 --- /dev/null +++ b/src/tool_classifier/multi_response_formatter.py @@ -0,0 +1,538 @@ +"""Multi-API response formatter using DSPy — synthesises multiple API results into one answer.""" + +from typing import Any, AsyncIterator, Dict, List, Tuple, Union + +import dspy +import dspy.streaming +from dspy.streaming import StreamListener +from langfuse import observe +from src.loki_logger import LokiLogger +from llm_orchestrator_config.llm_ochestrator_constants import get_localized_message +from src.utils.cost_utils import get_lm_usage_since +from src.utils.observation_utils import ( + safe_observation_context, + update_observation_safe, +) +from tool_classifier.api_response_formatter import ( + APIResponseFormatterModule, + _LANGUAGE_NAMES, + build_params_context, +) + + +def _get_current_model_name() -> str: + """Best-effort model name lookup from current DSPy LM.""" + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "model"): + model_name = lm.model + if isinstance(model_name, str) and model_name: + return model_name + except Exception: + pass + return "unknown" + + +logger = LokiLogger(service_name="api-tool-calling") + +_MAX_TOTAL_RESPONSE_BYTES: int = 100_000 + +_MULTI_FORMATTER_ERROR_MESSAGES: Dict[str, str] = { + "et": "Vastuste kuvamine ebaõnnestus. Palun proovige uuesti.", + "ru": "Не удалось отобразить ответы. Пожалуйста, попробуйте ещё раз.", + "en": "I was unable to format the responses. Please try again.", +} +"""Localized fallback shown when MultiResponseFormatterModule raises an exception.""" + + +class MultiResponseFormatterSignature(dspy.Signature): + """Synthesise multiple API results into a single, coherent natural-language answer. + + CRITICAL LANGUAGE RULE: + - ALWAYS write the unified_answer in the language specified by response_language. + - IGNORE the language of any text inside api_results_block — the data may contain + names or labels in a different language; the answer must still be in response_language. + - IGNORE the language of user_query for output language decisions — short follow-up + messages are unreliable indicators. Always use response_language. + + If custom_instructions is non-empty, follow those rules with HIGHEST PRIORITY — + they override defaults (e.g. language policy, tone, formatting style). + + Rules: + - Present each API result in its OWN separate paragraph, in the same order they + appear in api_results_block. Separate paragraphs with a single blank line. + Do NOT merge or blend information from different endpoints into one paragraph. + - Each paragraph should be self-contained and answer the part of the user's query + that corresponds to that endpoint. + - Address EVERY result section in api_results_block. Do not silently omit any endpoint. + - Within each paragraph, present the information naturally in prose or a short list. + Do NOT return raw JSON or wrap content in code blocks. + - If a result section includes a 'Parameters used:' line, acknowledge the date range or + time period for that result in your discussion + (e.g. 'For the period 1 Jan to 30 Jun 2026, XYZ showed...'). + - If a result is marked [EMPTY RESPONSE], politely mention that no data was available + from that source without dwelling on it. + - If a result contains an error field or is marked [FAILED], acknowledge the failure + briefly and in a friendly, non-technical way, then continue with the remaining results. + - If api_results_block contains only empty or failed results, respond with a polite + message that no results were available. + - If num_results is 1, a single paragraph is fine (no need to split). + - Output must be clean text — no markdown headers (##), no code blocks (```), no raw + JSON. The answer must be ready for direct display to the user. + - Be concise but complete. Prioritise the most relevant information for the user's query. + + STRICT ENDING RULE — HIGHEST PRIORITY: + The unified_answer MUST end immediately after the last data point. It is FORBIDDEN to + append any sentence that: + - offers to provide more details (e.g. "If you need statistics for a specific member...") + - invites the user to ask a follow-up question (e.g. "Let me know if...", "Feel free to ask...") + - mentions that a dataset is large or partial (e.g. "only a sample is shown here") + - suggests the user can specify a name, party, or other filter + The very last character of unified_answer must be part of the actual data, not a helper offer. + """ + + user_query: str = dspy.InputField( + desc="The user's original question or request, in Estonian, Russian, or English" + ) + api_results_block: str = dspy.InputField( + desc=( + "A labeled text block containing one section per API result. " + "Each section is headed by the endpoint name and description, " + "followed by the raw response string. " + "May contain EMPTY RESPONSE or FAILED markers for individual results." + ) + ) + response_language: str = dspy.InputField( + desc=( + "The language to write the answer in, detected from the user's first message: " + "'English', 'Estonian', or 'Russian'. " + "Always use this — do not infer language from api_results_block content." + ) + ) + custom_instructions: str = dspy.InputField( + desc=( + "Optional system-level instructions configured by the organisation " + "(e.g. 'Always respond in Estonian', 'Use structured format'). " + "Empty string when no custom config is active. " + "When non-empty, follow these rules with highest priority." + ) + ) + num_results: str = dspy.InputField( + desc=( + "The total number of API results included in api_results_block, " + "as a plain integer string (e.g. '3'). " + "Use this to verify all results are addressed." + ) + ) + + unified_answer: str = dspy.OutputField( + desc=( + "One paragraph per API result, separated by blank lines, in the same order " + "as api_results_block. Each paragraph is self-contained and covers exactly " + "one endpoint's data. Written entirely in the language specified by " + "response_language. No raw JSON, no code blocks, no markdown headers. " + "MUST end after the last data point. " + "FORBIDDEN: any closing sentence offering more help, inviting follow-up questions, " + "mentioning that the dataset is partial, or suggesting the user specify a name/party." + ) + ) + + +class MultiResponseFormatterModule(dspy.Module): + """DSPy Module that synthesises multiple API results into one natural-language answer.""" + + def __init__(self, custom_instructions: str = "") -> None: + """Initialise formatter with a direct DSPy Predict. + + Args: + custom_instructions: Optional organisation-level prompt rules (e.g. language + policy). Passed verbatim to the DSPy predictor on every call. Defaults + to empty string (no custom config). + """ + super().__init__() + self.formatter = dspy.Predict(MultiResponseFormatterSignature) + self._custom_instructions = custom_instructions + + @observe(name="multi_api_response_formatting_llm", as_type="generation") + def forward( + self, + user_query: str, + api_results: List[ + Tuple[str, str, Union[str, Dict[str, Any], List[Any]], Dict[str, Any]] + ], + detected_language: str = "en", + ) -> str: + """Synthesise multiple API results into a single natural-language answer. + + Args: + user_query: The user's original question. + api_results: A list of + ``(endpoint_name, endpoint_description, api_response, collected_params)`` + 4-tuples. ``api_response`` may be a JSON string, dict, or list. + ``collected_params`` is used to generate a date-range + acknowledgment inside each result section. + detected_language: ISO language code from the agentic loop session + ('en', 'et', 'ru'). Defaults to 'en'. This is the authoritative + language for the answer. + + Returns: + A clean, unified natural-language answer ready for display to the user. + """ + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning( + f"Failed to get LM history length for multi response formatting: {e}" + ) + + try: + results_block = self._build_results_block(api_results) + response_language = _LANGUAGE_NAMES.get(detected_language, "English") + + result = self.formatter( + user_query=user_query, + api_results_block=results_block, + response_language=response_language, + custom_instructions=self._custom_instructions, + num_results=str(len(api_results)), + ) + unified_answer = result.unified_answer + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_query": user_query, + "num_results": len(api_results), + "response_language": response_language, + }, + output_data={ + "unified_answer_preview": str(unified_answer)[:500], + }, + metadata={ + "model": _get_current_model_name(), + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": False, + }, + ) + return unified_answer # type: ignore[no-any-return] + + except Exception as e: + logger.error( + f"MultiResponseFormatterModule.forward failed: {e}", exc_info=True + ) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_query": user_query, + "num_results": len(api_results), + "detected_language": detected_language, + }, + output_data={"error": str(e)}, + metadata={ + "model": _get_current_model_name(), + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": False, + }, + ) + safe_language = ( + detected_language + if detected_language in _MULTI_FORMATTER_ERROR_MESSAGES + else "en" + ) + return get_localized_message(_MULTI_FORMATTER_ERROR_MESSAGES, safe_language) + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + async def stream_forward_multi( + self, + user_query: str, + api_results: List[ + Tuple[str, str, Union[str, Dict[str, Any], List[Any]], Dict[str, Any]] + ], + detected_language: str = "en", + ) -> AsyncIterator[str]: + """Stream unified_answer tokens using DSPy native streaming. + + API results are pre-resolved; only the LLM synthesis step is streamed. + Yields individual token strings as they arrive from the LLM. + + Fallback chain: + 1. DSPy ``StreamResponse`` tokens (true token-by-token streaming) + 2. Final ``dspy.Prediction.unified_answer`` (if streamify yields no tokens) + 3. Blocking ``forward()`` call (if no Prediction was received) + 4. Localized error message on any exception. + + Args: + user_query: The user's original question. + api_results: A list of + ``(endpoint_name, endpoint_description, api_response, collected_params)`` + 4-tuples. ``api_response`` may be a JSON string, dict, or list. + detected_language: ISO code ('en', 'et', 'ru'). Defaults to 'en'. + + Yields: + Token strings from the LLM ``unified_answer`` field. + """ + safe_language = ( + detected_language + if detected_language in _MULTI_FORMATTER_ERROR_MESSAGES + else "en" + ) + output_stream = None + + with safe_observation_context( + as_type="generation", + name="multi_api_response_formatting_streaming", + input={ + "user_query": user_query[:500], + "num_results": len(api_results), + "detected_language": detected_language, + }, + ) as generation: + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning( + f"Failed to get LM history length for multi response streaming: {e}" + ) + + try: + results_block = self._build_results_block(api_results) + response_language = _LANGUAGE_NAMES.get(detected_language, "English") + + # A fresh wrapper is created on every call because dspy.configure(lm=...) + # is called per request. A cached wrapper retains a stale LM reference and + # yields a bare dspy.Prediction instead of StreamResponse tokens. + logger.debug( + "MultiResponseFormatterModule: creating fresh streamify wrapper " + "for unified_answer field" + ) + listener = StreamListener(signature_field_name="unified_answer") + stream_predictor: Any = dspy.streamify( + self.formatter, stream_listeners=[listener] + ) + output_stream = stream_predictor( + user_query=user_query, + api_results_block=results_block, + response_language=response_language, + custom_instructions=self._custom_instructions, + num_results=str(len(api_results)), + ) + + stream_started = False + token_count = 0 + accumulated: list[str] = [] + final_prediction: dspy.Prediction | None = None + async for chunk in output_stream: + if isinstance(chunk, dspy.streaming.StreamResponse): + if chunk.signature_field_name == "unified_answer": + stream_started = True + token_count += 1 + accumulated.append(chunk.chunk) + yield chunk.chunk + elif isinstance(chunk, dspy.Prediction): + final_prediction = chunk + # dspy.streamify did not stream individual tokens — yield the + # full answer from the final Prediction as a single frame. + if not stream_started: + answer = getattr(chunk, "unified_answer", None) + if answer: + logger.info( + "MultiResponseFormatterModule.stream_forward_multi: " + "no StreamResponse tokens — yielding full Prediction answer" + ) + stream_started = True + accumulated.append(answer) + yield answer + + assembled_answer = "".join(accumulated) + + if stream_started and token_count > 0: + logger.debug( + f"MultiResponseFormatterModule.stream_forward_multi: " + f"streamed {token_count} tokens" + ) + # DSPy streaming can drop the last few tokens before EOS. + # The final dspy.Prediction holds the authoritative complete answer. + # Yield any tail that wasn't delivered as StreamResponse chunks. + if final_prediction is not None: + full_answer = getattr(final_prediction, "unified_answer", None) + if full_answer: + streamed_text = assembled_answer + if full_answer.startswith(streamed_text) and len( + full_answer + ) > len(streamed_text): + tail = full_answer[len(streamed_text) :] + if tail.strip(): + logger.debug( + f"MultiResponseFormatterModule.stream_forward_multi: " + f"yielding {len(tail)} missing tail chars from Prediction" + ) + assembled_answer += tail + yield tail + elif streamed_text and not full_answer.startswith( + streamed_text + ): + logger.warning( + "MultiResponseFormatterModule.stream_forward_multi: " + "streamed output is not a prefix of final Prediction; " + "skipping tail reconciliation" + ) + + if not stream_started: + # Last-resort fallback: blocking forward() — covers cases where + # dspy.streamify yields neither StreamResponse nor Prediction. + logger.warning( + "MultiResponseFormatterModule.stream_forward_multi: " + "streamify produced no tokens and no Prediction — using blocking forward()" + ) + result = self.forward( + user_query=user_query, + api_results=api_results, + detected_language=detected_language, + ) + assembled_answer = result + yield result + + usage = get_lm_usage_since(history_length_before) + try: + generation.update( + input={ + "user_query": user_query, + "num_results": len(api_results), + "detected_language": detected_language, + }, + output=assembled_answer, + usage_details={ + "input": usage.get("total_prompt_tokens", 0), + "output": usage.get("total_completion_tokens", 0), + "total": usage.get("total_tokens", 0), + }, + cost_details={ + "total": usage.get("total_cost", 0.0), + }, + metadata={ + "stream_started": stream_started, + "chunk_count": token_count, + "streaming": True, + }, + ) + except Exception as update_error: + logger.debug( + f"Langfuse generation update skipped for multi response streaming: {update_error}" + ) + + except Exception as e: + logger.error( + f"MultiResponseFormatterModule.stream_forward_multi failed: {e}", + exc_info=True, + ) + usage = get_lm_usage_since(history_length_before) + try: + generation.update( + input={ + "user_query": user_query, + "num_results": len(api_results), + "detected_language": detected_language, + }, + output={"error": str(e)}, + usage_details={ + "input": usage.get("total_prompt_tokens", 0), + "output": usage.get("total_completion_tokens", 0), + "total": usage.get("total_tokens", 0), + }, + cost_details={ + "total": usage.get("total_cost", 0.0), + }, + metadata={"streaming": True}, + ) + except Exception as update_error: + logger.debug( + f"Langfuse error update skipped for multi response streaming: {update_error}" + ) + yield get_localized_message( + _MULTI_FORMATTER_ERROR_MESSAGES, safe_language + ) + finally: + if output_stream is not None: + try: + await output_stream.aclose() + except Exception as cleanup_error: + logger.debug(f"Error during stream cleanup: {cleanup_error}") + + @staticmethod + def _build_results_block( + api_results: List[ + Tuple[str, str, Union[str, Dict[str, Any], List[Any]], Dict[str, Any]] + ], + ) -> str: + """Serialise a list of API results into a labeled text block for the LLM. + + Each result becomes a section headed by the endpoint name and description, + optionally followed by a ``Parameters used: ...`` line when date-typed + values are present in ``collected_params``, then the normalised, annotated, + and truncated response string. + A combined size guard caps the total block at ``_MAX_TOTAL_RESPONSE_BYTES`` + to prevent LLM context overflow. + + Args: + api_results: List of + ``(endpoint_name, endpoint_description, api_response, collected_params)`` + 4-tuples. + + Returns: + A multiline string with one clearly delimited section per result. + """ + if not api_results: + return "[NO RESULTS: No API results were provided]" + + sections: List[str] = [] + total_bytes = 0 + + for idx, (name, description, raw_response, collected_params) in enumerate( + api_results, start=1 + ): + normalized = APIResponseFormatterModule._normalize_response(raw_response) + normalized = APIResponseFormatterModule._annotate_empty(normalized) + normalized = APIResponseFormatterModule._truncate_if_needed(normalized) + + params_context = build_params_context(collected_params) + params_line = ( + f"Parameters used: {params_context}\n" if params_context else "" + ) + + section = ( + f"--- Result {idx}: {name} ---\n" + f"Description: {description}\n" + f"{params_line}" + f"Response:\n{normalized}\n" + ) + section_bytes = len(section.encode("utf-8")) + + if total_bytes + section_bytes > _MAX_TOTAL_RESPONSE_BYTES: + remaining = _MAX_TOTAL_RESPONSE_BYTES - total_bytes + if remaining > 0: + suffix = ( + "\n[NOTE: Combined results truncated due to total size limit]" + ) + suffix_bytes = len(suffix.encode("utf-8")) + encoded = section.encode("utf-8") + truncated = encoded[: max(0, remaining - suffix_bytes)].decode( + "utf-8", errors="ignore" + ) + sections.append(truncated + suffix) + total_bytes = _MAX_TOTAL_RESPONSE_BYTES + break + + sections.append(section) + total_bytes += section_bytes + + return "\n".join(sections) diff --git a/src/tool_classifier/param_extractor.py b/src/tool_classifier/param_extractor.py index 7e7fa945..a3c0cbe6 100644 --- a/src/tool_classifier/param_extractor.py +++ b/src/tool_classifier/param_extractor.py @@ -3,19 +3,27 @@ import asyncio import json import re +import time from datetime import datetime, timezone from typing import Any, Dict, List, Optional, TypedDict import dspy import dspy.streaming from dspy.streaming import StreamListener -from loguru import logger +from langfuse import observe +from src.utils.observation_utils import update_observation_safe +from src.loki_logger import LokiLogger +from src.utils.error_utils import generate_error_id + +from src.utils.cost_utils import get_lm_usage_since _TRUTHY_STRINGS = {"true", "yes", "jah", "1", "on", "õige", "да"} _FALSY_STRINGS = {"false", "no", "ei", "0", "off", "vale", "нет"} _MAX_HISTORY_TURNS = 5 +logger = LokiLogger(service_name="api-tool-calling") + # Regex patterns to strip format hints from parameter descriptions before # they are fed to the question-generation prompt. This prevents the LLM # from including format instructions (e.g. "YYYY-MM-DD") in its questions. @@ -30,7 +38,7 @@ ] -def _strip_format_hints(description: str) -> str: +def strip_format_hints(description: str) -> str: """Remove format hints from a parameter description. Strips patterns such as ``(YYYY-MM-DD)``, ``(ISO 8601)``, @@ -92,10 +100,36 @@ class ParamExtractionSignature(dspy.Signature): - After extraction, check whether ALL required params are now satisfied (i.e., present in already_collected OR just extracted). - If ALL required params are satisfied, return the literal string "none". - - If ONE OR MORE required params are still missing, generate ONE friendly question - that asks for ALL of those remaining missing params at once. - On the first turn this may cover many params; on follow-up turns it narrows - to only the params the user has not yet provided. + Do NOT add any acknowledgment or thank-you — just return "none". + - If ONE OR MORE required params are still missing, generate ONE friendly message + that first acknowledges what was just collected (if anything) then asks for + all remaining missing params in the same message. + + Acknowledgment rules (when params are still missing): + - CASE A — turn_count is not '0' AND extracted_params is non-empty + (the user was responding to a follow-up question and gave something): + Prepend a brief natural acknowledgment of the values just provided, then + ask for the remaining missing params. + Use the parameter's description field to describe what was received, NOT + the raw camelCase name (e.g. "Thank you for providing the group" not + "Thank you for providing group"). + Keep it concise: two sentences — one acknowledgment sentence, one question sentence. + Join the two clauses with a period and a space, never with a dash, em dash, + or any other punctuation connector. + VARY the acknowledgment opener each turn — do NOT repeat the same phrase. + Use different natural openers across turns, for example: + English: "Got it.", "Perfect.", "Got the address.", "Noted.", + "Thanks for that.", "Great.", "Received." + Estonian: "Selge.", "Hästi.", "Aitäh.", "Märgitud.", "Tubli." + Russian: "Принято.", "Понял.", "Отлично.", "Хорошо.", "Записал." + Match the tone and language to session_language. + - CASE B — turn_count is '0' (this is the user's original query): + Do NOT add any acknowledgment or thank-you, even if params were extracted + from the original message. Simply ask for the remaining missing params. + The user was not responding to a question — they just asked naturally. + - CASE C — turn_count is not '0' AND extracted_params is empty (the user did NOT answer what was asked): + Do NOT add any acknowledgment or thank-you. Ask for the missing params + again directly, rephrasing if helpful. - Use each missing parameter's description field to phrase the question naturally (e.g., "Which country and date would you like to use?" not "Provide countryIsoCode and startDate") - Never expose raw parameter names (camelCase identifiers) to the user @@ -104,6 +138,11 @@ class ParamExtractionSignature(dspy.Signature): "in the format...") in the question — only ask WHAT information is needed, not HOW it should be formatted. The system handles format conversion internally from any natural-language input the user provides. + - When intent_groups is a non-empty JSON array with 2 or more groups, structure + the question with a clear natural-language separation per group — e.g., + "Which address or place would you like to search for, and what is the unique + identifier of the initiative you want to check?" Use conjunctions like "and" + to link the groups naturally into one fluent sentence. """ user_message: str = dspy.InputField( @@ -130,6 +169,13 @@ class ParamExtractionSignature(dspy.Signature): "still extract the new value — corrections are allowed." ) ) + turn_count: str = dspy.InputField( + desc=( + "Zero-based turn index as a string. '0' means this is the user's original query " + "(not a response to a follow-up question). '1' or higher means the user is " + "responding to a clarifying question asked by the system." + ) + ) custom_instructions: str = dspy.InputField( desc=( "Optional system-level instructions configured by the organisation " @@ -138,6 +184,18 @@ class ParamExtractionSignature(dspy.Signature): "When non-empty, follow these rules with highest priority for the clarifying_question." ) ) + intent_groups: str = dspy.InputField( + desc=( + "Optional JSON array grouping missing required params by their owning intent/endpoint: " + '[{"intent": "", "missing_param_descriptions": ["desc1", ...]}, ...]. ' + "Empty JSON array [] when not in multi-intent mode or when no intents have " + "missing params. " + "When 1 group is present, use the intent name to add context to the question " + "(e.g., 'For the electricity market prices, what is the start and end date?'). " + "When 2+ groups are present, use this to phrase a single question " + "that clearly separates the needs of each intent with natural conjunctions." + ) + ) extracted_params: str = dspy.OutputField( desc='Valid JSON object of newly extracted parameters only: {"param_name": value}. Empty object {} if nothing new found.' @@ -147,10 +205,22 @@ class ParamExtractionSignature(dspy.Signature): ) clarifying_question: str = dspy.OutputField( desc=( - "A single natural-language question that asks for ALL missing parameters " - 'at once, or the literal string "none" if all required params are collected. ' + "One natural-language message in session_language. " + 'Return the literal string "none" if all required params are collected (no acknowledgment). ' + "If params are still missing and turn_count is '0' (original query): " + "skip any acknowledgment and ask for the missing params directly. " + "If params are still missing and turn_count > 0 and extracted_params is non-empty: " + "briefly acknowledge the values just provided (using natural descriptions, not raw param names), " + "then ask for all remaining missing params in the same message. " + "If params are still missing and turn_count > 0 and extracted_params is empty: " + "skip any acknowledgment and ask for the missing params directly. " 'Never include format instructions or examples (e.g. "YYYY-MM-DD", ' - '"ISO 8601", "2-letter code") — only ask what information is needed.' + '"ISO 8601", "2-letter code") — only ask what information is needed. ' + "When intent_groups contains 1 group, prefix the question with the intent " + "name for context (e.g. 'For the electricity market prices, what is the " + "start and end date?'). " + "When intent_groups contains 2+ groups, the question must use natural conjunctions " + "(e.g. 'and') to separate the missing params of each group into one fluent sentence." ) ) @@ -170,6 +240,7 @@ def __init__(self, custom_instructions: str = "") -> None: self.extractor = dspy.Predict(ParamExtractionSignature) self._custom_instructions = custom_instructions + @observe(name="api_param_extraction_llm", as_type="generation") def forward( self, user_message: str, @@ -177,6 +248,8 @@ def forward( conversation_history: Optional[List[Dict[str, Any]]] = None, already_collected: Optional[Dict[str, Any]] = None, session_language: str = "en", + turn_count: int = 0, + intent_groups: Optional[List[Dict[str, Any]]] = None, ) -> ParamExtractionResult: """ Extract parameter values from user message and conversation history. @@ -188,6 +261,15 @@ def forward( already_collected: Parameter values collected in prior turns (optional) session_language: Language code detected on turn 0 ('en', 'et', 'ru'). All clarifying questions will be generated in this language. + intent_groups: Optional list of intent-group dicts, each with an + ``"intent"`` key (endpoint name) and a + ``"missing_param_descriptions"`` key (list of human-readable + descriptions for the missing params of that intent). Pass this + when a single user query maps to multiple intents that each have + outstanding required parameters; the extractor will then phrase + one clarifying question that addresses every group's missing + params using natural conjunctions. Pass ``None`` (default) or + an empty list when operating in single-intent mode. Returns: ParamExtractionResult with extracted_params, missing_required, clarifying_question @@ -196,15 +278,27 @@ def forward( history_text = self._format_conversation_history(conversation_history) sanitized_schema = [ - {**p, "description": _strip_format_hints(p.get("description", ""))} + {**p, "description": strip_format_hints(p.get("description", ""))} if isinstance(p, dict) else p for p in params_schema ] params_schema_json = json.dumps(sanitized_schema, ensure_ascii=False) already_collected_json = json.dumps(already_collected, ensure_ascii=False) + intent_groups_json = json.dumps(intent_groups or [], ensure_ascii=False) + + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning( + f"Failed to get LM history length for parameter extraction: {e}" + ) result = None + _t0 = time.time() try: result = self.extractor( user_message=user_message, @@ -212,24 +306,92 @@ def forward( session_language=session_language, params_schema=params_schema_json, already_collected=already_collected_json, + turn_count=str(turn_count), custom_instructions=self._custom_instructions, + intent_groups=intent_groups_json, ) - return self._parse_prediction(result, params_schema, already_collected) - + _duration_ms = round((time.time() - _t0) * 1000, 1) + logger.debug( + f"ParamExtractionModule: LLM extraction complete" + f" | event_type=param_extraction_llm_complete" + f" turn_count={turn_count} duration_ms={_duration_ms}" + ) + parsed = self._parse_prediction(result, params_schema, already_collected) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_message": user_message, + "params_schema_count": len(params_schema), + }, + output_data={ + "missing_required_count": len(parsed["missing_required"]), + "clarifying_question": parsed["clarifying_question"], + }, + metadata={ + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "extracted_params_count": len(parsed["extracted_params"]), + "streaming": False, + }, + ) + return parsed except json.JSONDecodeError as e: - logger.error(f"Failed to parse param extraction JSON: {e}") - if result: - logger.error( - f"Raw extracted_params: {getattr(result, 'extracted_params', None)}" - ) - logger.error( - f"Raw missing_required: {getattr(result, 'missing_required', None)}" - ) - return self._safe_defaults(params_schema, already_collected) + _raw_ep = getattr(result, "extracted_params", None) + _raw_mr = getattr(result, "missing_required", None) + logger.error( + f"ParamExtractionModule: JSON parse error in forward" + f" | event_type=param_extraction_json_error" + f" error_id={generate_error_id()}" + f" exc_type={type(e).__name__} exc_msg={e!r}" + f" raw_extracted_params={_raw_ep!r}" + f" raw_missing_required={_raw_mr!r}" + ) + fallback = self._safe_defaults(params_schema, already_collected) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_message": user_message, + "params_schema_count": len(params_schema), + }, + output_data={ + "missing_required_count": len(fallback["missing_required"]), + "clarifying_question": fallback["clarifying_question"], + "error": f"JSON parse error: {e}", + }, + metadata={ + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": False, + }, + ) + return fallback except Exception as e: - logger.exception(f"Param extraction forward failed: {e}") - return self._safe_defaults(params_schema, already_collected) + logger.error( + f"ParamExtractionModule: forward failed" + f" | event_type=param_extraction_failed" + f" error_id={generate_error_id()}" + f" exc_type={type(e).__name__} exc_msg={e!r}" + ) + fallback = self._safe_defaults(params_schema, already_collected) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_message": user_message, + "params_schema_count": len(params_schema), + }, + output_data={ + "missing_required_count": len(fallback["missing_required"]), + "clarifying_question": fallback["clarifying_question"], + "error": str(e), + }, + metadata={ + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": False, + }, + ) + return fallback def _get_stream_predictor(self) -> Any: """Return a fresh streamified predictor for each call. @@ -245,6 +407,7 @@ def _get_stream_predictor(self) -> Any: listener = StreamListener(signature_field_name="clarifying_question") return dspy.streamify(self.extractor, stream_listeners=[listener]) + @observe(name="api_param_extraction_streaming", as_type="generation") async def stream_forward( self, user_message: str, @@ -252,6 +415,8 @@ async def stream_forward( conversation_history: Optional[List[Dict[str, Any]]] = None, already_collected: Optional[Dict[str, Any]] = None, session_language: str = "en", + turn_count: int = 0, + intent_groups: Optional[List[Dict[str, Any]]] = None, ) -> tuple[List[str], ParamExtractionResult]: """Stream clarifying_question tokens while returning the full extraction result. @@ -270,6 +435,15 @@ async def stream_forward( conversation_history: Recent conversation messages (optional). already_collected: Parameter values from prior turns (optional). session_language: Language code (``'en'``, ``'et'``, ``'ru'``). + intent_groups: Optional list of intent-group dicts, each with an + ``"intent"`` key (endpoint name) and a + ``"missing_param_descriptions"`` key (list of human-readable + descriptions for the missing params of that intent). Pass this + when a single user query maps to multiple intents that each have + outstanding required parameters; the extractor will then phrase + one clarifying question that addresses every group's missing + params using natural conjunctions, streamed token by token. + Pass ``None`` (default) or an empty list for single-intent mode. Returns: Tuple of ``(question_tokens, extraction_result)``. @@ -279,15 +453,28 @@ async def stream_forward( history_text = self._format_conversation_history(conversation_history) sanitized_schema = [ - {**p, "description": _strip_format_hints(p.get("description", ""))} + {**p, "description": strip_format_hints(p.get("description", ""))} if isinstance(p, dict) else p for p in params_schema ] params_schema_json = json.dumps(sanitized_schema, ensure_ascii=False) already_collected_json = json.dumps(already_collected, ensure_ascii=False) + intent_groups_json = json.dumps(intent_groups or [], ensure_ascii=False) output_stream = None + + history_length_before = 0 + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + history_length_before = len(lm.history) + except Exception as e: + logger.warning( + "Failed to get LM history length for parameter extraction " + f"streaming: {e}" + ) + try: stream_predictor = self._get_stream_predictor() output_stream = stream_predictor( @@ -296,7 +483,9 @@ async def stream_forward( session_language=session_language, params_schema=params_schema_json, already_collected=already_collected_json, + turn_count=str(turn_count), custom_instructions=self._custom_instructions, + intent_groups=intent_groups_json, ) tokens: List[str] = [] @@ -311,8 +500,8 @@ async def stream_forward( if prediction is None: logger.warning( - "ParamExtractionModule.stream_forward: no Prediction received — " - "falling back to blocking forward()" + "ParamExtractionModule: no Prediction received" + " | event_type=stream_no_prediction_received" ) result = await asyncio.to_thread( self.forward, @@ -321,6 +510,8 @@ async def stream_forward( conversation_history, already_collected, session_language, + turn_count, + intent_groups, ) fallback_token = result["clarifying_question"] return ( @@ -338,28 +529,99 @@ async def stream_forward( if tokens: logger.debug( - f"ParamExtractionModule.stream_forward: streamed {len(tokens)} tokens" + f"ParamExtractionModule: stream tokens collected" + f" | event_type=stream_tokens_collected token_count={len(tokens)}" ) + # Join all streamed token chunks into the complete assembled question for langfuse logging. + assembled_question = ( + "".join(tokens) if tokens else result["clarifying_question"] + ) + + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_message": user_message, + "params_schema_count": len(params_schema), + }, + output_data={ + "missing_required_count": len(result["missing_required"]), + "clarifying_question": assembled_question, + }, + metadata={ + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "extracted_params_count": len(result["extracted_params"]), + "streaming": True, + "question_token_count": len(tokens), + }, + ) + return tokens, result except json.JSONDecodeError as e: logger.error( - f"ParamExtractionModule.stream_forward failed to parse JSON: {e}" + f"ParamExtractionModule: JSON parse error in stream_forward" + f" | event_type=stream_json_error" + f" error_id={generate_error_id()}" + f" exc_type={type(e).__name__} exc_msg={e!r}" + ) + fallback = self._safe_defaults(params_schema, already_collected) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_message": user_message, + "params_schema_count": len(params_schema), + }, + output_data={ + "missing_required_count": len(fallback["missing_required"]), + "clarifying_question": fallback["clarifying_question"], + "error": f"JSON parse error: {e}", + }, + metadata={ + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": True, + }, ) - return [], self._safe_defaults(params_schema, already_collected) + return [], fallback except Exception as e: - logger.exception(f"ParamExtractionModule.stream_forward failed: {e}") - return [], self._safe_defaults(params_schema, already_collected) + logger.error( + f"ParamExtractionModule: stream_forward failed" + f" | event_type=stream_extraction_failed" + f" error_id={generate_error_id()}" + f" exc_type={type(e).__name__} exc_msg={e!r}" + ) + fallback = self._safe_defaults(params_schema, already_collected) + usage = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "user_message": user_message, + "params_schema_count": len(params_schema), + }, + output_data={ + "missing_required_count": len(fallback["missing_required"]), + "clarifying_question": fallback["clarifying_question"], + "error": str(e), + }, + metadata={ + "usage": usage, + "num_calls": usage.get("num_calls", 0), + "streaming": True, + }, + ) + return [], fallback finally: if output_stream is not None: try: await output_stream.aclose() - except Exception as cleanup_error: + except Exception as e: logger.debug( - f"Error during param extraction stream cleanup: {cleanup_error}" + f"ParamExtractionModule: stream cleanup error" + f" | event_type=stream_cleanup_error" + f" exc_type={type(e).__name__} exc_msg={e!r}" ) # ------------------------------------------------------------------ @@ -434,7 +696,10 @@ def _validate_param_type( return False, value # Unknown type — accept as string to avoid silent data loss - logger.warning(f"Unknown param type '{param_type}'; accepting as string") + logger.warning( + f"ParamExtractionModule: unknown param type" + f" | event_type=param_type_unknown param_type={param_type!r}" + ) return True, str_value def _format_conversation_history( @@ -532,36 +797,13 @@ def _parse_prediction( validated_params[param_name] = coerced else: logger.warning( - f"Extracted value for '{param_name}' failed type validation " - f"(expected {param_type}, got {raw_value!r})" + f"ParamExtractionModule: param value type mismatch" + f" | event_type=param_value_type_mismatch" + f" param_name={param_name!r} expected_type={param_type!r}" + f" raw_value={raw_value!r}" ) type_invalid_params.append(param_name) - # SINGLE-VALUE REASSIGNMENT: if the LLM assigned a value to a later same-type - # param while an earlier same-type param is still missing, move the value forward. - # This fixes the common case where a lone date like "2026-04-01" is extracted as - # endDate when startDate is still missing. - combined_after_extraction = {**already_collected, **validated_params} - required_schema_order = [ - p for p in params_schema if isinstance(p, dict) and p.get("required", False) - ] - for idx, missing_entry in enumerate(required_schema_order): - m_name = missing_entry["name"] - m_type = missing_entry.get("type", "string") - if m_name in combined_after_extraction: - continue # already satisfied - # Find the first later param with the same type that was just extracted - for later_entry in required_schema_order[idx + 1 :]: - l_name = later_entry["name"] - l_type = later_entry.get("type", "string") - if l_type == m_type and l_name in validated_params: - logger.debug( - f"ParamExtractor: reassigning '{l_name}' → '{m_name}' " - f"(single {m_type} value assigned to wrong param by LLM)" - ) - validated_params[m_name] = validated_params.pop(l_name) - break - # Re-derive missing required params after type validation. # validated_params (current turn) takes precedence over already_collected # so that explicit user corrections override prior values. @@ -591,8 +833,7 @@ def _parse_prediction( # LLM incorrectly returned "none" despite missing params — reset to empty # string so callers receive a reliable signal that a follow-up is needed. logger.warning( - "LLM returned clarifying_question='none' but required params are " - f"still missing: {missing_required}. Resetting to empty string." + f"ParamExtractionModule: LLM returned 'none' despite missing params: {missing_required}" ) clarifying_question = "" diff --git a/src/tool_classifier/workflows/api_tool_workflow.py b/src/tool_classifier/workflows/api_tool_workflow.py index bc1dc94b..b52c21fe 100644 --- a/src/tool_classifier/workflows/api_tool_workflow.py +++ b/src/tool_classifier/workflows/api_tool_workflow.py @@ -1,11 +1,13 @@ """API Tool Calling Workflow Executor — Layer 2 of the classification chain.""" import asyncio +import time from dataclasses import dataclass, field from typing import ( TYPE_CHECKING, Any, AsyncIterator, + Coroutine, Dict, List, Literal, @@ -15,21 +17,32 @@ cast, ) -from loguru import logger +from src.loki_logger import LokiLogger +from llm_orchestrator_config.feature_flags import FeatureFlags from models.request_models import ( OrchestrationRequest, OrchestrationResponse, TestOrchestrationResponse, ) -from models.session_models import APIToolSession +from models.session_models import APIToolSession, EndpointSessionState, LastCallContext from tool_classifier.agentic_loop import AgenticLoop from tool_classifier.api_caller import APICaller from tool_classifier.api_response_formatter import APIResponseFormatterModule from tool_classifier.base_workflow import BaseWorkflow -from tool_classifier.enums import AgenticLoopStatus +from tool_classifier.enums import AgenticLoopStatus, ExecutionMode from tool_classifier.param_extractor import ParamExtractionModule from utils.api_tool_session_store import APIToolSessionStore +from utils.conversation_history_helpers import get_conversation_history +from utils.conversation_history_store import ConversationHistoryStore +from utils.atc_cache_store import ATCCacheStore +from tool_classifier.constants import ATC_CACHE_DEFAULT_TTL_SECONDS +from tool_classifier.follow_up_detector import FollowUpDetectorModule +from tool_classifier.multi_agentic_loop import MultiEndpointAgenticLoop +from tool_classifier.multi_api_caller import MultiAPICaller +from tool_classifier.multi_response_formatter import MultiResponseFormatterModule + +logger = LokiLogger(service_name="api-tool-calling") if TYPE_CHECKING: from guardrails.nemo_rails_adapter import NeMoRailsAdapter @@ -57,6 +70,14 @@ async def handle_output_guardrails( """Check output guardrails and return (possibly replaced) response.""" ... + async def store_streaming_inference( + self, + request: OrchestrationRequest, + final_answer: str, + ) -> None: + """Store streaming inference data for production/testing environments.""" + ... + @dataclass class _LoopStep: @@ -64,20 +85,33 @@ class _LoopStep: ``kind`` drives both the sync and streaming execution paths: - * ``"api_call"`` — all params collected; call the external API and format. - * ``"question"`` — agentic loop needs more input; return ``question`` to user. - * ``"fallback"`` — nothing to do; caller should fall back to RAG. + * ``"api_call"`` — single endpoint; all params collected; call API and format. + Populates: ``endpoint``, ``collected_params``, ``user_query``. + * ``"multi_api_call"`` — parallel endpoints; call all APIs concurrently and merge results. + Populates: ``parallel_endpoints``, ``user_query``. + ``endpoint`` and ``collected_params`` are empty/unused. + * ``"question"`` — agentic loop needs more input; return ``question`` to user. + Populates: ``question``, ``question_tokens``. + * ``"fallback"`` — nothing to do; caller should fall back to RAG. + No additional fields are populated. + * ``"cached_response"`` — cache hit (L1 or L2); skip APICaller; format ``cached_raw_response``. + Populates: ``endpoint``, ``cached_raw_response``, ``collected_params``, ``user_query``, ``cache_source``. """ - kind: Literal["api_call", "question", "fallback"] + kind: Literal[ + "api_call", "multi_api_call", "question", "fallback", "cached_response" + ] chat_id: str = "" endpoint: Dict[str, Any] = field(default_factory=dict) + parallel_endpoints: List[EndpointSessionState] = field(default_factory=list) collected_params: Dict[str, Any] = field(default_factory=dict) detected_language: str = "en" user_query: str = "" question: str = "" question_tokens: List[str] = field(default_factory=list) custom_instructions: str = "" + cached_raw_response: Any = None + cache_source: str = "L1" class APIToolWorkflowExecutor(BaseWorkflow): @@ -111,18 +145,40 @@ def __init__( if orchestration_service is not None else None ) + self._background_tasks: set[asyncio.Task[None]] = set() # ------------------------------------------------------------------ # Internal helpers # ------------------------------------------------------------------ + def _create_background_task(self, coro: Coroutine[Any, Any, None]) -> None: + """Keep fire-and-forget tasks alive until they finish.""" + task = asyncio.create_task(coro) + self._background_tasks.add(task) + task.add_done_callback(self._discard_background_task) + + def _discard_background_task(self, task: asyncio.Task[None]) -> None: + """Remove a completed background task and log any uncaught exception.""" + self._background_tasks.discard(task) + if task.cancelled(): + return + exception = task.exception() + if exception is not None: + logger.warning(f"APIToolWorkflow: background task failed: {exception}") + def _get_session_store(self) -> Optional[APIToolSessionStore]: """Return the session store from the orchestration service, or None.""" if self.orchestration_service is None: return None return getattr(self.orchestration_service, "session_store", None) + def _get_conversation_history_store(self) -> Optional[ConversationHistoryStore]: + """Return the conversation history store from the orchestration service, or None.""" + if self.orchestration_service is None: + return None + return getattr(self.orchestration_service, "conversation_history_store", None) + def _get_guardrails_adapter( self, environment: str, connection_id: Optional[str] = None ) -> Optional["NeMoRailsAdapter"]: @@ -212,6 +268,22 @@ def _language_from_custom_instructions(custom_instructions: str) -> Optional[str def _required_params(params: List[Dict[str, Any]]) -> List[Dict[str, Any]]: return [p for p in params if isinstance(p, dict) and p.get("required", False)] + @staticmethod + def _missing_required_params( + schema: List[Dict[str, Any]], collected: Dict[str, Any] + ) -> List[str]: + """Return names of required params from *schema* not yet present in *collected*.""" + missing: list[str] = [] + for p in schema: + if not isinstance(p, dict) or not p.get("required", False): + continue + name = p.get("name") + if not isinstance(name, str) or not name: + continue + if name not in collected: + missing.append(name) + return missing + async def _execute_api_and_format( self, chat_id: str, @@ -231,11 +303,33 @@ async def _execute_api_and_format( method = endpoint.get("method", "GET") description = endpoint.get("description", "") + # L1 cache check — collected_params is complete here, so the hash matches + # what was stored on the previous successful call with the same params. + if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED and endpoint.get("cacheable", True): + _cache_store = ATCCacheStore() + _cached = await _cache_store.get_l1( + chat_id, endpoint.get("name", ""), collected_params + ) + if _cached is not None: + logger.info( + f"[{chat_id}] ATC cache: L1 hit for {endpoint.get('name')!r} " + f"— skipping API call" + ) + return await self._format_cached_response( + chat_id=chat_id, + endpoint=endpoint, + user_query=user_query, + detected_language=detected_language, + cached_raw_response=_cached, + collected_params=collected_params, + custom_instructions=custom_instructions, + cache_source="L1", + ) + logger.info( f"[{chat_id}] APIToolWorkflow: calling API " f"{method} {url} with params={list(collected_params.keys())}" ) - api_result = await self._api_caller.call( url=url, method=method, @@ -248,6 +342,43 @@ async def _execute_api_and_format( f"[{chat_id}] APIToolWorkflow: API call succeeded " f"(status={api_result.status_code})" ) + if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED and endpoint.get( + "cacheable", True + ): + _ep_name = endpoint.get("name", "") + _ttl_override = endpoint.get("cache_ttl_seconds") + _ttl = ( + _ttl_override + if isinstance(_ttl_override, int) and _ttl_override > 0 + else ATC_CACHE_DEFAULT_TTL_SECONDS + ) + _resp_data = api_result.response_data + + async def _write_l1_l2() -> None: + try: + _cs = ATCCacheStore() + await _cs.set_l1( + chat_id, _ep_name, collected_params, _resp_data, _ttl + ) + await _cs.set_l2( + chat_id, + [ + LastCallContext( + api_name=_ep_name, + endpoint=endpoint, + collected_params=collected_params, + raw_response=_resp_data, + original_query=user_query, + timestamp=time.time(), + ) + ], + ) + except Exception as _exc: + logger.warning( + f"[{chat_id}] ATC cache: background write failed: {_exc}" + ) + + self._create_background_task(_write_l1_l2()) formatter = APIResponseFormatterModule( custom_instructions=custom_instructions ) @@ -257,6 +388,7 @@ async def _execute_api_and_format( api_response=api_result.response_data, endpoint_description=description, detected_language=detected_language, + collected_params=collected_params, ) else: logger.warning( @@ -283,6 +415,258 @@ def _build_question_response(chat_id: str, question: str) -> OrchestrationRespon content=question, ) + async def _execute_multi_api_and_format( + self, + chat_id: str, + parallel_endpoints: List[EndpointSessionState], + user_query: str, + detected_language: str, + custom_instructions: str = "", + ) -> OrchestrationResponse: + """Call all parallel endpoints concurrently and return a merged natural-language response. + + Builds a ``call_params``-keyed payload for each :class:`EndpointSessionState`, + dispatches all calls concurrently via :class:`MultiAPICaller`, then synthesises + the results into one unified answer with :class:`MultiResponseFormatterModule`. + """ + call_payloads = [ + {**state.endpoint, "call_params": state.collected_params} + for state in parallel_endpoints + ] + ep_names = [ + state.endpoint.get("name", "") for state in parallel_endpoints + ] + logger.info( + f"[{chat_id}] APIToolWorkflow: parallel — calling {len(call_payloads)} APIs " + f"concurrently: {ep_names}" + ) + + multi_caller = MultiAPICaller(self._api_caller) + multi_result = await multi_caller.call_all( + call_payloads, language=detected_language + ) + + logger.info( + f"[{chat_id}] APIToolWorkflow: parallel batch complete — " + f"{sum(r.success for r in multi_result.results)}/{len(multi_result.results)} succeeded" + ) + + _pairs = list(zip(parallel_endpoints, multi_result.results, strict=True)) + + if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED: + + async def _write_multi_cache() -> None: + try: + _cs = ATCCacheStore() + _ctxs: list[LastCallContext] = [] + for _state, _result in _pairs: + if _result.success and _state.endpoint.get("cacheable", True): + _ttl_override = _state.endpoint.get("cache_ttl_seconds") + _ttl = ( + _ttl_override + if isinstance(_ttl_override, int) and _ttl_override > 0 + else ATC_CACHE_DEFAULT_TTL_SECONDS + ) + await _cs.set_l1( + chat_id, + _state.endpoint.get("name", ""), + _state.collected_params, + _result.response_data, + _ttl, + ) + _ctxs.append( + LastCallContext( + api_name=_state.endpoint.get("name", ""), + endpoint=_state.endpoint, + collected_params=_state.collected_params, + raw_response=_result.response_data, + original_query=user_query, + timestamp=time.time(), + ) + ) + if _ctxs: + await _cs.set_l2(chat_id, _ctxs) + except Exception as _exc: + logger.warning( + f"[{chat_id}] ATC cache: background multi-write failed: {_exc}" + ) + + self._create_background_task(_write_multi_cache()) + + api_results = [ + ( + state.endpoint.get("name", ""), + state.endpoint.get("description", ""), + result.response_data if result.success else result.error or "", + state.collected_params, + ) + for state, result in zip( + parallel_endpoints, multi_result.results, strict=True + ) + ] + + formatter = MultiResponseFormatterModule( + custom_instructions=custom_instructions + ) + content = await asyncio.to_thread( + formatter.forward, + user_query=user_query, + api_results=api_results, + detected_language=detected_language, + ) + + return OrchestrationResponse( + chatId=chat_id, + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content=content, + ) + + async def _stream_multi_api_and_format( + self, + chat_id: str, + parallel_endpoints: List[EndpointSessionState], + user_query: str, + detected_language: str, + orchestration_service: OrchestrationServiceProtocol, + request: OrchestrationRequest, + costs_metric: Optional[Dict[str, Any]] = None, + custom_instructions: str = "", + ) -> AsyncIterator[str]: + """Call all parallel APIs concurrently, then stream the merged answer token by token. + + API calls are pre-resolved before streaming starts so the LLM synthesis step + receives all results at once. Uses the same buffer-first guardrails approach as + :meth:`_stream_api_and_format` — the full response is assembled and validated + before any token is sent to the client. + """ + call_payloads = [ + {**state.endpoint, "call_params": state.collected_params} + for state in parallel_endpoints + ] + ep_names = [ + state.endpoint.get("name", "") for state in parallel_endpoints + ] + logger.info( + f"[{chat_id}] APIToolWorkflow (streaming): parallel — calling {len(call_payloads)} APIs " + f"concurrently: {ep_names}" + ) + + multi_caller = MultiAPICaller(self._api_caller) + multi_result = await multi_caller.call_all( + call_payloads, language=detected_language + ) + + logger.info( + f"[{chat_id}] APIToolWorkflow (streaming): parallel batch complete — " + f"{sum(r.success for r in multi_result.results)}/{len(multi_result.results)} succeeded" + ) + + _pairs = list(zip(parallel_endpoints, multi_result.results, strict=True)) + + if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED: + + async def _write_multi_cache() -> None: + try: + _cs = ATCCacheStore() + _ctxs: list[LastCallContext] = [] + for _state, _result in _pairs: + if _result.success and _state.endpoint.get("cacheable", True): + _ttl_override = _state.endpoint.get("cache_ttl_seconds") + _ttl = ( + _ttl_override + if isinstance(_ttl_override, int) and _ttl_override > 0 + else ATC_CACHE_DEFAULT_TTL_SECONDS + ) + await _cs.set_l1( + chat_id, + _state.endpoint.get("name", ""), + _state.collected_params, + _result.response_data, + _ttl, + ) + _ctxs.append( + LastCallContext( + api_name=_state.endpoint.get("name", ""), + endpoint=_state.endpoint, + collected_params=_state.collected_params, + raw_response=_result.response_data, + original_query=user_query, + timestamp=time.time(), + ) + ) + if _ctxs: + await _cs.set_l2(chat_id, _ctxs) + except Exception as _exc: + logger.warning( + f"[{chat_id}] ATC cache: background multi-write failed: {_exc}" + ) + + self._create_background_task(_write_multi_cache()) + + api_results = [ + ( + state.endpoint.get("name", ""), + state.endpoint.get("description", ""), + result.response_data if result.success else result.error or "", + state.collected_params, + ) + for state, result in zip( + parallel_endpoints, multi_result.results, strict=True + ) + ] + + formatter = MultiResponseFormatterModule( + custom_instructions=custom_instructions + ) + buffered_tokens = [ + token + async for token in formatter.stream_forward_multi( + user_query=user_query, + api_results=api_results, + detected_language=detected_language, + ) + ] + + full_response = "".join(buffered_tokens) + final_answer = full_response # Track what gets sent to user + + guardrails_passed = True + if orchestration_service is not None: + guardrails_adapter = self._get_guardrails_adapter( + request.environment, request.connection_id + ) + if guardrails_adapter is not None: + dummy_response = OrchestrationResponse( + chatId=chat_id, + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content=full_response, + ) + checked = await orchestration_service.handle_output_guardrails( + guardrails_adapter, + dummy_response, + request, + costs_metric if costs_metric is not None else {}, + ) + if checked.content != full_response: + logger.warning( + f"[{chat_id}] APIToolWorkflow (streaming): " + f"parallel output blocked by guardrails" + ) + yield orchestration_service.format_sse(chat_id, checked.content) + final_answer = checked.content + guardrails_passed = False + + if guardrails_passed: + for token in buffered_tokens: + yield orchestration_service.format_sse(chat_id, token) + + await orchestration_service.store_streaming_inference(request, final_answer) + yield orchestration_service.format_sse(chat_id, "END") + # ------------------------------------------------------------------ # Core loop handler — shared by async and streaming paths # ------------------------------------------------------------------ @@ -323,13 +707,35 @@ async def _compute_loop_step( await session_store.delete(chat_id) return _LoopStep(kind="fallback", chat_id=chat_id) - logger.info( - f"[{chat_id}] APIToolWorkflow: resuming session " - f"(turn={session.turn_count}, endpoint={endpoint.get('name')!r})" - ) + if session.execution_mode == ExecutionMode.PARALLEL.value: + ep_names = [s.endpoint.get("name") for s in session.parallel_endpoints] + logger.info( + f"[{chat_id}] APIToolWorkflow: resuming parallel session " + f"(turn={session.turn_count}, endpoints={ep_names})" + ) + else: + logger.info( + f"[{chat_id}] APIToolWorkflow: resuming session " + f"(turn={session.turn_count}, endpoint={endpoint.get('name')!r})" + ) else: - # New-session path — endpoint must come from classifier context - endpoint = context.get("matched_endpoint") + # New-session path — endpoint must come from classifier context. + # For parallel execution_mode, drive param collection for the first + # endpoint (Phase 3 MultiEndpointAgenticLoop will advance the index). + all_matched: list[dict[str, Any]] = [] + if context.get("execution_mode") == ExecutionMode.PARALLEL: + all_matched = context.get("matched_endpoints", []) + endpoint = all_matched[0] if all_matched else None + if endpoint: + logger.info( + f"[{chat_id}] APIToolWorkflow: parallel mode — " + f"starting param collection for first endpoint " + f"{endpoint.get('name')!r} " + f"({len(all_matched)} endpoints total)" + ) + else: + endpoint = context.get("matched_endpoint") + if not endpoint: logger.warning( f"[{chat_id}] APIToolWorkflow: no matched_endpoint in context " @@ -339,8 +745,185 @@ async def _compute_loop_step( params_schema: List[Dict[str, Any]] = endpoint.get("params", []) - # Fast path: no required params — call API immediately - if not self._required_params(params_schema): + # L1 + L2 cache checks — single-mode new queries only; both guarded by + # the ATC_RESPONSE_CACHE_ENABLED kill-switch. + if ( + FeatureFlags.ATC_RESPONSE_CACHE_ENABLED + and not all_matched + and endpoint.get("cacheable", True) + ): + _cache_store = ATCCacheStore() + + # ── L1: exact param-hash hit ───────────────────────────────── + _cached = await _cache_store.get_l1( + chat_id, + endpoint.get("name", ""), + context.get("pre_extracted_params", {}), + ) + if _cached is not None: + logger.info( + f"[{chat_id}] ATC cache: L1 hit for {endpoint.get('name')!r}" + ) + return _LoopStep( + kind="cached_response", + chat_id=chat_id, + endpoint=endpoint, + cached_raw_response=_cached, + detected_language=getattr(request, "_detected_language", "en"), + user_query=request.message, + custom_instructions=custom_instructions, + collected_params=context.get("pre_extracted_params", {}), + ) + + # ── L2: follow-up routing based on last call context ───────── + _last_calls = await _cache_store.get_l2(chat_id) + if _last_calls: + _matching = next( + (c for c in _last_calls if c.api_name == endpoint.get("name")), + None, + ) + if _matching is not None: + try: + _detector = FollowUpDetectorModule() + _det_result = await asyncio.to_thread( + _detector.forward, + user_query=request.message, + previous_query=_matching.original_query, + previous_params=_matching.collected_params, + params_schema=endpoint.get("params", []), + ) + if _det_result["follow_up_type"] == "response_question": + logger.info( + f"[{chat_id}] ATC cache: L2 follow-up — response_question" + ) + return _LoopStep( + kind="cached_response", + chat_id=chat_id, + endpoint=endpoint, + cached_raw_response=_matching.raw_response, + detected_language=getattr( + request, "_detected_language", "en" + ), + user_query=request.message, + custom_instructions=custom_instructions, + cache_source="L2", + collected_params=_matching.collected_params, + ) + elif _det_result["follow_up_type"] == "param_update": + _updated = _det_result["updated_params"] + _merged = {**_matching.collected_params, **_updated} + _missing = self._missing_required_params( + endpoint.get("params", []), _merged + ) + if not _missing: + # If the merged params are identical to the previous + # call's params (e.g. no genuine new params survived + # schema validation), the user is re-requesting the + # same data → serve from L2 directly without an API + # call or L1 lookup. + if ATCCacheStore._param_hash( + _merged + ) == ATCCacheStore._param_hash( + _matching.collected_params + ): + # Params unchanged — prefer L1 (exact hash hit + # with the actual previously-collected params). + # L1 uses a TTL so it may have expired; fall + # back to L2 raw_response in that case. + _l1_data = await _cache_store.get_l1( + chat_id, + endpoint.get("name", ""), + _matching.collected_params, + ) + if _l1_data is not None: + logger.info( + f"[{chat_id}] ATC cache: L2 follow-up — param_update " + f"(params unchanged), L1 hit — serving from L1" + ) + return _LoopStep( + kind="cached_response", + chat_id=chat_id, + endpoint=endpoint, + cached_raw_response=_l1_data, + detected_language=getattr( + request, "_detected_language", "en" + ), + user_query=request.message, + custom_instructions=custom_instructions, + cache_source="L1", + collected_params=_matching.collected_params, + ) + logger.info( + f"[{chat_id}] ATC cache: L2 follow-up — param_update " + f"(params unchanged), L1 miss — serving from L2" + ) + return _LoopStep( + kind="cached_response", + chat_id=chat_id, + endpoint=endpoint, + cached_raw_response=_matching.raw_response, + detected_language=getattr( + request, "_detected_language", "en" + ), + user_query=request.message, + custom_instructions=custom_instructions, + cache_source="L2", + collected_params=_matching.collected_params, + ) + logger.info( + f"[{chat_id}] ATC cache: L2 follow-up — param_update " + f"(all params present), calling API directly" + ) + return _LoopStep( + kind="api_call", + chat_id=chat_id, + endpoint=endpoint, + collected_params=_merged, + detected_language=getattr( + request, "_detected_language", "en" + ), + user_query=request.message, + custom_instructions=custom_instructions, + ) + else: + logger.info( + f"[{chat_id}] ATC cache: L2 follow-up — param_update " + f"(missing: {_missing}), seeding agentic loop" + ) + context["seeded_params"] = _merged + # new_intent → fall through to normal agentic-loop path + except Exception as _exc: + logger.warning( + f"[{chat_id}] ATC cache: L2 follow-up detection failed: " + f"{_exc} — falling through to normal path" + ) + + # Fast path — skip the agentic loop when no required params need collecting. + # user_query falls back to request.message on the fast path because no session + # exists yet (original_query is only stored once a session is created). + user_query_for_fast_path = request.message + if all_matched: + # Parallel mode: fast-path only when ALL endpoints have no required params. + if all( + not self._required_params(ep.get("params", [])) + for ep in all_matched + ): + logger.info( + f"[{chat_id}] APIToolWorkflow: parallel fast path — " + f"all {len(all_matched)} endpoints have no required params" + ) + return _LoopStep( + kind="multi_api_call", + chat_id=chat_id, + parallel_endpoints=[ + EndpointSessionState(endpoint=e) for e in all_matched + ], + detected_language=getattr(request, "_detected_language", "en"), + user_query=user_query_for_fast_path, + custom_instructions=custom_instructions, + ) + elif not self._required_params(params_schema): + # Single mode: the only endpoint has no required params. logger.info( f"[{chat_id}] APIToolWorkflow: endpoint {endpoint.get('name')!r} " f"has no required params — fast path" @@ -351,22 +934,29 @@ async def _compute_loop_step( endpoint=endpoint, collected_params={}, detected_language=getattr(request, "_detected_language", "en"), - user_query=request.message, + user_query=user_query_for_fast_path, custom_instructions=custom_instructions, ) # Create a new session before running the first loop turn + _seeded_params = context.get("seeded_params", {}) if session_store is not None: new_session = APIToolSession( chat_id=chat_id, state="collecting_params", selected_endpoint=endpoint, - collected_params={}, + collected_params=_seeded_params, turn_count=0, max_turns=5, awaiting_continuation=False, detected_language=getattr(request, "_detected_language", "en"), original_query=request.message, + execution_mode=ExecutionMode.PARALLEL.value + if all_matched + else ExecutionMode.SINGLE.value, + parallel_endpoints=[ + EndpointSessionState(endpoint=e) for e in all_matched + ], ) await session_store.save(new_session) session = new_session @@ -379,12 +969,18 @@ async def _compute_loop_step( chat_id=chat_id, state="collecting_params", selected_endpoint=endpoint, - collected_params={}, + collected_params=_seeded_params, turn_count=0, max_turns=5, awaiting_continuation=False, detected_language=getattr(request, "_detected_language", "en"), original_query=request.message, + execution_mode=ExecutionMode.PARALLEL.value + if all_matched + else ExecutionMode.SINGLE.value, + parallel_endpoints=[ + EndpointSessionState(endpoint=e) for e in all_matched + ], ) # ── Run one loop turn ───────────────────────────────────────────── @@ -394,8 +990,6 @@ async def _compute_loop_step( f"agentic loop running without persistence" ) - loop = self._build_agentic_loop(session_store, custom_instructions) # type: ignore[arg-type] - # If custom_instructions contain a language directive (e.g. "respond in English" # inside an Estonian-language prompt), use that directed language for all # user-facing questions — both the LLM-generated clarifying questions and the @@ -405,26 +999,109 @@ async def _compute_loop_step( or session.detected_language ) + _atc_conversation_summary: Optional[str] = None + conversation_history_for_loop: List[Dict[str, Any]] + if session.turn_count == 0: + # On the first ATC turn there is no prior ATC exchange to pass. + conversation_history_for_loop = [] + else: + # On subsequent turns prefer Redis as the authoritative history source. + _redis_history, _atc_conversation_summary = await get_conversation_history( + chat_id=chat_id, + store=self._get_conversation_history_store(), + fallback=list(request.conversationHistory or []), + ) + conversation_history_for_loop = [ + {"authorRole": item.authorRole, "message": item.message} + for item in _redis_history + ] + + # Incorporate any Redis-supplied conversation summary into custom_instructions + # so the param extractor and response formatter have full conversational context. + if _atc_conversation_summary: + _summary_prefix = ( + f"Summary of earlier conversation: {_atc_conversation_summary}" + ) + custom_instructions = ( + f"{_summary_prefix}\n\n{custom_instructions}".strip() + if custom_instructions + else _summary_prefix + ) + + if session.execution_mode == ExecutionMode.PARALLEL.value: + # Parallel path: MultiEndpointAgenticLoop operates on the full + # parallel_endpoints list and builds a merged schema internally so + # only one clarifying question is asked per turn. + multi_loop = MultiEndpointAgenticLoop( + session_store=session_store, + param_extractor=ParamExtractionModule( + custom_instructions=custom_instructions + ), + ) + result, question_tokens = await multi_loop.stream_run_turn( + chat_id=chat_id, + user_message=request.message, + conversation_history=conversation_history_for_loop, + endpoint_states=session.parallel_endpoints, + turn_count=session.turn_count, + awaiting_continuation=session.awaiting_continuation, + session_language=effective_session_language, + ) + + if result.status == AgenticLoopStatus.COMPLETED: + logger.info( + f"[{chat_id}] APIToolWorkflow: parallel params fully collected " + f"(turns={result.turn_count}, " + f"endpoints={[s.endpoint.get('name') for s in session.parallel_endpoints]})" + ) + if session_store is not None: + await session_store.delete(chat_id) + return _LoopStep( + kind="multi_api_call", + chat_id=chat_id, + parallel_endpoints=session.parallel_endpoints, + detected_language=effective_session_language, + user_query=session.original_query or request.message, + custom_instructions=custom_instructions, + ) + + if result.status == AgenticLoopStatus.MAX_TURNS_REACHED: + logger.info( + f"[{chat_id}] APIToolWorkflow: parallel max turns reached — deleting session" + ) + if session_store is not None: + await session_store.delete(chat_id) + return _LoopStep(kind="fallback", chat_id=chat_id) + + # NEEDS_INPUT or AWAITING_CONTINUATION_DECISION + logger.info( + f"[{chat_id}] APIToolWorkflow: parallel — asking for more info " + f"(status={result.status.value}, turn={result.turn_count})" + ) + return _LoopStep( + kind="question", + chat_id=chat_id, + question=result.clarifying_question, + question_tokens=question_tokens, + custom_instructions=custom_instructions, + ) + + # Single path: AgenticLoop (untouched) + loop = self._build_agentic_loop(session_store, custom_instructions) # type: ignore[arg-type] + result, question_tokens = await loop.stream_run_turn( chat_id=chat_id, user_message=request.message, - conversation_history=( - [] - if session.turn_count == 0 - else [ - {"authorRole": item.authorRole, "message": item.message} - for item in (request.conversationHistory or []) - ] - ), + conversation_history=conversation_history_for_loop, params_schema=endpoint.get("params", []), collected_params=session.collected_params, turn_count=session.turn_count, max_turns=session.max_turns, awaiting_continuation=session.awaiting_continuation, session_language=effective_session_language, + seeded_params=context.get("seeded_params"), ) - # ── Translate result into a _LoopStep ───────────────────────────── if result.status == AgenticLoopStatus.COMPLETED: logger.info( f"[{chat_id}] APIToolWorkflow: all params collected " @@ -463,6 +1140,117 @@ async def _compute_loop_step( custom_instructions=custom_instructions, ) + async def _format_cached_response( + self, + chat_id: str, + endpoint: Dict[str, Any], + user_query: str, + detected_language: str, + cached_raw_response: Any, + collected_params: Dict[str, Any], + custom_instructions: str = "", + cache_source: str = "L1", + ) -> OrchestrationResponse: + """Format a cached API response without calling the external API. + + Used when a cache hit (L1 or L2) is detected for a matched endpoint. + Runs the raw response through APIResponseFormatterModule, bypassing APICaller. + The collected_params are passed to the formatter to maintain query context consistency. + """ + description = endpoint.get("description", "") + logger.info( + f"[{chat_id}] ATC cache: formatting {cache_source} hit for {endpoint.get('name')!r}" + ) + formatter = APIResponseFormatterModule(custom_instructions=custom_instructions) + content = await asyncio.to_thread( + formatter.forward, + user_query=user_query, + api_response=cached_raw_response, + endpoint_description=description, + detected_language=detected_language, + collected_params=collected_params, + ) + return OrchestrationResponse( + chatId=chat_id, + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content=content, + ) + + async def _stream_cached_response( + self, + chat_id: str, + endpoint: Dict[str, Any], + user_query: str, + detected_language: str, + cached_raw_response: Any, + collected_params: Dict[str, Any], + orchestration_service: OrchestrationServiceProtocol, + request: OrchestrationRequest, + costs_metric: Optional[Dict[str, Any]] = None, + custom_instructions: str = "", + cache_source: str = "L1", + ) -> AsyncIterator[str]: + """Stream a cached API response without calling the external API. + + Used when a cache hit (L1 or L2) is detected for a matched endpoint. + Runs the raw response through APIResponseFormatterModule streaming, bypassing APICaller. + The collected_params are passed to the formatter to maintain query context consistency. + """ + description = endpoint.get("description", "") + logger.info( + f"[{chat_id}] ATC cache: streaming {cache_source} hit for {endpoint.get('name')!r}" + ) + formatter = APIResponseFormatterModule(custom_instructions=custom_instructions) + buffered_tokens = [ + token + async for token in formatter.stream_forward( + user_query=user_query, + api_response=cached_raw_response, + endpoint_description=description, + detected_language=detected_language, + collected_params=collected_params, + ) + ] + + full_response = "".join(buffered_tokens) + final_answer = full_response # Track what gets sent to user + + guardrails_passed = True + if orchestration_service is not None: + guardrails_adapter = self._get_guardrails_adapter( + request.environment, request.connection_id + ) + if guardrails_adapter is not None: + dummy_response = OrchestrationResponse( + chatId=chat_id, + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content=full_response, + ) + checked = await orchestration_service.handle_output_guardrails( + guardrails_adapter, + dummy_response, + request, + costs_metric if costs_metric is not None else {}, + ) + if checked.content != full_response: + logger.warning( + f"[{chat_id}] ATC cache: streaming {cache_source} hit blocked by guardrails" + ) + yield orchestration_service.format_sse(chat_id, checked.content) + final_answer = checked.content + guardrails_passed = False + + if guardrails_passed: + for token in buffered_tokens: + yield orchestration_service.format_sse(chat_id, token) + + await orchestration_service.store_streaming_inference(request, final_answer) + yield orchestration_service.format_sse(chat_id, "END") + async def _run( self, request: OrchestrationRequest, @@ -477,6 +1265,59 @@ async def _run( if step.kind == "question": return self._build_question_response(step.chat_id, step.question) + # multi_api_call — concurrent batch execution + multi-formatter + if step.kind == "multi_api_call": + response = await self._execute_multi_api_and_format( + chat_id=step.chat_id, + parallel_endpoints=step.parallel_endpoints, + user_query=step.user_query, + detected_language=step.detected_language, + custom_instructions=step.custom_instructions, + ) + if response.llmServiceActive and self.orchestration_service is not None: + guardrails_adapter = self._get_guardrails_adapter( + request.environment, request.connection_id + ) + if guardrails_adapter is not None: + response = cast( + OrchestrationResponse, + await self.orchestration_service.handle_output_guardrails( + guardrails_adapter, + response, + request, + {}, + ), + ) + return response + + # cached_response — cache hit (L1 or L2): skip APICaller, format cached data + if step.kind == "cached_response": + response = await self._format_cached_response( + chat_id=step.chat_id, + endpoint=step.endpoint, + user_query=step.user_query, + detected_language=step.detected_language, + collected_params=step.collected_params, + custom_instructions=step.custom_instructions, + cached_raw_response=step.cached_raw_response, + cache_source=step.cache_source, + ) + if response.llmServiceActive and self.orchestration_service is not None: + guardrails_adapter = self._get_guardrails_adapter( + request.environment, request.connection_id + ) + if guardrails_adapter is not None: + response = cast( + OrchestrationResponse, + await self.orchestration_service.handle_output_guardrails( + guardrails_adapter, + response, + request, + {}, + ), + ) + return response + # api_call — blocking response = await self._execute_api_and_format( chat_id=step.chat_id, @@ -492,15 +1333,16 @@ async def _run( guardrails_adapter = self._get_guardrails_adapter( request.environment, request.connection_id ) - response = cast( - OrchestrationResponse, - await self.orchestration_service.handle_output_guardrails( - guardrails_adapter, - response, - request, - {}, - ), - ) + if guardrails_adapter is not None: + response = cast( + OrchestrationResponse, + await self.orchestration_service.handle_output_guardrails( + guardrails_adapter, + response, + request, + {}, + ), + ) return response @@ -537,11 +1379,38 @@ async def _stream_api_and_format( method = endpoint.get("method", "GET") description = endpoint.get("description", "") + # L1 cache check — collected_params is complete here, so the hash matches + # what was stored on the previous successful call with the same params. + if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED and endpoint.get("cacheable", True): + _cache_store = ATCCacheStore() + _cached = await _cache_store.get_l1( + chat_id, endpoint.get("name", ""), collected_params + ) + if _cached is not None: + logger.info( + f"[{chat_id}] ATC cache: L1 hit for {endpoint.get('name')!r} " + f"— skipping API call (streaming)" + ) + async for token in self._stream_cached_response( + chat_id=chat_id, + endpoint=endpoint, + user_query=user_query, + detected_language=detected_language, + cached_raw_response=_cached, + collected_params=collected_params, + orchestration_service=orchestration_service, + request=request, + costs_metric=costs_metric, + custom_instructions=custom_instructions, + cache_source="L1", + ): + yield token + return + logger.info( f"[{chat_id}] APIToolWorkflow (streaming): calling API " f"{method} {url} with params={list(collected_params.keys())}" ) - api_result = await self._api_caller.call( url=url, method=method, @@ -554,6 +1423,43 @@ async def _stream_api_and_format( f"[{chat_id}] APIToolWorkflow (streaming): API call succeeded " f"(status={api_result.status_code}), streaming formatted response" ) + if FeatureFlags.ATC_RESPONSE_CACHE_ENABLED and endpoint.get( + "cacheable", True + ): + _ep_name = endpoint.get("name", "") + _ttl_override = endpoint.get("cache_ttl_seconds") + _ttl = ( + _ttl_override + if isinstance(_ttl_override, int) and _ttl_override > 0 + else ATC_CACHE_DEFAULT_TTL_SECONDS + ) + _resp_data = api_result.response_data + + async def _write_l1_l2() -> None: + try: + _cs = ATCCacheStore() + await _cs.set_l1( + chat_id, _ep_name, collected_params, _resp_data, _ttl + ) + await _cs.set_l2( + chat_id, + [ + LastCallContext( + api_name=_ep_name, + endpoint=endpoint, + collected_params=collected_params, + raw_response=_resp_data, + original_query=user_query, + timestamp=time.time(), + ) + ], + ) + except Exception as _exc: + logger.warning( + f"[{chat_id}] ATC cache: background write failed: {_exc}" + ) + + self._create_background_task(_write_l1_l2()) # Buffer all tokens first, then validate with output guardrails before # streaming to the client (validate-first approach). formatter = APIResponseFormatterModule( @@ -566,10 +1472,12 @@ async def _stream_api_and_format( api_response=api_result.response_data, endpoint_description=description, detected_language=detected_language, + collected_params=collected_params, ) ] full_response = "".join(buffered_tokens) + final_answer_to_store = full_response # Track what gets sent to user # Run output guardrails on the complete response guardrails_passed = True @@ -599,6 +1507,7 @@ async def _stream_api_and_format( f"output blocked by guardrails" ) yield orchestration_service.format_sse(chat_id, checked.content) + final_answer_to_store = checked.content guardrails_passed = False if guardrails_passed: @@ -609,8 +1518,12 @@ async def _stream_api_and_format( f"[{chat_id}] APIToolWorkflow (streaming): API call failed " f"(status={api_result.status_code}, error={api_result.error!r})" ) + final_answer_to_store = api_result.error or "" yield orchestration_service.format_sse(chat_id, api_result.error or "") + await orchestration_service.store_streaming_inference( + request, final_answer_to_store + ) yield orchestration_service.format_sse(chat_id, "END") async def execute_streaming( @@ -643,12 +1556,47 @@ async def execute_streaming( if step.kind == "question": async def _stream_question() -> AsyncIterator[str]: + accumulated_question: list[str] = [] for token in step.question_tokens or [step.question]: + accumulated_question.append(token) yield orchestration_service.format_sse(step.chat_id, token) + full_question = "".join(accumulated_question) + await orchestration_service.store_streaming_inference( + request, full_question + ) yield orchestration_service.format_sse(step.chat_id, "END") return _stream_question() + # multi_api_call — stream the merged answer from all parallel APIs + if step.kind == "multi_api_call": + return self._stream_multi_api_and_format( + chat_id=step.chat_id, + parallel_endpoints=step.parallel_endpoints, + user_query=step.user_query, + detected_language=step.detected_language, + orchestration_service=orchestration_service, + request=request, + costs_metric=context.get("costs_metric"), + custom_instructions=step.custom_instructions, + ) + + # cached_response — cache hit (L1 or L2): skip APICaller, stream formatted cached data + if step.kind == "cached_response": + return self._stream_cached_response( + chat_id=step.chat_id, + endpoint=step.endpoint, + user_query=step.user_query, + detected_language=step.detected_language, + collected_params=step.collected_params, + cached_raw_response=step.cached_raw_response, + orchestration_service=orchestration_service, + request=request, + costs_metric=context.get("costs_metric"), + custom_instructions=step.custom_instructions, + cache_source=step.cache_source, + ) + # api_call — stream the LLM-formatted answer token by token return self._stream_api_and_format( chat_id=step.chat_id, diff --git a/src/tool_classifier/workflows/context_workflow.py b/src/tool_classifier/workflows/context_workflow.py index f3369a6c..5352bd32 100644 --- a/src/tool_classifier/workflows/context_workflow.py +++ b/src/tool_classifier/workflows/context_workflow.py @@ -3,9 +3,12 @@ from typing import Any, AsyncIterator, Dict, Optional, cast import time import dspy -from loguru import logger +from langfuse import observe +from src.loki_logger import LokiLogger +from src.models.conversation_history_models import ConversationHistoryState from src.models.request_models import OrchestrationRequest, OrchestrationResponse +from src.utils.conversation_history_store import ConversationHistoryStore from tool_classifier.base_workflow import BaseWorkflow from tool_classifier.context_analyzer import ContextAnalyzer, ContextDetectionResult from tool_classifier.workflows.service_workflow import LLMServiceProtocol @@ -13,11 +16,18 @@ from src.llm_orchestrator_config.llm_manager import LLMManager from src.utils.cost_utils import get_lm_usage_since from src.utils.language_detector import detect_language +from src.utils.observation_utils import ( + safe_observation_context, + update_observation_safe, +) from src.llm_orchestrator_config.llm_ochestrator_constants import ( GUARDRAILS_BLOCKED_PHRASES, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE, ) +# Initialize Loki logger +logger = LokiLogger(service_name="context-workflow") + class ContextWorkflowExecutor(BaseWorkflow): """ @@ -46,6 +56,7 @@ def __init__( self, llm_manager: LLMManager, orchestration_service: Optional[LLMServiceProtocol] = None, + conversation_history_store: Optional[ConversationHistoryStore] = None, ) -> None: """ Initialize context workflow executor. @@ -53,15 +64,68 @@ def __init__( Args: llm_manager: LLM manager for context analysis orchestration_service: Reference to LLMOrchestrationService for cost logging + conversation_history_store: Redis-backed conversation history store; + when provided, history is read from Redis (canonical source) rather + than from ``request.conversationHistory``. """ self.llm_manager = llm_manager self.orchestration_service = orchestration_service + self.conversation_history_store = conversation_history_store self.context_analyzer = ContextAnalyzer(llm_manager) logger.info("Context workflow executor initialized") - @staticmethod - def _build_history(request: OrchestrationRequest) -> list[Dict[str, Any]]: - return [ + async def _build_history( + self, request: OrchestrationRequest + ) -> tuple[list[Dict[str, Any]], Optional[str]]: + """Fetch conversation history, preferring Redis over request payload. + + When a :class:`ConversationHistoryStore` is wired in and the session has + stored rounds, the Redis state is used as the canonical source of truth and + the GUI-provided ``request.conversationHistory`` is ignored. Any running + summary stored alongside the rounds is returned as ``pre_computed_summary`` + so that downstream callers can skip an expensive LLM summarisation step. + + If the store is absent, raises, or returns no rounds the method falls back + to the conversation history supplied in the request. + + Returns: + A ``(history_dicts, pre_computed_summary)`` tuple where + *history_dicts* is a list of ``{"authorRole", "message", "timestamp"}`` + dicts and *pre_computed_summary* is the Redis summary string or ``None``. + """ + if self.conversation_history_store is not None: + try: + state: ConversationHistoryState = ( + await self.conversation_history_store.get_context(request.chatId) + ) + if state.rounds: + history: list[Dict[str, Any]] = [] + for round_ in state.rounds: + history.append( + { + "authorRole": "user", + "message": round_.user_message, + "timestamp": str(round_.timestamp), + } + ) + history.append( + { + "authorRole": "bot", + "message": round_.bot_message, + "timestamp": str(round_.timestamp), + } + ) + logger.debug( + f"[{request.chatId}] Using Redis history: {len(state.rounds)} rounds, summary={'present' if state.summary else 'absent'}" + ) + return history, state.summary + except Exception as exc: + logger.warning( + f"[{request.chatId}] Redis history fetch failed, falling back to request history: {exc}" + ) + + # Fallback: use the conversation history supplied in the request + request_history: list[Dict[str, Any]] = [ { "authorRole": item.authorRole, "message": item.message, @@ -69,20 +133,28 @@ def _build_history(request: OrchestrationRequest) -> list[Dict[str, Any]]: } for item in request.conversationHistory ] + return request_history, None + @observe(name="context_workflow_detect", as_type="generation") async def _detect( self, message: str, history: list[Dict[str, Any]], time_metric: Dict[str, float], costs_metric: Dict[str, Dict[str, Any]], + pre_computed_summary: Optional[str] = None, ) -> Optional[ContextDetectionResult]: """Phase 1: run context detection with summary fallback. Checks the last 10 conversation turns first. If the query cannot be - answered from those and the history exceeds 10 turns, falls back to a - summary-based check over the older turns. Returns None on error so the - caller falls through to RAG. + answered from those and the history exceeds 10 turns (or a Redis summary + is available), falls back to a summary-based check. Returns None on error + so the caller falls through to RAG. + + Args: + pre_computed_summary: Running conversation summary retrieved from Redis. + When provided, the expensive LLM summarisation step is skipped and + this value is used directly. """ try: start = time.time() @@ -90,13 +162,29 @@ async def _detect( result, cost, ) = await self.context_analyzer.detect_context_with_summary_fallback( - query=message, conversation_history=history + query=message, + conversation_history=history, + pre_computed_summary=pre_computed_summary, ) time_metric["context.detection"] = time.time() - start costs_metric["context_detection"] = cost + update_observation_safe( + input_data={"query": message, "history_length": len(history)}, + output_data={ + "is_greeting": result.is_greeting if result else None, + "can_answer_from_context": result.can_answer_from_context + if result + else None, + }, + metadata={"usage": cost}, + ) return result except Exception as e: logger.error(f"Phase 1 detection failed: {e}", exc_info=True) + update_observation_safe( + output_data={"error": str(e)}, + metadata={"usage": {}}, + ) return None def _log_costs(self, costs_metric: Dict[str, Dict[str, Any]]) -> None: @@ -113,6 +201,7 @@ def _is_guardrail_violation(chunk: str) -> bool: for phrase in GUARDRAILS_BLOCKED_PHRASES ) + @observe(name="context_workflow_generate_response", as_type="generation") async def _generate_response_async( self, request: OrchestrationRequest, @@ -128,8 +217,21 @@ async def _generate_response_async( ) time_metric["context.generation"] = time.time() - start costs_metric["context_response"] = cost + update_observation_safe( + input_data={ + "chat_id": request.chatId, + "query": request.message, + "context_snippet_length": len(context_snippet), + }, + output_data={"has_answer": bool(answer)}, + metadata={"usage": cost}, + ) except Exception as e: logger.error(f"Phase 2 generation failed: {e}", exc_info=True) + update_observation_safe( + output_data={"error": str(e)}, + metadata={"usage": {}}, + ) self._log_costs(costs_metric) return None @@ -174,6 +276,7 @@ async def _stream_history_generator( history_length_before: int, guardrails_adapter: NeMoRailsAdapter, costs_metric: Dict[str, Dict[str, Any]], + request: OrchestrationRequest, ) -> AsyncIterator[str]: """Async generator: stream history answer through NeMo Guardrails.""" bot_generator = self.context_analyzer.stream_context_response( @@ -182,27 +285,71 @@ async def _stream_history_generator( orchestration_service = self.orchestration_service if orchestration_service is None: return - async for validated_chunk in guardrails_adapter.stream_with_guardrails( - user_message=query, bot_message_generator=bot_generator - ): - if isinstance(validated_chunk, str) and self._is_guardrail_violation( - validated_chunk + accumulated_response: list[str] = [] + with safe_observation_context( + as_type="generation", + name="context_workflow_streaming", + input={"query": query, "chat_id": chat_id}, + ) as _generation: + async for validated_chunk in guardrails_adapter.stream_with_guardrails( + user_message=query, bot_message_generator=bot_generator ): - logger.warning(f"[{chat_id}] Guardrails violation in context streaming") - yield orchestration_service.format_sse( - chat_id, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE - ) - yield orchestration_service.format_sse(chat_id, "END") - costs_metric["context_response"] = get_lm_usage_since( - history_length_before - ) - orchestration_service.log_costs(costs_metric) - return - yield orchestration_service.format_sse(chat_id, validated_chunk) - yield orchestration_service.format_sse(chat_id, "END") - logger.info(f"[{chat_id}] Context streaming complete") - costs_metric["context_response"] = get_lm_usage_since(history_length_before) - orchestration_service.log_costs(costs_metric) + if isinstance(validated_chunk, str) and self._is_guardrail_violation( + validated_chunk + ): + logger.warning( + f"[{chat_id}] Guardrails violation in context streaming" + ) + yield orchestration_service.format_sse( + chat_id, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE + ) + await orchestration_service.store_streaming_inference( + request, OUTPUT_GUARDRAIL_VIOLATION_MESSAGE + ) + yield orchestration_service.format_sse(chat_id, "END") + costs_metric["context_response"] = get_lm_usage_since( + history_length_before + ) + _usage = costs_metric["context_response"] + try: + if _generation is not None: + _generation.update( + usage_details={ + "input": _usage.get("total_prompt_tokens", 0), + "output": _usage.get("total_completion_tokens", 0), + "total": _usage.get("total_tokens", 0), + }, + cost_details={"total": _usage.get("total_cost", 0.0)}, + output={"guardrail_violation": True}, + ) + except Exception as _e: + logger.debug( + f"Langfuse streaming observation update skipped: {_e}" + ) + orchestration_service.log_costs(costs_metric) + return + accumulated_response.append(validated_chunk) + yield orchestration_service.format_sse(chat_id, validated_chunk) + final_answer = "".join(accumulated_response) + await orchestration_service.store_streaming_inference(request, final_answer) + yield orchestration_service.format_sse(chat_id, "END") + logger.info(f"[{chat_id}] Context streaming complete") + costs_metric["context_response"] = get_lm_usage_since(history_length_before) + _usage = costs_metric["context_response"] + try: + if _generation is not None: + _generation.update( + usage_details={ + "input": _usage.get("total_prompt_tokens", 0), + "output": _usage.get("total_completion_tokens", 0), + "total": _usage.get("total_tokens", 0), + }, + cost_details={"total": _usage.get("total_cost", 0.0)}, + output={"answer_preview": final_answer[:500]}, + ) + except Exception as _e: + logger.debug(f"Langfuse streaming observation update skipped: {_e}") + orchestration_service.log_costs(costs_metric) async def _create_history_stream( self, @@ -251,8 +398,10 @@ async def _create_history_stream( history_length_before=history_length_before, guardrails_adapter=guardrails_adapter, costs_metric=costs_metric, + request=request, ) + @observe(name="context_workflow_execute_async", as_type="span") async def execute_async( self, request: OrchestrationRequest, @@ -277,7 +426,7 @@ async def execute_async( time_metric = {} language = detect_language(request.message) - history = self._build_history(request) + history, pre_computed_summary = await self._build_history(request) # Check if analysis is pre-computed (e.g. from classifier classify step) pre_computed = context.get("analysis_result") @@ -295,9 +444,18 @@ async def execute_async( ) else: _detected = await self._detect( - request.message, history, time_metric, costs_metric + request.message, + history, + time_metric, + costs_metric, + pre_computed_summary, ) if _detected is None: + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) self._log_costs(costs_metric) context["costs_dict"] = costs_metric return None @@ -316,6 +474,11 @@ async def execute_async( ) self._log_costs(costs_metric) context["costs_dict"] = costs_metric + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "greeting_response"}, + metadata={"costs": costs_metric}, + ) return OrchestrationResponse( chatId=request.chatId, llmServiceActive=True, @@ -329,6 +492,11 @@ async def execute_async( and detection_result.context_snippet ): context["costs_dict"] = costs_metric + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "context_answer_generation"}, + metadata={"costs": costs_metric}, + ) return await self._generate_response_async( request, detection_result.context_snippet, time_metric, costs_metric ) @@ -336,10 +504,20 @@ async def execute_async( logger.warning( f"[{request.chatId}] Cannot answer from context — falling back to RAG" ) + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) self._log_costs(costs_metric) context["costs_dict"] = costs_metric return None + @observe( + name="context_workflow_execute_streaming", + as_type="span", + capture_output=False, + ) async def execute_streaming( self, request: OrchestrationRequest, @@ -364,12 +542,17 @@ async def execute_streaming( time_metric = {} language = detect_language(request.message) - history = self._build_history(request) + history, pre_computed_summary = await self._build_history(request) detection_result = await self._detect( - request.message, history, time_metric, costs_metric + request.message, history, time_metric, costs_metric, pre_computed_summary ) if detection_result is None: + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) self._log_costs(costs_metric) return None @@ -392,15 +575,26 @@ async def execute_streaming( async def _stream_greeting() -> AsyncIterator[str]: yield orchestration_service.format_sse(chat_id, greeting) + await orchestration_service.store_streaming_inference(request, greeting) yield orchestration_service.format_sse(chat_id, "END") orchestration_service.log_costs(costs_metric) + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "greeting_stream"}, + metadata={"costs": costs_metric}, + ) return _stream_greeting() if ( detection_result.can_answer_from_context and detection_result.context_snippet ): + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "context_stream_generation"}, + metadata={"costs": costs_metric}, + ) return await self._create_history_stream( request, detection_result.context_snippet, costs_metric ) @@ -408,5 +602,10 @@ async def _stream_greeting() -> AsyncIterator[str]: logger.warning( f"[{request.chatId}] Cannot answer from context — falling back to RAG" ) + update_observation_safe( + input_data={"chat_id": request.chatId, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) self._log_costs(costs_metric) return None diff --git a/src/tool_classifier/workflows/ood_workflow.py b/src/tool_classifier/workflows/ood_workflow.py index ed923879..4f4bab5f 100644 --- a/src/tool_classifier/workflows/ood_workflow.py +++ b/src/tool_classifier/workflows/ood_workflow.py @@ -1,11 +1,14 @@ """OOD workflow executor - Layer 4: Out-of-domain fallback.""" from typing import Any, AsyncIterator, Dict, Optional -from loguru import logger +from src.loki_logger import LokiLogger from models.request_models import OrchestrationRequest, OrchestrationResponse from tool_classifier.base_workflow import BaseWorkflow +# Initialize Loki logger +logger = LokiLogger(service_name="ood-workflow") + class OODWorkflowExecutor(BaseWorkflow): """ @@ -22,13 +25,6 @@ class OODWorkflowExecutor(BaseWorkflow): - "Tell me a joke" (not government service) - Questions with no relevant knowledge - Implementation Status: SKELETON - Returns None (will implement to return OOD message) - - TODO - Implementation (Simple): - - Return localized OUT_OF_SCOPE_MESSAGE - - Set questionOutOfLLMScope flag to True - - For streaming: chunk message and stream for UX consistency """ def __init__(self) -> None: @@ -80,7 +76,7 @@ async def execute_async( f"(not implemented - returning None for now)" ) - # TODO: Implement OOD response logic here + # Implement OOD response logic here # For now, return None (will be implemented as simple message return) return None @@ -129,6 +125,6 @@ async def stream_ood_message(): f"(not implemented - returning None for now)" ) - # TODO: Implement OOD streaming logic here + # Implement OOD streaming logic here # For now, return None (will be implemented as simple message streaming) return None diff --git a/src/tool_classifier/workflows/rag_workflow.py b/src/tool_classifier/workflows/rag_workflow.py index 9b3c588f..d84e5ae2 100644 --- a/src/tool_classifier/workflows/rag_workflow.py +++ b/src/tool_classifier/workflows/rag_workflow.py @@ -1,7 +1,7 @@ """RAG workflow executor - Layer 3: Knowledge base retrieval.""" from typing import Any, AsyncIterator, Dict, Optional, Union, cast, TYPE_CHECKING -from loguru import logger +from src.loki_logger import LokiLogger from models.request_models import ( OrchestrationRequest, @@ -14,6 +14,9 @@ if TYPE_CHECKING: from llm_orchestration_service import LLMOrchestrationService +# Initialize Loki logger +logger = LokiLogger(service_name="rag-workflow") + class RAGWorkflowExecutor(BaseWorkflow): """ diff --git a/src/tool_classifier/workflows/service_workflow.py b/src/tool_classifier/workflows/service_workflow.py index a021aba9..2b1f8f57 100644 --- a/src/tool_classifier/workflows/service_workflow.py +++ b/src/tool_classifier/workflows/service_workflow.py @@ -5,12 +5,16 @@ import dspy import httpx -from loguru import logger +from langfuse import observe +from src.loki_logger import LokiLogger from llm_orchestrator_config.llm_manager import LLMManager from src.guardrails.nemo_rails_adapter import NeMoRailsAdapter +from src.utils.conversation_history_helpers import get_conversation_history +from src.utils.conversation_history_store import ConversationHistoryStore from src.utils.cost_utils import get_lm_usage_since +from src.utils.observation_utils import update_observation_safe from models.request_models import ( ChoiceButton, @@ -38,6 +42,11 @@ from tool_classifier.intent_detector import IntentDetectionModule import time +# Initialize Loki logger +logger = LokiLogger(service_name="service-workflow") + +SERVICE_INTENT_DETECTION_METRIC = "service.intent_detection" + class LLMServiceProtocol(Protocol): """Protocol defining interface for LLM service embedding operations.""" @@ -104,6 +113,14 @@ async def handle_output_guardrails( """Apply output guardrails to the generated response.""" ... + async def store_streaming_inference( + self, + request: OrchestrationRequest, + final_answer: str, + ) -> None: + """Store streaming inference data for production/testing environments.""" + ... + class ServiceWorkflowExecutor(BaseWorkflow): """Executes external service calls via Ruuter endpoints (Layer 1).""" @@ -117,6 +134,12 @@ def __init__( self.llm_manager = llm_manager self.orchestration_service = orchestration_service + def _get_conversation_history_store(self) -> Optional[ConversationHistoryStore]: + """Return the conversation history store from the orchestration service, or None.""" + if self.orchestration_service is None: + return None + return getattr(self.orchestration_service, "conversation_history_store", None) + async def _semantic_search_services( self, query: str, @@ -243,15 +266,27 @@ async def _call_service_discovery(self, chat_id: str) -> Optional[Dict[str, Any] logger.error(f"[{chat_id}] Service discovery failed: {e}", exc_info=True) return None + @observe(name="service_intent_detection_orchestration", as_type="generation") async def _detect_service_intent( self, user_query: str, services: List[Dict[str, Any]], conversation_history: List[Any], chat_id: str, + conversation_summary: Optional[str] = None, ) -> tuple[Optional[Dict[str, Any]], Dict[str, Any]]: """Use DSPy + LLMManager to detect service intent and extract entities. + Args: + user_query: The user's query string. + services: List of available service dicts. + conversation_history: Recent conversation turns (``ConversationItem`` objects). + chat_id: Chat identifier for logging. + conversation_summary: Optional summary of earlier conversation rounds + evicted from Redis. When provided it is prepended to the history + passed to the intent detector so the LLM has additional context + without a separate summarisation call. + Returns: Tuple of (intent_result, usage_info): - intent_result: Intent detection result dict (or None on error) @@ -270,11 +305,21 @@ async def _detect_service_intent( ) intent_module = IntentDetectionModule() - history_dicts = [ - {"authorRole": msg.authorRole, "message": msg.message} - for msg in conversation_history - if hasattr(msg, "authorRole") and hasattr(msg, "message") - ] + history_dicts: List[Dict[str, str]] = [] + if conversation_summary: + history_dicts.append( + { + "authorRole": "system", + "message": f"Summary of earlier conversation: {conversation_summary}", + } + ) + history_dicts.extend( + [ + {"authorRole": msg.authorRole, "message": msg.message} + for msg in conversation_history + if hasattr(msg, "authorRole") and hasattr(msg, "message") + ] + ) with self.llm_manager.use_task_local(): intent_result = intent_module.forward( @@ -284,11 +329,38 @@ async def _detect_service_intent( ) usage_info = get_lm_usage_since(history_length_before) + update_observation_safe( + input_data={ + "chat_id": chat_id, + "query": user_query, + "services_count": len(services), + }, + output_data={ + "matched_service_id": ( + intent_result.get("matched_service_id") + if intent_result + else None + ), + "confidence": intent_result.get("confidence", 0.0) + if intent_result + else 0.0, + }, + metadata={"usage": usage_info}, + ) return intent_result, usage_info except Exception as e: logger.error(f"[{chat_id}] Intent detection failed: {e}", exc_info=True) + update_observation_safe( + input_data={ + "chat_id": chat_id, + "query": user_query, + "services_count": len(services), + }, + output_data={"matched_service_id": None, "error": str(e)}, + metadata={"usage": {}}, + ) return None, {} def _validate_detected_service( @@ -331,11 +403,17 @@ async def _process_intent_detection( context: Context dict to populate with results costs_metric: Dictionary to track LLM costs """ + conversation_history, conversation_summary = await get_conversation_history( + chat_id=request.chatId, + store=self._get_conversation_history_store(), + fallback=request.conversationHistory, + ) intent_result, intent_usage = await self._detect_service_intent( user_query=request.message, services=services, - conversation_history=request.conversationHistory, + conversation_history=conversation_history, chat_id=chat_id, + conversation_summary=conversation_summary, ) costs_metric["intent_detection"] = intent_usage @@ -698,6 +776,7 @@ async def _log_request_details( else: logger.warning(f"[{chat_id}] Service discovery failed") + @observe(name="service_workflow_execute_async", as_type="span") async def execute_async( self, request: OrchestrationRequest, @@ -746,7 +825,7 @@ async def execute_async( context=context, costs_metric=costs_metric, ) - time_metric["service.intent_detection"] = time.time() - start_time + time_metric[SERVICE_INTENT_DETECTION_METRIC] = time.time() - start_time if not context.get("service_data"): context["service_id"] = matched.get("service_id") @@ -768,7 +847,7 @@ async def execute_async( context=context, costs_metric=costs_metric, ) - time_metric["service.intent_detection"] = time.time() - start_time + time_metric[SERVICE_INTENT_DETECTION_METRIC] = time.time() - start_time else: start_time = time.time() @@ -779,11 +858,21 @@ async def execute_async( if not context.get("service_id"): logger.info(f"[{chat_id}] No service matched, falling back") + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) return None start_time = time.time() service_metadata = self._extract_service_metadata(context, chat_id) if not service_metadata: + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={"workflow_result": "missing_service_metadata"}, + metadata={"costs": costs_metric}, + ) return None logger.info( @@ -835,8 +924,22 @@ async def execute_async( if service_result is None: logger.warning(f"[{chat_id}] Service endpoint call failed, falling back") + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) return None + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={ + "workflow_result": "service_response", + "service_id": context.get("service_id"), + }, + metadata={"costs": costs_metric}, + ) + service_content = service_result["content"] service_buttons = service_result["buttons"] buttons_list = [ @@ -854,6 +957,11 @@ async def execute_async( buttons=buttons_list if buttons_list else None, ) + @observe( + name="service_workflow_execute_streaming", + as_type="span", + capture_output=False, + ) async def execute_streaming( self, request: OrchestrationRequest, @@ -899,7 +1007,7 @@ async def execute_streaming( context=context, costs_metric=costs_metric, ) - time_metric["service.intent_detection"] = time.time() - start_time + time_metric[SERVICE_INTENT_DETECTION_METRIC] = time.time() - start_time if not context.get("service_data"): context["service_id"] = matched.get("service_id") @@ -921,7 +1029,7 @@ async def execute_streaming( context=context, costs_metric=costs_metric, ) - time_metric["service.intent_detection"] = time.time() - start_time + time_metric[SERVICE_INTENT_DETECTION_METRIC] = time.time() - start_time else: start_time = time.time() @@ -932,10 +1040,20 @@ async def execute_streaming( if not context.get("service_id"): logger.info(f"[{chat_id}] No service matched, falling back") + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) return None service_metadata = self._extract_service_metadata(context, chat_id) if not service_metadata: + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={"workflow_result": "missing_service_metadata"}, + metadata={"costs": costs_metric}, + ) return None logger.info( @@ -981,6 +1099,11 @@ async def execute_streaming( if service_result is None: logger.warning(f"[{chat_id}] Service endpoint call failed, falling back") + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={"workflow_result": "fallback_to_rag"}, + metadata={"costs": costs_metric}, + ) return None if self.orchestration_service is None: @@ -994,9 +1117,20 @@ async def service_stream() -> AsyncIterator[str]: yield orchestration_service.format_sse( chat_id, service_content, service_buttons or None ) + await orchestration_service.store_streaming_inference( + request, service_content + ) yield orchestration_service.format_sse(chat_id, "END") orchestration_service.log_costs(costs_metric) + update_observation_safe( + input_data={"chat_id": chat_id, "query": request.message}, + output_data={ + "workflow_result": "service_stream", + "service_id": context.get("service_id"), + }, + metadata={"costs": costs_metric}, + ) return service_stream() async def execute_direct_step( @@ -1122,6 +1256,9 @@ async def step_stream() -> AsyncIterator[str]: yield orchestration_service.format_sse( chat_id, service_content, service_buttons or None ) + await orchestration_service.store_streaming_inference( + request, service_content + ) yield orchestration_service.format_sse(chat_id, "END") return step_stream() diff --git a/src/utils/api_tool_session_store.py b/src/utils/api_tool_session_store.py index 3a06e5ec..807e6de8 100644 --- a/src/utils/api_tool_session_store.py +++ b/src/utils/api_tool_session_store.py @@ -3,12 +3,14 @@ from typing import Any, Optional from fastapi import HTTPException, Request, status -from loguru import logger +from src.loki_logger import LokiLogger from redis import WatchError -from src.models.session_models import APIToolSession +from models.session_models import APIToolSession from src.utils.redis_client import get_redis_client +logger = LokiLogger(service_name="api-tool-session-store") + _SESSION_KEY_PREFIX = "session:" _SESSION_TTL_SECONDS = 1800 # 30 minutes, sliding _UPDATE_MAX_RETRIES = 3 @@ -36,9 +38,7 @@ async def get(self, chat_id: str) -> Optional[APIToolSession]: """ client = get_redis_client() if client is None: - logger.warning( - "[SessionStore] Redis unavailable - get({}) skipped", chat_id - ) + logger.warning(f"[SessionStore] Redis unavailable - get({chat_id}) skipped") return None try: @@ -47,7 +47,7 @@ async def get(self, chat_id: str) -> Optional[APIToolSession]: return None return APIToolSession.model_validate_json(raw) except Exception as exc: - logger.error("[SessionStore] get({}) failed: {}", chat_id, exc) + logger.error(f"[SessionStore] get({chat_id}) failed: {exc}") return None async def save(self, session: APIToolSession) -> None: @@ -59,7 +59,7 @@ async def save(self, session: APIToolSession) -> None: client = get_redis_client() if client is None: logger.warning( - "[SessionStore] Redis unavailable - save({}) skipped", session.chat_id + f"[SessionStore] Redis unavailable - save({session.chat_id}) skipped" ) return @@ -69,9 +69,9 @@ async def save(self, session: APIToolSession) -> None: session.model_dump_json(), ex=_SESSION_TTL_SECONDS, ) - logger.debug("[SessionStore] Session saved for chat_id={}", session.chat_id) + logger.debug(f"[SessionStore] Session saved for chat_id={session.chat_id}") except Exception as exc: - logger.error("[SessionStore] save({}) failed: {}", session.chat_id, exc) + logger.error(f"[SessionStore] save({session.chat_id}) failed: {exc}") async def update(self, chat_id: str, **fields: Any) -> Optional[APIToolSession]: """Atomically update a session using optimistic locking (WATCH/MULTI/EXEC). @@ -97,7 +97,7 @@ async def update(self, chat_id: str, **fields: Any) -> Optional[APIToolSession]: client = get_redis_client() if client is None: logger.warning( - "[SessionStore] Redis unavailable - update({}) skipped", chat_id + f"[SessionStore] Redis unavailable - update({chat_id}) skipped" ) return None @@ -112,8 +112,7 @@ async def update(self, chat_id: str, **fields: Any) -> Optional[APIToolSession]: if raw is None: await pipe.unwatch() logger.warning( - "[SessionStore] update({}) - session not found, skipping", - chat_id, + f"[SessionStore] update({chat_id}) - session not found, skipping" ) return None @@ -125,27 +124,22 @@ async def update(self, chat_id: str, **fields: Any) -> Optional[APIToolSession]: await pipe.execute() logger.debug( - "[SessionStore] Session updated for chat_id={}", chat_id + f"[SessionStore] Session updated for chat_id={chat_id}" ) return updated except WatchError: logger.debug( - "[SessionStore] update({}) - concurrent modification detected, " - "retrying (attempt {}/{})", - chat_id, - attempt + 1, - _UPDATE_MAX_RETRIES, + f"[SessionStore] update({chat_id}) - concurrent modification detected, " + f"retrying (attempt {attempt + 1}/{_UPDATE_MAX_RETRIES})" ) continue except Exception as exc: - logger.error("[SessionStore] update({}) failed: {}", chat_id, exc) + logger.error(f"[SessionStore] update({chat_id}) failed: {exc}") return None logger.error( - "[SessionStore] update({}) - exhausted {} retries due to concurrent writes", - chat_id, - _UPDATE_MAX_RETRIES, + f"[SessionStore] update({chat_id}) - exhausted {_UPDATE_MAX_RETRIES} retries due to concurrent writes" ) return None @@ -158,15 +152,15 @@ async def delete(self, chat_id: str) -> None: client = get_redis_client() if client is None: logger.warning( - "[SessionStore] Redis unavailable - delete({}) skipped", chat_id + f"[SessionStore] Redis unavailable - delete({chat_id}) skipped" ) return try: await client.delete(_key(chat_id)) - logger.debug("[SessionStore] Session deleted for chat_id={}", chat_id) + logger.debug(f"[SessionStore] Session deleted for chat_id={chat_id}") except Exception as exc: - logger.error("[SessionStore] delete({}) failed: {}", chat_id, exc) + logger.error(f"[SessionStore] delete({chat_id}) failed: {exc}") async def exists(self, chat_id: str) -> bool: """Check whether a session exists for the given chat_id. @@ -181,7 +175,7 @@ async def exists(self, chat_id: str) -> bool: try: return bool(await client.exists(_key(chat_id))) except Exception as exc: - logger.error("[SessionStore] exists({}) failed: {}", chat_id, exc) + logger.error(f"[SessionStore] exists({chat_id}) failed: {exc}") return False @@ -197,8 +191,7 @@ def require_session_store(request: Request) -> APIToolSessionStore: ) if store is None: logger.error( - "[SessionStore] Session store unavailable — returning 503 for {}", - request.url.path, + f"[SessionStore] Session store unavailable — returning 503 for {request.url.path}" ) raise HTTPException( status_code=status.HTTP_503_SERVICE_UNAVAILABLE, diff --git a/src/utils/atc_cache_store.py b/src/utils/atc_cache_store.py new file mode 100644 index 00000000..9e203709 --- /dev/null +++ b/src/utils/atc_cache_store.py @@ -0,0 +1,197 @@ +"""Two-tier Redis response cache for the ATC workflow.""" + +import hashlib +import json +from typing import Any + +from src.loki_logger import LokiLogger + +from models.session_models import LastCallContext +from tool_classifier.constants import ( + ATC_CACHE_DEFAULT_TTL_SECONDS, + ATC_CACHE_KEY_PREFIX, + ATC_LAST_CALL_KEY_PREFIX, + ATC_LAST_CALL_TTL_SECONDS, +) +from src.utils.redis_client import get_redis_client + +logger = LokiLogger(service_name="atc-cache-store") + + +class ATCCacheStore: + """Two-tier response cache for the ATC workflow. + + Tier 1 (L1) — Exact response cache + Key: atc:cache:{chat_id}:{api_name}:{param_hash} + Value: raw API response JSON + TTL: per-endpoint ``cache_ttl_seconds`` or ``ATC_CACHE_DEFAULT_TTL_SECONDS`` + + Tier 2 (L2) — Last call context + Key: atc:last:{chat_id} + Value: JSON ``list[LastCallContext]`` + TTL: ``ATC_LAST_CALL_TTL_SECONDS`` (sliding — reset on every write) + + All public methods are async and fail-open: a Redis unavailability or + serialisation error is logged as a warning and the caller receives + ``None`` / silent no-op rather than an exception. + """ + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + @staticmethod + def _normalise_params(params: dict[str, Any]) -> dict[str, Any]: + """Normalise param values so equivalent inputs produce the same hash. + + Rules applied per string value: + - Strip leading/trailing whitespace. + - If the stripped string is purely numeric (``"2026"``), cast to ``int``. + - If the stripped string is all-alpha with no spaces (enum-like), lowercase it. + """ + result: dict[str, Any] = {} + for k, v in params.items(): + if isinstance(v, str): + stripped = v.strip() + if stripped.isdigit(): + result[k] = int(stripped) + elif stripped.isalpha(): + result[k] = stripped.lower() + else: + result[k] = stripped + else: + result[k] = v + return result + + @staticmethod + def _param_hash(params: dict[str, Any]) -> str: + """Return a 16-char hex digest of the normalised, sorted params dict.""" + normalised = ATCCacheStore._normalise_params(params) + serialised = json.dumps(normalised, sort_keys=True) + return hashlib.sha256(serialised.encode()).hexdigest()[:16] + + @staticmethod + def _l1_key(chat_id: str, api_name: str, params: dict[str, Any]) -> str: + return f"{ATC_CACHE_KEY_PREFIX}:{chat_id}:{api_name}:{ATCCacheStore._param_hash(params)}" + + @staticmethod + def _l2_key(chat_id: str) -> str: + return f"{ATC_LAST_CALL_KEY_PREFIX}:{chat_id}" + + # ------------------------------------------------------------------ + # L1 — Exact response cache + # ------------------------------------------------------------------ + + async def get_l1( + self, chat_id: str, api_name: str, params: dict[str, Any] + ) -> Any | None: + """Return the cached raw API response, or ``None`` on cache miss or error.""" + client = get_redis_client() + if client is None: + return None + try: + raw = await client.get(self._l1_key(chat_id, api_name, params)) + if raw is None: + return None + return json.loads(raw) + except Exception as exc: + logger.warning( + f"[ATCCache] get_l1 failed: chat_id={chat_id} api_name={api_name!r} error={exc}" + ) + return None + + async def set_l1( + self, + chat_id: str, + api_name: str, + params: dict[str, Any], + raw_response: Any, + ttl: int = ATC_CACHE_DEFAULT_TTL_SECONDS, + ) -> None: + """Serialise and store a raw API response in L1. + + Args: + chat_id: Conversation identifier. + api_name: Endpoint name (snake_case, matches ``EnrichedEndpoint.name``). + params: The collected param values used for this API call. + raw_response: Parsed API JSON (dict or list). + ttl: Cache TTL in seconds. Defaults to ``ATC_CACHE_DEFAULT_TTL_SECONDS``. + """ + client = get_redis_client() + if client is None: + return + try: + await client.set( + self._l1_key(chat_id, api_name, params), + json.dumps(raw_response), + ex=ttl, + ) + logger.debug( + f"[ATCCache] L1 set: chat_id={chat_id} api_name={api_name!r} ttl={ttl}s" + ) + except Exception as exc: + logger.warning( + f"[ATCCache] set_l1 failed: chat_id={chat_id} api_name={api_name!r} error={exc}" + ) + + # ------------------------------------------------------------------ + # L2 — Last call context + # ------------------------------------------------------------------ + + async def get_l2(self, chat_id: str) -> list[LastCallContext] | None: + """Return the last-call context list, or ``None`` on miss or error.""" + client = get_redis_client() + if client is None: + return None + try: + raw = await client.get(self._l2_key(chat_id)) + if raw is None: + return None + data: list[dict[str, Any]] = json.loads(raw) + return [LastCallContext.model_validate(item) for item in data] + except Exception as exc: + logger.warning(f"[ATCCache] get_l2 failed: chat_id={chat_id} error={exc}") + return None + + async def set_l2(self, chat_id: str, contexts: list[LastCallContext]) -> None: + """Serialise and store the last-call context list in L2. + + The TTL is reset on every write (sliding expiry) so an active conversation + always has a live L2 entry while the session exists. + + For single-intent calls pass a one-element list; for multi-intent calls + pass one ``LastCallContext`` per succeeded endpoint. + """ + client = get_redis_client() + if client is None: + return + try: + payload = json.dumps([ctx.model_dump() for ctx in contexts]) + await client.set( + self._l2_key(chat_id), + payload, + ex=ATC_LAST_CALL_TTL_SECONDS, + ) + logger.debug( + f"[ATCCache] L2 set: chat_id={chat_id} entries={len(contexts)}" + ) + except Exception as exc: + logger.warning(f"[ATCCache] set_l2 failed: chat_id={chat_id} error={exc}") + + async def invalidate_l2(self, chat_id: str) -> None: + """Delete the L2 last-call context key for a chat session. + + Called on intent switch so the follow-up detector does not carry stale + context into a brand-new query. L1 entries are **not** deleted — they + are param-hash scoped and expire on their own TTL. + """ + client = get_redis_client() + if client is None: + return + try: + await client.delete(self._l2_key(chat_id)) + logger.debug(f"[ATCCache] L2 invalidated: chat_id={chat_id}") + except Exception as exc: + logger.warning( + f"[ATCCache] invalidate_l2 failed: chat_id={chat_id} error={exc}" + ) diff --git a/src/utils/budget_tracker.py b/src/utils/budget_tracker.py index 48507de3..58d356b2 100644 --- a/src/utils/budget_tracker.py +++ b/src/utils/budget_tracker.py @@ -1,68 +1,41 @@ """Budget tracking utility for LLM connection usage.""" from typing import Optional, Dict, Any, cast, List -from loguru import logger +from src.loki_logger import LokiLogger import requests from ..llm_orchestrator_config.llm_ochestrator_constants import RAG_SEARCH_RESQL -from .connection_id_fetcher import get_connection_id_fetcher + +# Initialize Loki logger +logger = LokiLogger(service_name="budget-tracker") class BudgetTracker: - """Handles budget updates for LLM connections.""" + """Handles budget updates for LLM connections using vault_uuid.""" def __init__(self) -> None: - """Initialize the budget tracker with Resql and Ruuter endpoints.""" + """Initialize the budget tracker with Resql endpoint.""" # Use Resql directly for budget updates self.resql_base = RAG_SEARCH_RESQL self.update_endpoint = f"{self.resql_base}/update-llm-connection-used-budget" self.timeout = 5 # seconds - # Use centralized connection ID fetcher - self.connection_fetcher = get_connection_id_fetcher() - - def _validate_connection_id(self, connection_id: Optional[str]) -> Optional[int]: - """ - Validate and convert connection_id to integer. - - Args: - connection_id: The connection ID to validate - - Returns: - Integer connection ID, or None if invalid - """ - if not connection_id: - logger.debug("No connection_id provided, skipping budget update") - return None - - try: - return int(connection_id) - except (ValueError, TypeError): - logger.warning( - f"Connection ID '{connection_id}' is not numeric. " - f"Budget tracking requires numeric database IDs. " - f"Skipping budget update for this request." - ) - return None - def _make_budget_update_request( - self, connection_id_int: int, usage_cost: float + self, vault_uuid: str, usage_cost: float ) -> Dict[str, Any]: """ Make the actual budget update API request. Args: - connection_id_int: The integer connection ID + vault_uuid: The vault UUID identifying the connection usage_cost: The cost to add Returns: Dictionary containing the response or error """ - payload = {"connection_id": connection_id_int, "usage": usage_cost} - logger.info( - f"Updating budget for connection_id={connection_id_int}, usage={usage_cost}" - ) + payload = {"vault_uuid": vault_uuid, "usage": usage_cost} + logger.info(f"Updating budget for vault_uuid={vault_uuid}, usage={usage_cost}") response = requests.post( self.update_endpoint, json=payload, timeout=self.timeout @@ -82,10 +55,6 @@ def _make_budget_update_request( else: data = response_data - logger.info( - f"Budget updated successfully for connection_id={connection_id_int}" - ) - # Check if budget was exceeded budget_exceeded: bool = False if isinstance(data, dict): @@ -96,7 +65,7 @@ def _make_budget_update_request( if budget_exceeded: logger.warning( - f"Budget threshold exceeded for connection_id={connection_id_int}. " + f"Budget threshold exceeded for vault_uuid={vault_uuid}. " f"Connection may have been deactivated." ) @@ -107,7 +76,7 @@ def _make_budget_update_request( } else: logger.error( - f"Failed to update budget for connection_id={connection_id_int}. " + f"Failed to update budget for vault_uuid={vault_uuid}. " f"Status: {response.status_code}, Response: {response.text}" ) return { @@ -124,38 +93,21 @@ def update_budget( Update the used budget for an LLM connection. Args: - connection_id: The LLM connection ID (can be numeric ID or string identifier) + connection_id: The vault_uuid identifying the LLM connection usage_cost: The cost to add to the used budget Returns: Dictionary containing the response from the update endpoint or an error indicator if the update failed """ - # If no connection ID provided, try to fetch production connection ID + # Validate connection_id (vault_uuid) is provided if not connection_id: logger.debug( - "No connection_id provided, attempting to fetch production connection ID" + "No connection_id (vault_uuid) provided, skipping budget update" ) - try: - fetched_id = self.connection_fetcher.fetch_connection_id_sync( - "production" - ) - if fetched_id is not None: - connection_id = str(fetched_id) - logger.debug( - f"Using fetched production connection_id: {connection_id}" - ) - except Exception as e: - logger.warning(f"Failed to fetch production connection ID: {str(e)}") - - # Validate connection_id - connection_id_int = self._validate_connection_id(connection_id) - if connection_id_int is None: return { "success": False, - "reason": "no_connection_id" - if not connection_id - else "non_numeric_connection_id", + "reason": "no_connection_id", "connection_id": connection_id, } @@ -165,23 +117,23 @@ def update_budget( return {"success": False, "reason": "zero_or_negative_cost"} try: - return self._make_budget_update_request(connection_id_int, usage_cost) + return self._make_budget_update_request(connection_id, usage_cost) except requests.exceptions.Timeout: logger.error( - f"Timeout while updating budget for connection_id={connection_id}" + f"Timeout while updating budget for vault_uuid={connection_id}" ) return {"success": False, "reason": "timeout"} except requests.exceptions.RequestException as e: logger.error( - f"Request error while updating budget for connection_id={connection_id}: {str(e)}" + f"Request error while updating budget for vault_uuid={connection_id}: {str(e)}" ) return {"success": False, "reason": "request_error", "error": str(e)} except Exception as e: logger.error( - f"Unexpected error while updating budget for connection_id={connection_id}: {str(e)}" + f"Unexpected error while updating budget for vault_uuid={connection_id}: {str(e)}" ) return {"success": False, "reason": "unexpected_error", "error": str(e)} diff --git a/src/utils/connection_id_fetcher.py b/src/utils/connection_id_fetcher.py index f9e17cca..ac17f6cc 100644 --- a/src/utils/connection_id_fetcher.py +++ b/src/utils/connection_id_fetcher.py @@ -8,12 +8,15 @@ import asyncio import threading from typing import Optional, Dict, Any -from loguru import logger +from src.loki_logger import LokiLogger import requests import aiohttp from src.llm_orchestrator_config.llm_ochestrator_constants import RAG_SEARCH_RESQL +# Initialize Loki logger +logger = LokiLogger(service_name="connection-id-fetcher") + class ConnectionIdFetcher: """ @@ -29,8 +32,8 @@ def __init__(self) -> None: self.resql_base = RAG_SEARCH_RESQL self.timeout = 5 # seconds - # Cache connection IDs to avoid repeated requests - self._connection_cache: Dict[str, int] = {} + # Cache connection IDs and vault UUIDs to avoid repeated requests + self._connection_cache: Dict[str, int | str] = {} # Thread-safe lock for cache access self._cache_lock = threading.Lock() @@ -87,7 +90,7 @@ def fetch_connection_id_sync(self, environment: str) -> Optional[int]: # Thread-safe cache check with self._cache_lock: if cache_key in self._connection_cache: - cached_value = self._connection_cache[cache_key] + cached_value = int(self._connection_cache[cache_key]) logger.debug( f"Using cached connection_id for {environment}: {cached_value}" ) @@ -147,7 +150,7 @@ async def fetch_connection_id_async(self, environment: str) -> Optional[int]: # Thread-safe cache check with self._cache_lock: if cache_key in self._connection_cache: - cached_value = self._connection_cache[cache_key] + cached_value = int(self._connection_cache[cache_key]) logger.debug( f"Using cached connection_id for {environment}: {cached_value}" ) @@ -212,13 +215,97 @@ def clear_cache(self, environment: Optional[str] = None) -> None: with self._cache_lock: if environment: cache_key = f"{environment}_connection_id" + vault_key = f"{environment}_vault_uuid" if cache_key in self._connection_cache: del self._connection_cache[cache_key] - logger.debug(f"Cleared cache for {environment} connection_id") + if vault_key in self._connection_cache: + del self._connection_cache[vault_key] + logger.debug(f"Cleared cache for {environment}") else: self._connection_cache.clear() logger.debug("Cleared all connection_id cache") + def fetch_vault_uuid_sync(self, environment: str) -> Optional[str]: + """ + Synchronously fetch the vault_uuid for specified environment. + + Args: + environment: The deployment environment ("production" or "testing") + + Returns: + The vault_uuid (string) or None if unavailable + """ + cache_key = f"{environment}_vault_uuid" + + with self._cache_lock: + if cache_key in self._connection_cache: + cached_value = self._connection_cache[cache_key] + logger.debug( + f"Using cached vault_uuid for {environment}: {cached_value}" + ) + return str(cached_value) + + try: + logger.debug(f"Fetching {environment} vault_uuid from Resql (sync)...") + + endpoint = f"{self.resql_base}/get-{environment}-connection" + response = requests.post(endpoint, json={}, timeout=self.timeout) + + if response.status_code == 200: + data = response.json() + vault_uuid = self._extract_vault_uuid_from_response(data) + + if vault_uuid is not None: + with self._cache_lock: + self._connection_cache[cache_key] = vault_uuid + logger.info( + f"{environment.capitalize()} vault_uuid fetched: {vault_uuid}" + ) + return vault_uuid + else: + logger.warning(f"No {environment} vault_uuid found in response") + return None + else: + logger.error( + f"Failed to fetch {environment} vault_uuid. " + f"Status: {response.status_code}, Response: {response.text}" + ) + return None + + except requests.exceptions.Timeout: + logger.error(f"Timeout while fetching {environment} vault_uuid") + return None + except Exception as e: + logger.error(f"Error fetching {environment} vault_uuid: {str(e)}") + return None + + def _extract_vault_uuid_from_response( + self, data: dict[str, Any] | list[Any] + ) -> Optional[str]: + """ + Extract vault_uuid from API response data. + + Args: + data: The JSON response data (dict or list) + + Returns: + The vault_uuid as string, or None if not found + """ + if isinstance(data, dict): + response_data: Any = data.get("response", data) + else: + response_data = data + + if isinstance(response_data, list): + if len(response_data) > 0 and isinstance(response_data[0], dict): + vault_uuid = response_data[0].get("vaultUuid") + return str(vault_uuid) if vault_uuid else None + elif isinstance(response_data, dict): + vault_uuid = response_data.get("vaultUuid") + return str(vault_uuid) if vault_uuid else None + + return None + # Singleton instance for reuse across modules _connection_id_fetcher: Optional[ConnectionIdFetcher] = None diff --git a/src/utils/conversation_history_helpers.py b/src/utils/conversation_history_helpers.py new file mode 100644 index 00000000..7ceba47b --- /dev/null +++ b/src/utils/conversation_history_helpers.py @@ -0,0 +1,78 @@ +"""Shared helper for fetching conversation history from Redis. + +Mirrors the pattern established by ``ContextWorkflowExecutor._build_history()`` +so that all workflow entry points (RAG, service, ATC) resolve history in the +same way: Redis is preferred over the GUI-supplied request payload, and any +running conversation summary stored alongside the rounds is surfaced to callers +so they can skip an expensive LLM summarisation step. +""" + +from typing import List, Optional + +from src.loki_logger import LokiLogger + +from src.models.conversation_history_models import ConversationHistoryState +from models.request_models import ConversationItem +from src.utils.conversation_history_store import ConversationHistoryStore + +logger = LokiLogger(service_name="conversation-history") + + +async def get_conversation_history( + chat_id: str, + store: Optional[ConversationHistoryStore], + fallback: List[ConversationItem], +) -> tuple[List[ConversationItem], Optional[str]]: + """Fetch conversation history, preferring Redis over the request payload. + + When *store* is provided and the Redis session has rounds, those rounds are + returned as the authoritative history and *fallback* is ignored. Any running + summary attached to the stored state is returned as the second tuple element + so that callers can skip an expensive LLM summarisation step. + + If the store is absent, raises, or contains no rounds the function returns + *fallback* with ``None`` as the summary. + + Args: + chat_id: The conversation identifier. + store: Optional Redis-backed conversation history store. + fallback: ``request.conversationHistory`` — used when Redis is unavailable + or has no rounds for this session. + + Returns: + ``(history, summary)`` where *history* is a list of + :class:`~src.models.request_models.ConversationItem` objects (two per + stored round: one ``"user"`` and one ``"bot"`` item) and *summary* is + the Redis summary string or ``None``. + """ + if store is not None: + try: + state: ConversationHistoryState = await store.get_context(chat_id) + if state.rounds: + history: List[ConversationItem] = [] + for round_ in state.rounds: + history.append( + ConversationItem( + authorRole="user", + message=round_.user_message, + timestamp=str(round_.timestamp), + ) + ) + history.append( + ConversationItem( + authorRole="bot", + message=round_.bot_message, + timestamp=str(round_.timestamp), + ) + ) + logger.debug( + f"[{chat_id}] Using Redis history: {len(state.rounds)} rounds, " + f"summary={'present' if state.summary else 'absent'}" + ) + return history, state.summary + except Exception as exc: + logger.warning( + f"[{chat_id}] Redis history fetch failed, falling back to request history: {exc}" + ) + + return fallback, None diff --git a/src/utils/conversation_history_store.py b/src/utils/conversation_history_store.py new file mode 100644 index 00000000..85fb84a6 --- /dev/null +++ b/src/utils/conversation_history_store.py @@ -0,0 +1,364 @@ +"""Redis-backed conversation history store.""" + +import asyncio +import json +import weakref +from typing import TYPE_CHECKING, Optional, Union, cast + +from src.loki_logger import LokiLogger +from redis import WatchError + +from src.models.conversation_history_models import ( + ConversationHistoryState, + ConversationRound, +) +from src.utils.redis_client import get_redis_client + +if TYPE_CHECKING: + from src.models.request_models import ( + OrchestrationResponse, + TestOrchestrationResponse, + ) + from src.utils.conversation_summary_generator import SummarizerCallable + +logger = LokiLogger(service_name="conversation-history") + +_HISTORY_KEY_PREFIX = "conv:" +_SUMMARY_KEY_PREFIX = "conv:summary:" +_HISTORY_TTL_SECONDS = 1800 # 30 minutes, sliding +_MAX_ROUNDS = 10 +_APPEND_MAX_RETRIES = 3 + + +def _history_key(chat_id: str) -> str: + return f"{_HISTORY_KEY_PREFIX}{chat_id}" + + +def _summary_key(chat_id: str) -> str: + return f"{_SUMMARY_KEY_PREFIX}{chat_id}" + + +class ConversationHistoryStore: + """CRUD store for per-session conversation history backed by Redis. + + All operations are async and safe to call from FastAPI handlers. + The TTL is reset (sliding expiry) on every write that touches a key. + + Key layout (db=1, same as session store): + ``conv:{chat_id}`` — JSON list of up to 10 ``ConversationRound`` objects + ``conv:summary:{chat_id}`` — plain string summary (optional) + + An optional *summarizer* callable is injected at construction time. When + trimming evicts rounds (``len(rounds) > _MAX_ROUNDS``), a background + ``asyncio.Task`` is created to merge the evicted rounds into the existing + summary via the summarizer. If *summarizer* is ``None``, trimming still + occurs but no summary is generated. + """ + + def __init__( + self, + summarizer: Optional["SummarizerCallable"] = None, + ) -> None: + """Initialise the store. + + Args: + summarizer: Optional async callable that merges evicted rounds into + the running conversation summary. See + :func:`~src.utils.conversation_summary_generator.create_incremental_summarizer` + for a factory that creates one. + """ + self._summarizer = summarizer + # Hold strong references to background tasks to prevent GC collection + # before they complete (asyncio tasks are only weakly referenced by the + # event loop). + self._pending_tasks: set[asyncio.Task[None]] = set() + # Per-chat locks to serialize summary updates and prevent concurrent + # write races when multiple save_round() calls trigger evictions. + # WeakValueDictionary allows entries to be GC'd once no task holds a + # strong reference, preventing unbounded growth under high-cardinality + # chat_ids. + self._summary_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = ( + weakref.WeakValueDictionary() + ) + + async def save_round(self, chat_id: str, round: ConversationRound) -> None: + """Append a round to the history, trim to ``_MAX_ROUNDS``, reset both TTLs. + + Uses optimistic locking (WATCH/MULTI/EXEC) to detect concurrent writes and + retries up to ``_APPEND_MAX_RETRIES`` times on conflict. + + Args: + chat_id: The conversation identifier. + round: The completed user+bot exchange to persist. + """ + client = get_redis_client() + if client is None: + logger.warning( + f"[ConversationHistoryStore] Redis unavailable - save_round({chat_id}) skipped" + ) + return + + hkey = _history_key(chat_id) + skey = _summary_key(chat_id) + + for attempt in range(_APPEND_MAX_RETRIES): + try: + async with client.pipeline(transaction=True) as pipe: + await pipe.watch(hkey) + + raw = await pipe.get(hkey) + if raw is not None: + rounds: list[dict] = json.loads(raw) + else: + rounds = [] + + rounds.append(round.model_dump()) + + # Capture rounds that will be evicted before trimming. + evicted: list[ConversationRound] = [] + if len(rounds) > _MAX_ROUNDS: + evicted = [ + ConversationRound.model_validate(r) + for r in rounds[: len(rounds) - _MAX_ROUNDS] + ] + rounds = rounds[-_MAX_ROUNDS:] + + pipe.multi() + pipe.set(hkey, json.dumps(rounds), ex=_HISTORY_TTL_SECONDS) + # Reset summary TTL without overwriting its value + pipe.expire(skey, _HISTORY_TTL_SECONDS) + await pipe.execute() + + logger.debug( + f"[ConversationHistoryStore] Round saved for chat_id={chat_id}" + ) + + # Fire-and-forget incremental summary generation for evicted rounds. + if evicted and self._summarizer is not None: + task = asyncio.create_task( + self._run_incremental_summary(chat_id, evicted) + ) + self._pending_tasks.add(task) + task.add_done_callback(self._pending_tasks.discard) + + return + + except WatchError: + logger.debug( + f"[ConversationHistoryStore] save_round({chat_id}) - concurrent modification " + f"detected, retrying (attempt {attempt + 1}/{_APPEND_MAX_RETRIES})" + ) + continue + except Exception as exc: + logger.error( + f"[ConversationHistoryStore] save_round({chat_id}) failed: {exc}" + ) + return + + logger.error( + f"[ConversationHistoryStore] save_round({chat_id}) - exhausted {_APPEND_MAX_RETRIES} retries due to " + f"concurrent writes" + ) + + def _get_summary_lock(self, chat_id: str) -> asyncio.Lock: + """Get or create a lock for serializing summary updates for this chat_id. + + Args: + chat_id: The conversation identifier. + + Returns: + An asyncio.Lock that serializes summary merges for this chat_id. + """ + if chat_id not in self._summary_locks: + self._summary_locks[chat_id] = asyncio.Lock() + return self._summary_locks[chat_id] + + async def _run_incremental_summary( + self, + chat_id: str, + evicted_rounds: list[ConversationRound], + ) -> None: + """Background task: merge *evicted_rounds* into the stored summary. + + Acquires a per-chat lock to serialize summary updates for the same + chat_id. This ensures that concurrent evictions do not lose information + due to simultaneous reads and writes. Only one summarizer runs at a time + for each chat_id, making summary merges deterministic. + + Fetches the current summary, calls the injected summarizer, and persists + the result. All exceptions are caught so the task never propagates. + + Args: + chat_id: The conversation identifier. + evicted_rounds: Rounds that were just trimmed from the active window. + """ + lock = self._get_summary_lock(chat_id) + try: + async with lock: + existing_summary = await self.get_summary(chat_id) + updated = await self._summarizer(existing_summary, evicted_rounds) # type: ignore[misc] + if updated: + await self.save_summary(chat_id, updated) + logger.debug( + f"[ConversationHistoryStore] Incremental summary updated for chat_id={chat_id}" + ) + except Exception as exc: + logger.error( + f"[ConversationHistoryStore] _run_incremental_summary({chat_id}) failed: {exc}" + ) + + async def get_history(self, chat_id: str) -> list[ConversationRound]: + """Retrieve the stored rounds for a conversation. + + Returns: + Ordered list of ``ConversationRound`` objects (newest last), + or an empty list if the key is missing or Redis is unavailable. + """ + client = get_redis_client() + if client is None: + logger.warning( + f"[ConversationHistoryStore] Redis unavailable - get_history({chat_id}) skipped" + ) + return [] + + try: + raw = await client.get(_history_key(chat_id)) + if raw is None: + return [] + return [ConversationRound.model_validate(r) for r in json.loads(raw)] + except Exception as exc: + logger.error( + f"[ConversationHistoryStore] get_history({chat_id}) failed: {exc}" + ) + return [] + + async def get_summary(self, chat_id: str) -> Optional[str]: + """Retrieve the optional summary for a conversation. + + Returns: + The summary string, or None if not set or Redis is unavailable. + """ + client = get_redis_client() + if client is None: + logger.warning( + f"[ConversationHistoryStore] Redis unavailable - get_summary({chat_id}) skipped" + ) + return None + + try: + raw = await client.get(_summary_key(chat_id)) + return raw if raw is not None else None + except Exception as exc: + logger.error( + f"[ConversationHistoryStore] get_summary({chat_id}) failed: {exc}" + ) + return None + + async def save_summary(self, chat_id: str, summary: str) -> None: + """Persist a summary string and reset the TTL on both keys. + + Args: + chat_id: The conversation identifier. + summary: The condensed text to store. + """ + client = get_redis_client() + if client is None: + logger.warning( + f"[ConversationHistoryStore] Redis unavailable - save_summary({chat_id}) skipped" + ) + return + + hkey = _history_key(chat_id) + skey = _summary_key(chat_id) + + try: + async with client.pipeline(transaction=False) as pipe: + pipe.set(skey, summary, ex=_HISTORY_TTL_SECONDS) + # Reset history key TTL to keep both keys in sync + pipe.expire(hkey, _HISTORY_TTL_SECONDS) + await pipe.execute() + logger.debug( + f"[ConversationHistoryStore] Summary saved for chat_id={chat_id}" + ) + except Exception as exc: + logger.error( + f"[ConversationHistoryStore] save_summary({chat_id}) failed: {exc}" + ) + + async def get_context(self, chat_id: str) -> ConversationHistoryState: + """Return the full conversation context (rounds + summary) for a chat. + + Fetches both keys concurrently via ``asyncio.gather``. + + Returns: + A ``ConversationHistoryState`` instance. Always succeeds — both + fields fall back to safe defaults if Redis is unavailable. + """ + rounds, summary = await asyncio.gather( + self.get_history(chat_id), + self.get_summary(chat_id), + ) + return ConversationHistoryState( + chat_id=chat_id, + rounds=rounds, + summary=summary, + ) + + +def should_save_history( + conversation_history_store: Optional["ConversationHistoryStore"], + response: Union["OrchestrationResponse", "TestOrchestrationResponse"], + excluded_messages: frozenset[str], +) -> bool: + """Return True when a successful exchange should be persisted to history. + + Args: + conversation_history_store: The active store instance, or None when Redis is unavailable. + response: The response produced by the orchestration pipeline. + excluded_messages: Set of content strings that must never be persisted (OOS, error, etc.). + """ + if conversation_history_store is None: + return False + # Use duck typing to distinguish response types, avoiding isinstance() issues + # caused by import path aliasing (models.request_models vs src.models.request_models). + # OrchestrationResponse has chatId; TestOrchestrationResponse does not. + if not hasattr(response, "chatId"): + # TestOrchestrationResponse (testing env) — skip history. + return False + + # After hasattr check, safely access chatId via cast for type safety + orch_response = cast("OrchestrationResponse", response) + if orch_response.chatId is None: + return False + if response.inputGuardFailed or response.questionOutOfLLMScope: + return False + if response.content in excluded_messages: + return False + return True + + +async def save_history_round( + store: ConversationHistoryStore, + chat_id: str, + user_message: str, + bot_message: str, +) -> None: + """Persist a completed user+bot exchange to Redis. Never raises. + + Args: + store: The active ConversationHistoryStore. + chat_id: Conversation identifier. + user_message: The user's original message. + bot_message: The bot's full response. + """ + try: + round_ = ConversationRound( + user_message=user_message, + bot_message=bot_message, + ) + await store.save_round(chat_id, round_) + logger.debug( + f"[{chat_id}] Conversation history round saved ({len(bot_message)} chars)" + ) + except Exception as exc: + logger.warning(f"[{chat_id}] Failed to save conversation history round: {exc}") diff --git a/src/utils/conversation_summary_generator.py b/src/utils/conversation_summary_generator.py new file mode 100644 index 00000000..71e1c078 --- /dev/null +++ b/src/utils/conversation_summary_generator.py @@ -0,0 +1,96 @@ +"""Factory for creating incremental conversation summarizer callables.""" + +from __future__ import annotations + +import json +from typing import Any, Protocol + +import dspy +from src.loki_logger import LokiLogger + +from src.models.conversation_history_models import ConversationRound +from tool_classifier.context_analyzer import IncrementalSummarySignature + +logger = LokiLogger(service_name="context-workflow") + + +class SummarizerCallable(Protocol): + """Protocol for incremental summary callables injected into the history store.""" + + async def __call__( + self, + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + """Merge *evicted_rounds* into *existing_summary* and return the result. + + Args: + existing_summary: The current summary string, or None / empty string + when no summary exists yet. + evicted_rounds: The rounds that were just trimmed from the active + history window. + + Returns: + The updated summary string, or an empty string on failure. + """ + ... + + +def _format_rounds_as_json(rounds: list[ConversationRound]) -> str: + """Serialise *rounds* to a compact JSON string suitable for LLM prompts.""" + return json.dumps( + [r.model_dump() for r in rounds], + ensure_ascii=False, + separators=(",", ":"), + ) + + +def create_incremental_summarizer(llm_manager: Any) -> SummarizerCallable: # noqa: ANN401 + """Return an async callable that merges evicted rounds into a running summary. + + The returned callable is safe to use as a fire-and-forget background task. + Any exception is caught and logged; the caller always receives either a + non-empty updated summary or an empty string (graceful degradation). + + Args: + llm_manager: The application-wide LLM manager instance. + + Returns: + An async callable matching the ``SummarizerCallable`` protocol. + """ + _module: dspy.Module | None = None + + async def _summarize( + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + nonlocal _module + try: + rounds_json = _format_rounds_as_json(evicted_rounds) + current_summary = existing_summary or "" + + llm_manager.ensure_global_config() + with llm_manager.use_task_local(): + if _module is None: + _module = dspy.ChainOfThought(IncrementalSummarySignature) + response = _module( + existing_summary=current_summary, + new_rounds=rounds_json, + ) + + updated: str = response.updated_summary + if not updated or not updated.strip(): + logger.warning( + "[IncrementalSummarizer] LLM returned empty summary; " + "keeping existing summary unchanged." + ) + return "" + return updated.strip() + + except Exception as exc: + logger.error( + f"[IncrementalSummarizer] Failed to generate incremental summary: {exc}" + ) + return "" + + return _summarize # type: ignore[return-value] diff --git a/src/utils/cost_utils.py b/src/utils/cost_utils.py index d890c07f..e6376b91 100644 --- a/src/utils/cost_utils.py +++ b/src/utils/cost_utils.py @@ -1,10 +1,60 @@ """Cost calculation utilities for LLM usage tracking.""" -from typing import Dict, Any, List -import logging +from typing import Dict, Any, List, Tuple +from src.loki_logger import LokiLogger import dspy -logger = logging.getLogger(__name__) +# Initialize Loki logger for cost tracking +logger = LokiLogger(service_name="cost-utils") + + +def _to_float(value: str | int | float | bytes | bytearray | None) -> float: + """Best-effort float conversion for cost values.""" + try: + if value is None: + return 0.0 + return float(value) + except (TypeError, ValueError): + return 0.0 + + +def _extract_cost_with_fallback( + item: Dict[str, Any], usage: Dict[str, Any] +) -> Tuple[float, str]: + """Extract cost from history entry with streaming-safe fallbacks.""" + # Primary source used by DSPy for non-streaming calls. + item_cost = _to_float(item.get("cost")) + if item_cost > 0.0: + return item_cost, "item.cost" + + # Some providers put cost directly into usage for streaming responses. + usage_cost = _to_float(usage.get("cost")) if isinstance(usage, dict) else 0.0 + if usage_cost > 0.0: + return usage_cost, "usage.cost" + + # Final fallback: estimate from model + tokens when available. + prompt_tokens = int(usage.get("prompt_tokens", 0)) if isinstance(usage, dict) else 0 + completion_tokens = ( + int(usage.get("completion_tokens", 0)) if isinstance(usage, dict) else 0 + ) + model_name = item.get("model") + + if not model_name or (prompt_tokens == 0 and completion_tokens == 0): + return 0.0, "missing_model_or_tokens" + + try: + from litellm.cost_calculator import cost_per_token + + input_cost, output_cost = cost_per_token( + model=model_name, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + ) + estimated_cost = input_cost + output_cost + return _to_float(estimated_cost), "litellm.cost_per_token" + except Exception as e: + logger.debug(f"Cost fallback failed for model '{model_name}': {e}") + return 0.0, "fallback_error" def extract_cost_from_lm_history(lm_history: List[Dict[str, Any]]) -> Dict[str, Any]: @@ -32,13 +82,13 @@ def extract_cost_from_lm_history(lm_history: List[Dict[str, Any]]) -> Dict[str, for item in lm_history: num_calls += 1 - # Extract cost (may be None or 0 for some providers) - cost = item.get("cost", 0.0) - if cost is not None: - total_cost += float(cost) - # Extract usage information usage = item.get("usage", {}) + + # Extract cost with fallback path for streaming entries. + entry_cost, _ = _extract_cost_with_fallback(item, usage) + total_cost += entry_cost + if usage: total_prompt_tokens += usage.get("prompt_tokens", 0) total_completion_tokens += usage.get("completion_tokens", 0) @@ -113,6 +163,71 @@ def get_lm_usage_since(history_length_before: int) -> Dict[str, Any]: return usage_info +# Every guardrails self-check prompt opens with this phrase (see the +# self_check_input / self_check_output tasks in src/guardrails/rails_config.yaml). +# Matching on it lets us bill guardrail traffic separately from generation, which +# otherwise hides inside the same LM history window. +_GUARDRAIL_PROMPT_MARKER = "you are tasked with evaluating if a" + + +def _history_entry_text(item: Dict[str, Any]) -> str: + """Best-effort extraction of the prompt text from an LM history entry.""" + prompt = item.get("prompt") + if isinstance(prompt, str): + return prompt + + messages = item.get("messages") + if isinstance(messages, list): + parts = [] + for message in messages: + if isinstance(message, dict): + content = message.get("content") + if isinstance(content, str): + parts.append(content) + return "\n".join(parts) + + return "" + + +def _is_guardrail_entry(item: Dict[str, Any]) -> bool: + """Whether an LM history entry is a guardrails self-check call.""" + return _GUARDRAIL_PROMPT_MARKER in _history_entry_text(item).lower() + + +def get_lm_usage_since_split( + history_length_before: int, +) -> Tuple[Dict[str, Any], Dict[str, Any]]: + """ + Extract usage since a point, split into generation and guardrails buckets. + + During streaming, output-rail validation calls are interleaved with the + generation call in the same LM history window. Reporting them as one figure + makes runaway guardrail spend invisible, so we attribute each entry by its + prompt. + + Args: + history_length_before: The history length to measure from + + Returns: + ``(generation_usage, guardrails_usage)`` + """ + generation_usage = get_default_usage_dict() + guardrails_usage = get_default_usage_dict() + + try: + lm = dspy.settings.lm + if lm and hasattr(lm, "history"): + new_history = lm.history[history_length_before:] + guardrail_entries = [i for i in new_history if _is_guardrail_entry(i)] + generation_entries = [i for i in new_history if not _is_guardrail_entry(i)] + generation_usage = extract_cost_from_lm_history(generation_entries) + guardrails_usage = extract_cost_from_lm_history(guardrail_entries) + except Exception as e: + logger.warning(f"Failed to split usage info: {str(e)}") + + return generation_usage, guardrails_usage + + def get_default_usage_dict() -> Dict[str, Any]: """ Return a default usage dictionary with zero values. diff --git a/src/utils/decrypt_vault_secrets.py b/src/utils/decrypt_vault_secrets.py index 6a47507f..e0246230 100644 --- a/src/utils/decrypt_vault_secrets.py +++ b/src/utils/decrypt_vault_secrets.py @@ -12,15 +12,15 @@ from cryptography.hazmat.primitives import serialization, hashes from cryptography.hazmat.primitives.asymmetric import padding, rsa from cryptography.hazmat.backends import default_backend -from loguru import logger - -# Configure logger to write ONLY to stderr (stdout is reserved for decrypted value) -logger.remove() # Remove default handler -logger.add( - sys.stderr, - format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name} - {message}", - level="INFO", -) + +try: + # In cron-manager: loki_logger.py is mounted into /app/src/vector_indexer/ + from loki_logger import LokiLogger +except ModuleNotFoundError: + # In the main llm-service: import via the full src package path. + from src.loki_logger import LokiLogger + +logger = LokiLogger(service_name="vault-secrets-decryptor") def base64url_to_bytes(base64url_string: str) -> bytes: diff --git a/src/utils/input_sanitizer.py b/src/utils/input_sanitizer.py index b0bd146f..c0c1bae9 100644 --- a/src/utils/input_sanitizer.py +++ b/src/utils/input_sanitizer.py @@ -3,7 +3,10 @@ import re import html from typing import Optional, List, Dict, Any -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="input-sanitizer") class InputSanitizer: diff --git a/src/utils/language_detector.py b/src/utils/language_detector.py index db79988a..92755836 100644 --- a/src/utils/language_detector.py +++ b/src/utils/language_detector.py @@ -5,7 +5,10 @@ import re from typing import Literal -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="language-detector") LanguageCode = Literal["et", "ru", "en"] diff --git a/src/utils/observation_utils.py b/src/utils/observation_utils.py new file mode 100644 index 00000000..213ec4fd --- /dev/null +++ b/src/utils/observation_utils.py @@ -0,0 +1,120 @@ +"""Langfuse observation utilities with graceful degradation. + +Two canonical tracing patterns are used throughout this codebase. Always use +the helpers defined here rather than calling ``get_client()`` directly. + +Pattern A — Non-streaming (single LLM call wrapped with ``@observe``): + Use ``@observe(name="...", as_type="generation")`` on the method, then call + ``update_observation_safe()`` at the end to attach input/output/cost data. + + NOTE: Use ``as_type="generation"`` (not ``"chain"``) for any method that makes + exactly one LLM call. ``"chain"`` extends the async observation context longer + than a single LLM call, which can cause DSPy history entries from *this* step + to bleed into the baseline capture of the *next* step, inflating cost reports. + Reserve ``"chain"`` for orchestration methods that call multiple + ``@observe``-decorated children (Langfuse auto-aggregates their costs). + +Pattern B — Streaming (``async def`` that ``yield``s tokens): + Do NOT put ``@observe`` on async generators — the decorator interferes with + async iteration. Use ``safe_observation_context()`` as a context manager + instead, then call ``generation.update()`` once after streaming finishes. + +Pattern C — Orchestration calling multiple observed children: + Use ``@observe(name="...", as_type="chain")``. No manual update needed; + Langfuse auto-aggregates child generations. + +Pattern D — Pipeline step with no LLM call: + Use ``@observe(name="...", as_type="span")``. Optionally call + ``get_client().update_current_span(metadata={...})`` for extra context. +""" + +from contextlib import AbstractContextManager, nullcontext +from typing import Any, Dict, Optional + +from langfuse import get_client +from loguru import logger + + +def safe_observation_context(**kwargs: Any) -> AbstractContextManager[Any]: + """Return a Langfuse generation/span context manager with a no-op fallback. + + Use this for **streaming** paths (Pattern B) where ``@observe`` cannot be + used on an async generator. Wraps + ``get_client().start_as_current_observation()``; returns ``nullcontext()`` + when Langfuse is unavailable and streaming continues uninterrupted. + + The object bound via ``as`` will be ``None`` on fallback — all call sites + already guard ``.update()`` with try/except. + + Args: + **kwargs: Forwarded to ``start_as_current_observation()`` + (e.g. ``as_type``, ``name``, ``input``). + """ + try: + return get_client().start_as_current_observation(**kwargs) + except Exception as e: + logger.debug(f"Langfuse observation unavailable, using no-op context: {e}") + return nullcontext() + + +def update_observation_safe( + *, + input_data: Optional[Dict[str, Any]] = None, + output_data: Optional[Dict[str, Any]] = None, + metadata: Optional[Dict[str, Any]] = None, +) -> None: + """Attach input/output/cost data to the active Langfuse span or generation. + + Use this for **non-streaming** paths (Pattern A) inside ``@observe``-decorated + methods. Silently no-ops on any failure so tracing never blocks a response. + + Dispatches to ``update_current_generation()`` when ``metadata["usage"]`` is a + dict (i.e. an LLM call was made), otherwise to ``update_current_span()``. + The method must be decorated with ``@observe(as_type="generation")`` for + ``update_current_generation()`` to have an active target. + + Args: + input_data: Dict to set as the observation ``input``. + output_data: Dict to set as the observation ``output``. + metadata: Dict that may contain ``"model"`` (str) and ``"usage"`` (dict + with ``total_prompt_tokens``, ``total_completion_tokens``, + ``total_tokens``, ``total_cost``). All other keys are forwarded + as Langfuse ``metadata``. + """ + try: + payload: Dict[str, Any] = {} + if input_data is not None: + payload["input"] = input_data + if output_data is not None: + payload["output"] = output_data + + model_name = None + usage = None + metadata_payload: Dict[str, Any] = {} + if metadata is not None: + metadata_payload = dict(metadata) + model_name = metadata_payload.pop("model", None) + usage = metadata_payload.pop("usage", None) + + if metadata_payload: + payload["metadata"] = metadata_payload + + if isinstance(usage, dict): + get_client().update_current_generation( + model=model_name + if isinstance(model_name, str) and model_name + else None, + usage_details={ + "input": usage.get("total_prompt_tokens", 0), + "output": usage.get("total_completion_tokens", 0), + "total": usage.get("total_tokens", 0), + }, + cost_details={ + "total": usage.get("total_cost", 0.0), + }, + **payload, + ) + else: + get_client().update_current_span(**payload) + except Exception as e: + logger.debug(f"Langfuse observation update skipped: {e}") diff --git a/src/utils/production_store.py b/src/utils/production_store.py index 69026035..742803b4 100644 --- a/src/utils/production_store.py +++ b/src/utils/production_store.py @@ -8,14 +8,17 @@ from typing import Dict, List, Any, Optional from datetime import datetime import json -from loguru import logger +from src.loki_logger import LokiLogger import requests import aiohttp -from src.utils.connection_id_fetcher import get_connection_id_fetcher + from src.llm_orchestrator_config.llm_ochestrator_constants import ( RAG_SEARCH_RUUTER_PUBLIC, ) +# Initialize Loki logger +logger = LokiLogger(service_name="production-store") + class ProductionInferenceStore: """ @@ -26,7 +29,6 @@ def __init__(self) -> None: """Initialize the production inference store with Ruuter configuration.""" self.store_endpoint = f"{RAG_SEARCH_RUUTER_PUBLIC}/inference/results/store" self.timeout = 10 # seconds - self.connection_fetcher = get_connection_id_fetcher() def _create_payload( self, @@ -38,7 +40,7 @@ def _create_payload( embedding_scores: List[float], final_answer: str, environment: str, - connection_id: Optional[int], + vault_uuid: Optional[str], ) -> Dict[str, Any]: """Create the payload for storing inference results.""" return { @@ -50,7 +52,7 @@ def _create_payload( "embedding_scores": json.dumps(embedding_scores), "final_answer": final_answer, "environment": environment, - "llm_connection_id": connection_id, + "vault_uuid": vault_uuid, "created_at": datetime.now().isoformat(), } @@ -105,7 +107,7 @@ def store_inference_result( embedding_scores: List[float], final_answer: str, environment: str, - connection_id: Optional[int] = None, + vault_uuid: Optional[str] = None, ) -> Dict[str, Any]: """ Store production inference result with comprehensive data. @@ -119,7 +121,7 @@ def store_inference_result( embedding_scores: Distance scores for each chunk final_answer: LLM's final generated answer environment: Deployment environment (production/testing) - connection_id: LLM connection ID (optional, will be fetched if not provided) + vault_uuid: LLM connection vault UUID (used to look up connection_id in DB) Returns: Dict containing: @@ -128,17 +130,6 @@ def store_inference_result( - error (Optional[str]): Error message if failed """ try: - # Fetch connection ID if not provided - if connection_id is None: - logger.debug(f"Fetching {environment} connection ID...") - connection_id = self.connection_fetcher.fetch_connection_id_sync( - environment - ) - if connection_id is None: - logger.warning( - f"Could not fetch {environment} connection ID, storing without it" - ) - # Prepare the request payload payload = self._create_payload( chat_id, @@ -149,7 +140,7 @@ def store_inference_result( embedding_scores, final_answer, environment, - connection_id, + vault_uuid, ) logger.debug( @@ -218,7 +209,7 @@ async def store_inference_result_async( embedding_scores: List[float], final_answer: str, environment: str = "production", - connection_id: Optional[int] = None, + vault_uuid: Optional[str] = None, ) -> Dict[str, Any]: """ Async version of store_inference_result for streaming pipelines. @@ -232,7 +223,7 @@ async def store_inference_result_async( embedding_scores: Distance scores for each chunk final_answer: LLM's final generated answer environment: Deployment environment (production/testing) - connection_id: LLM connection ID (optional, will be fetched if not provided) + vault_uuid: LLM connection vault UUID (used to look up connection_id in DB) Returns: Dict containing: @@ -241,17 +232,6 @@ async def store_inference_result_async( - error (Optional[str]): Error message if failed """ try: - # Fetch connection ID if not provided - if connection_id is None: - logger.debug(f"Fetching {environment} connection ID (async)...") - connection_id = await self.connection_fetcher.fetch_connection_id_async( - environment - ) - if connection_id is None: - logger.warning( - f"Could not fetch {environment} connection ID, storing without it" - ) - # Prepare the request payload payload = self._create_payload( chat_id, @@ -262,7 +242,7 @@ async def store_inference_result_async( embedding_scores, final_answer, environment, - connection_id, + vault_uuid, ) logger.debug( diff --git a/src/utils/prompt_config_loader.py b/src/utils/prompt_config_loader.py index bd977657..f4b00053 100644 --- a/src/utils/prompt_config_loader.py +++ b/src/utils/prompt_config_loader.py @@ -7,7 +7,10 @@ import time import threading from enum import Enum -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="prompt-config-loader") class PromptConfigLoadError(Exception): diff --git a/src/utils/rate_limiter.py b/src/utils/rate_limiter.py index 074a0f6f..b875715d 100644 --- a/src/utils/rate_limiter.py +++ b/src/utils/rate_limiter.py @@ -4,12 +4,14 @@ from collections import defaultdict, deque from typing import Dict, Deque, Optional, Any from threading import Lock - -from loguru import logger from pydantic import BaseModel, Field, ConfigDict +from src.loki_logger import LokiLogger from src.llm_orchestrator_config.stream_config import StreamConfig +# Initialize Loki logger +logger = LokiLogger(service_name="rate-limiter") + class RateLimitResult(BaseModel): """Result of rate limit check.""" diff --git a/src/utils/redis_client.py b/src/utils/redis_client.py index 960a9752..ffeca9f8 100644 --- a/src/utils/redis_client.py +++ b/src/utils/redis_client.py @@ -4,7 +4,9 @@ from typing import Any, Optional import redis.asyncio as aioredis -from loguru import logger +from src.loki_logger import LokiLogger + +logger = LokiLogger(service_name="redis_client") _redis_client: Optional[aioredis.Redis] = None # type: ignore[type-arg] @@ -69,7 +71,7 @@ async def init_redis_client() -> aioredis.Redis: # Verify connectivity await _redis_client.ping() logger.info( - "Redis session store connected (db={})", os.getenv("REDIS_SESSION_DB", "1") + f"Redis session store connected (db={os.getenv('REDIS_SESSION_DB', '1')})" ) return _redis_client diff --git a/src/utils/sse_utils.py b/src/utils/sse_utils.py new file mode 100644 index 00000000..e0c0810e --- /dev/null +++ b/src/utils/sse_utils.py @@ -0,0 +1,16 @@ +"""Utilities for parsing Server-Sent Events (SSE) messages.""" + +import json +from typing import Optional + + +def extract_content_from_sse(sse_chunk: str) -> Optional[str]: + """Parse an SSE chunk and return payload.content, or None on failure.""" + if not sse_chunk.startswith("data: "): + return None + json_part = sse_chunk[len("data: ") :].strip() + try: + parsed = json.loads(json_part) + return parsed.get("payload", {}).get("content") + except (json.JSONDecodeError, AttributeError): + return None diff --git a/src/utils/stream_manager.py b/src/utils/stream_manager.py index e12296ea..8486797c 100644 --- a/src/utils/stream_manager.py +++ b/src/utils/stream_manager.py @@ -4,13 +4,16 @@ from datetime import datetime from contextlib import asynccontextmanager import asyncio -from loguru import logger +from src.loki_logger import LokiLogger from pydantic import BaseModel, Field, ConfigDict from src.llm_orchestrator_config.stream_config import StreamConfig from src.llm_orchestrator_config.exceptions import StreamError from src.utils.error_utils import generate_error_id +# Initialize Loki logger +logger = LokiLogger(service_name="stream-manager") + class StreamContext(BaseModel): """Context for tracking a single stream's lifecycle.""" diff --git a/src/utils/stream_timeout.py b/src/utils/stream_timeout.py index 3278b7bb..e382272f 100644 --- a/src/utils/stream_timeout.py +++ b/src/utils/stream_timeout.py @@ -2,10 +2,40 @@ import asyncio from contextlib import asynccontextmanager -from typing import AsyncIterator +from typing import AsyncIterator, Optional, Union from src.llm_orchestrator_config.exceptions import StreamTimeoutError +# An SSE comment frame. Ignored by EventSource and by the notification server's +# relay (which only forwards `data: ` lines), but it is bytes on the wire, which +# is what keeps proxy idle timers from closing a slow stream. +HEARTBEAT_FRAME = ": ping\n\n" + + +class _StreamExhausted: + """Sentinel type marking a normally-completed source iterator. + + A dedicated class rather than a bare ``object()`` so that an ``isinstance`` + check narrows the value back to ``str`` for type checkers. + """ + + +_STREAM_EXHAUSTED = _StreamExhausted() + + +async def _next_or_sentinel( + iterator: AsyncIterator[str], +) -> Union[str, _StreamExhausted]: + """Advance an async iterator, returning a sentinel instead of raising at the end. + + Returning a sentinel keeps StopAsyncIteration out of the asyncio.Task that + wraps this call, where it would be an awkward special case. + """ + try: + return await iterator.__anext__() + except StopAsyncIteration: + return _STREAM_EXHAUSTED + @asynccontextmanager async def stream_timeout(seconds: int) -> AsyncIterator[None]: @@ -30,3 +60,66 @@ async def stream_timeout(seconds: int) -> AsyncIterator[None]: raise StreamTimeoutError( f"Stream exceeded maximum duration of {seconds} seconds" ) from e + + +async def with_heartbeat( + source: AsyncIterator[str], + heartbeat_interval: float, + idle_timeout: float, +) -> AsyncIterator[str]: + """ + Relay a stream, emitting SSE comment frames during quiet periods. + + Two problems are solved together. A long pause between chunks lets any proxy + on the path close the connection, so we keep writing; and a stream that has + genuinely stalled should fail fast rather than sit until the total-duration + cap expires, so we enforce an idle budget. + + Both timers measure the gap *between* chunks - a long answer that keeps + producing is never interrupted, however long it runs in total. + + Args: + source: The upstream chunk iterator. + heartbeat_interval: Seconds of quiet before emitting a heartbeat frame. + idle_timeout: Seconds of continuous quiet before giving up. + + Yields: + Chunks from ``source``, interleaved with ``HEARTBEAT_FRAME``. + + Raises: + StreamTimeoutError: If no chunk arrives for ``idle_timeout`` seconds. + """ + iterator = source.__aiter__() + pending: Optional["asyncio.Task[Union[str, _StreamExhausted]]"] = None + + try: + while True: + pending = asyncio.ensure_future(_next_or_sentinel(iterator)) + idle_elapsed = 0.0 + + while True: + try: + # Shielded so a heartbeat timeout does not cancel the pull; + # the same task is awaited again on the next pass. + chunk = await asyncio.wait_for( + asyncio.shield(pending), heartbeat_interval + ) + break + except asyncio.TimeoutError: + idle_elapsed += heartbeat_interval + if idle_elapsed >= idle_timeout: + pending.cancel() + raise StreamTimeoutError( + f"Stream produced no output for {idle_elapsed:.1f} " + f"seconds (idle limit {idle_timeout:.1f}s)" + ) from None + yield HEARTBEAT_FRAME + + if isinstance(chunk, _StreamExhausted): + return + + yield chunk + finally: + # Covers early consumer exit (client disconnect) as well as errors. + if pending is not None and not pending.done(): + pending.cancel() diff --git a/src/utils/time_tracker.py b/src/utils/time_tracker.py index 619030d7..5f9a4816 100644 --- a/src/utils/time_tracker.py +++ b/src/utils/time_tracker.py @@ -1,7 +1,10 @@ """Simple time tracking for orchestration service steps.""" from typing import Dict, Optional -from loguru import logger +from src.loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="time-tracker") def log_step_timings( diff --git a/src/vector_indexer/api_client.py b/src/vector_indexer/api_client.py index fa4f2d29..0f43e824 100644 --- a/src/vector_indexer/api_client.py +++ b/src/vector_indexer/api_client.py @@ -3,11 +3,14 @@ import asyncio from typing import List, Dict, Any, Optional, Union import httpx -from loguru import logger from typing_extensions import Self +from loki_logger import LokiLogger from vector_indexer.config.config_loader import VectorIndexerConfig +# Initialize Loki logger +logger = LokiLogger(service_name="api-client") + class LLMOrchestrationAPIClient: """Client for calling LLM Orchestration Service API endpoints.""" diff --git a/src/vector_indexer/config/config_loader.py b/src/vector_indexer/config/config_loader.py index 24af5d76..be107577 100644 --- a/src/vector_indexer/config/config_loader.py +++ b/src/vector_indexer/config/config_loader.py @@ -4,7 +4,7 @@ from pathlib import Path from typing import Optional, List, Dict, Any from pydantic import BaseModel, Field, field_validator, model_validator -from loguru import logger +from loki_logger import LokiLogger from vector_indexer.constants import ( DocumentConstants, @@ -13,6 +13,9 @@ ProcessingConstants, ) +# Initialize Loki logger +logger = LokiLogger(service_name="config-loader") + class ChunkingConfig(BaseModel): """Configuration for document chunking operations""" diff --git a/src/vector_indexer/contextual_processor.py b/src/vector_indexer/contextual_processor.py index b225cf30..9b69641c 100644 --- a/src/vector_indexer/contextual_processor.py +++ b/src/vector_indexer/contextual_processor.py @@ -3,7 +3,7 @@ import asyncio import tiktoken from typing import List, Dict, Any, Optional -from loguru import logger +from loki_logger import LokiLogger from vector_indexer.config.config_loader import VectorIndexerConfig from vector_indexer.models import ProcessingDocument, BaseChunk, ContextualChunk @@ -11,6 +11,9 @@ from vector_indexer.error_logger import ErrorLogger from vector_indexer.constants import ChunkingConstants, ProcessingConstants +# Initialize Loki logger +logger = LokiLogger(service_name="contextual-processor") + class ContextualProcessor: """Processes documents into contextual chunks using Anthropic methodology.""" @@ -41,7 +44,7 @@ def __init__( async def process_document( self, document: ProcessingDocument - ) -> List[ContextualChunk]: + ) -> tuple[List[ContextualChunk], int]: """ Process single document into contextual chunks. @@ -49,7 +52,8 @@ async def process_document( document: Document to process Returns: - List of contextual chunks with embeddings + Tuple of (contextual chunks with embeddings, number of chunks + dropped due to context-generation failure) """ logger.info( f"Processing document {document.document_hash} ({len(document.content)} characters)" @@ -69,11 +73,13 @@ async def process_document( # Step 3: Create contextual chunks (filter out failed context generations) contextual_chunks: List[ContextualChunk] = [] valid_contextual_contents: List[str] = [] + failed_chunks = 0 for i, (base_chunk, context) in enumerate( zip(base_chunks, contexts, strict=True) ): if isinstance(context, Exception): + failed_chunks += 1 self.error_logger.log_context_generation_failure( document.document_hash, i, str(context), self.config.max_retries ) @@ -128,7 +134,7 @@ async def process_document( logger.error( f"No valid chunks created for document {document.document_hash}" ) - return [] + return [], failed_chunks # Step 4: Create embeddings for all valid contextual chunks try: @@ -154,9 +160,10 @@ async def process_document( raise logger.info( - f"Successfully processed document {document.document_hash}: {len(contextual_chunks)} chunks" + f"Successfully processed document {document.document_hash}: " + f"{len(contextual_chunks)} chunks ({failed_chunks} dropped)" ) - return contextual_chunks + return contextual_chunks, failed_chunks except Exception as e: logger.error( diff --git a/src/vector_indexer/dataset_download.py b/src/vector_indexer/dataset_download.py index ebd95901..c3a3741f 100644 --- a/src/vector_indexer/dataset_download.py +++ b/src/vector_indexer/dataset_download.py @@ -4,7 +4,10 @@ import tempfile from pathlib import Path import requests -from loguru import logger +from loki_logger import LokiLogger + +# Initialize Loki logger +logger = LokiLogger(service_name="dataset-download") def download_and_extract_dataset(signed_url: str) -> tuple[str, int]: diff --git a/src/vector_indexer/diff_identifier/DIFF_IDENTIFIER_FLOW.md b/src/vector_indexer/diff_identifier/DIFF_IDENTIFIER_FLOW.md index 57a48d20..34b23134 100644 --- a/src/vector_indexer/diff_identifier/DIFF_IDENTIFIER_FLOW.md +++ b/src/vector_indexer/diff_identifier/DIFF_IDENTIFIER_FLOW.md @@ -94,7 +94,7 @@ class VersionManager: "last_run_modified_files": 1, "last_run_deleted_files": 1, "last_cleanup_deleted_chunks": 15, - "last_run_timestamp": "2025-10-17T00:00:46Z" + "last_run_timestamp": "2025-10-17T00:00:46Z", }, "processed_files": { "sha256_hash": { @@ -103,9 +103,15 @@ class VersionManager: "file_size": 15234, "processed_at": "2025-10-17T00:00:46Z", "chunk_count": 5, # Track chunk count for validation - "chunk_ids": ["uuid1", "uuid2", "uuid3", "uuid4", "uuid5"] # Track exact chunks + "chunk_ids": [ + "uuid1", + "uuid2", + "uuid3", + "uuid4", + "uuid5", + ], # Track exact chunks } - } + }, } ``` @@ -138,27 +144,39 @@ class ProcessedFileInfo(BaseModel): chunk_count: int = 0 # NEW: Track number of chunks chunk_ids: List[str] = Field(default_factory=list) # NEW: Track chunk IDs + class DiffResult(BaseModel): # File change detection new_files: List[str] = Field(..., description="Files to process for first time") - modified_files: List[str] = Field(default_factory=list, description="Files with changed content") - deleted_files: List[str] = Field(default_factory=list, description="Files removed from dataset") - unchanged_files: List[str] = Field(default_factory=list, description="Files with same content") - + modified_files: List[str] = Field( + default_factory=list, description="Files with changed content" + ) + deleted_files: List[str] = Field( + default_factory=list, description="Files removed from dataset" + ) + unchanged_files: List[str] = Field( + default_factory=list, description="Files with same content" + ) + # Statistics total_files_scanned: int previously_processed_count: int is_first_run: bool - + # NEW: Cleanup metadata - chunks_to_delete: Dict[str, List[str]] = Field(default_factory=dict) # document_hash -> chunk_ids + chunks_to_delete: Dict[str, List[str]] = Field( + default_factory=dict + ) # document_hash -> chunk_ids estimated_cleanup_count: int = Field(default=0) # Total chunks to be removed + class VersionState(BaseModel): last_updated: str processed_files: Dict[str, ProcessedFileInfo] total_processed: int - processing_stats: Dict[str, Any] = Field(default_factory=dict) # NEW: Enhanced stats + processing_stats: Dict[str, Any] = Field( + default_factory=dict + ) # NEW: Enhanced stats ``` ## Enhanced Processing Flow @@ -247,8 +265,8 @@ if not files_to_process: ```python # NEW: Track chunk information in metadata await diff_detector.mark_files_processed( - processed_paths, - chunks_info=collected_chunk_information # Future enhancement + processed_paths, + chunks_info=collected_chunk_information, # Future enhancement ) ``` @@ -269,7 +287,7 @@ await diff_detector.mark_files_processed( # Efficient chunk identification for cleanup chunks_to_delete = { "document_hash_123": ["chunk_uuid_1", "chunk_uuid_2", "chunk_uuid_3"], - "document_hash_456": ["chunk_uuid_4", "chunk_uuid_5"] + "document_hash_456": ["chunk_uuid_4", "chunk_uuid_5"], } # Cleanup execution per collection @@ -395,7 +413,9 @@ diff_result = await diff_detector.get_changed_files() logger.info(f"Cleanup metadata: {diff_result.chunks_to_delete}") # Test cleanup operations -cleanup_count = await main_indexer._execute_cleanup_operations(qdrant_manager, diff_result) +cleanup_count = await main_indexer._execute_cleanup_operations( + qdrant_manager, diff_result +) logger.info(f"Total cleanup: {cleanup_count} chunks") ``` @@ -408,20 +428,22 @@ logger.info(f"Total cleanup: {cleanup_count} chunks") async def process_all_documents(self) -> ProcessingStats: # 1. Enhanced diff detection diff_result = await diff_detector.get_changed_files() - + # 2. NEW: Automatic cleanup execution if diff_result.chunks_to_delete: - cleanup_count = await self._execute_cleanup_operations(qdrant_manager, diff_result) - + cleanup_count = await self._execute_cleanup_operations( + qdrant_manager, diff_result + ) + # 3. Selective document processing files_to_process = diff_result.new_files + diff_result.modified_files if not files_to_process: return self.stats # Early exit - + # 4. Standard processing pipeline documents = self._filter_documents_by_paths(files_to_process) results = await self._process_documents(documents) - + # 5. Enhanced metadata update await diff_detector.mark_files_processed(processed_paths, chunks_info) ``` @@ -565,17 +587,18 @@ if is_first_run: current_files = version_manager.scan_current_files() # Returns: Dict[content_hash, file_path] for all discovered files + def scan_current_files(self) -> Dict[str, str]: file_hash_map = {} for root, _, files in os.walk(self.config.datasets_path): for file in files: file_path = os.path.join(root, file) relative_path = os.path.relpath(file_path, self.config.datasets_path) - + # Calculate content hash for change detection content_hash = self._calculate_file_hash(file_path) file_hash_map[content_hash] = relative_path - + return file_hash_map ``` @@ -592,23 +615,24 @@ def scan_current_files(self) -> Dict[str, str]: processed_metadata = await s3_ferry_client.download_metadata() # Downloads from: s3://rag-search/resources/datasets/processed-metadata.json + def download_metadata(self) -> Optional[Dict[str, Any]]: # Create temporary file for S3Ferry transfer - with tempfile.NamedTemporaryFile(suffix='.json', delete=False) as temp_file: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as temp_file: temp_file_path = temp_file.name - + # Transfer S3 → FS via S3Ferry API response = self._retry_with_backoff( lambda: self.s3_ferry.transfer_file( destinationFilePath=temp_file_path, - destinationStorageType="FS", + destinationStorageType="FS", sourceFilePath=self.config.metadata_s3_path, - sourceStorageType="S3" + sourceStorageType="S3", ) ) - + if response.status_code == 200: - with open(temp_file_path, 'r') as f: + with open(temp_file_path, "r") as f: return json.load(f) elif response.status_code == 404: return None # First run - no metadata exists yet @@ -624,19 +648,23 @@ def download_metadata(self) -> Optional[Dict[str, Any]]: #### Phase 5: Differential Analysis ```python # 6. Change Detection Algorithm (version_manager.py) -changed_files = version_manager.identify_changed_files(current_files, processed_metadata) +changed_files = version_manager.identify_changed_files( + current_files, processed_metadata +) -def identify_changed_files(self, current_files: Dict[str, str], - processed_state: Optional[Dict]) -> Set[str]: + +def identify_changed_files( + self, current_files: Dict[str, str], processed_state: Optional[Dict] +) -> Set[str]: if not processed_state: return set(current_files.values()) # All files are "new" - - processed_hashes = set(processed_state.get('processed_files', {}).keys()) + + processed_hashes = set(processed_state.get("processed_files", {}).keys()) current_hashes = set(current_files.keys()) - + # Identify new and modified files new_or_changed_hashes = current_hashes - processed_hashes - + # Convert hashes back to file paths return {current_files[hash_val] for hash_val in new_or_changed_hashes} ``` @@ -654,8 +682,8 @@ def identify_changed_files(self, current_files: Dict[str, str], return DiffResult( new_files=list(changed_files), total_files_scanned=len(current_files), - previously_processed_count=len(processed_state.get('processed_files', {})), - is_first_run=is_first_run + previously_processed_count=len(processed_state.get("processed_files", {})), + is_first_run=is_first_run, ) ``` @@ -706,11 +734,15 @@ Dataset Download → [shared-volume] → diff_identifier → [datasets mount] if diff_result.new_files: # Process only changed files documents = self._filter_documents_by_paths(diff_result.new_files) - logger.info(f"Processing {len(documents)} documents from {len(diff_result.new_files)} changed files") + logger.info( + f"Processing {len(documents)} documents from {len(diff_result.new_files)} changed files" + ) else: # No changes detected - skip processing entirely logger.info("No changes detected. Skipping processing phase.") - return ProcessingResult(processed_count=0, skipped_count=diff_result.total_files_scanned) + return ProcessingResult( + processed_count=0, skipped_count=diff_result.total_files_scanned + ) # Continue with existing vector generation pipeline... ``` @@ -727,25 +759,26 @@ else: async def mark_files_processed(self, file_paths: List[str]) -> bool: # Update processed files metadata new_metadata = self._create_updated_metadata(file_paths) - + # Upload to S3 via S3Ferry success = await self.s3_ferry_client.upload_metadata(new_metadata) - + # Commit DVC state (optional - for advanced versioning) if success: self.version_manager.commit_dvc_state(f"Processed {len(file_paths)} files") - + return success + def _create_updated_metadata(self, file_paths: List[str]) -> Dict[str, Any]: current_files = self.version_manager.scan_current_files() - + metadata = { "last_updated": datetime.utcnow().isoformat(), - "total_processed": len(file_paths), - "processed_files": {} + "total_processed": len(file_paths), + "processed_files": {}, } - + # Add file metadata for each processed file for file_path in file_paths: file_hash = self._get_file_hash(file_path) @@ -753,9 +786,9 @@ def _create_updated_metadata(self, file_paths: List[str]) -> Dict[str, Any]: content_hash=file_hash, original_path=file_path, file_size=os.path.getsize(file_path), - processed_at=datetime.utcnow().isoformat() + processed_at=datetime.utcnow().isoformat(), ).dict() - + return metadata ``` @@ -1067,14 +1100,14 @@ diff_detector = DiffDetector(diff_config) # Passes to main orchestrator # diff_detector.py - Configuration factory config = DiffConfig( - s3_ferry_url=s3_ferry_url, # → Used by S3FerryClient - metadata_s3_path=metadata_s3_path, # → Used for S3Ferry operations - datasets_path=datasets_path, # → Used for file scanning - metadata_filename=metadata_filename, # → Used to build paths - dvc_remote_url=dvc_remote_url, # → Used by DVC setup - s3_endpoint_url=str(s3_endpoint_url), # → Used by DVC S3 config - s3_access_key_id=str(s3_access_key_id), # → Used by DVC authentication - s3_secret_access_key=str(s3_secret_access_key) # → Used by DVC authentication + s3_ferry_url=s3_ferry_url, # → Used by S3FerryClient + metadata_s3_path=metadata_s3_path, # → Used for S3Ferry operations + datasets_path=datasets_path, # → Used for file scanning + metadata_filename=metadata_filename, # → Used to build paths + dvc_remote_url=dvc_remote_url, # → Used by DVC setup + s3_endpoint_url=str(s3_endpoint_url), # → Used by DVC S3 config + s3_access_key_id=str(s3_access_key_id), # → Used by DVC authentication + s3_secret_access_key=str(s3_secret_access_key), # → Used by DVC authentication ) ``` @@ -1118,7 +1151,7 @@ response = self.s3_ferry.transfer_file( destinationFilePath="resources/datasets/processed-metadata.json", destinationStorageType="S3", sourceFilePath="/tmp/tmpABC123.json", # Temporary file - sourceStorageType="FS" + sourceStorageType="FS", ) ``` @@ -1145,7 +1178,7 @@ response = self.s3_ferry.transfer_file( destinationFilePath="/tmp/tmpDEF456.json", # Temporary file destinationStorageType="FS", sourceFilePath="resources/datasets/processed-metadata.json", - sourceStorageType="S3" + sourceStorageType="S3", ) ``` @@ -1535,7 +1568,7 @@ DiffResult( new_files=["datasets/collection1/abc123/cleaned.txt"], total_files_scanned=100, previously_processed_count=99, - is_first_run=False + is_first_run=False, ) ``` diff --git a/src/vector_indexer/diff_identifier/diff_detector.py b/src/vector_indexer/diff_identifier/diff_detector.py index 46edd3dd..51f77b72 100644 --- a/src/vector_indexer/diff_identifier/diff_detector.py +++ b/src/vector_indexer/diff_identifier/diff_detector.py @@ -3,7 +3,7 @@ import os from pathlib import Path from typing import List, Optional, Dict, Any -from loguru import logger +from loki_logger import LokiLogger import hashlib from diff_identifier.diff_models import DiffConfig, DiffError, DiffResult @@ -12,6 +12,9 @@ load_dotenv(".env") +# Initialize Loki logger +logger = LokiLogger(service_name="diff-detector") + class DiffDetector: """Main orchestrator for diff identification.""" @@ -144,13 +147,14 @@ async def mark_files_processed( logger.info(f"Marking {len(processed_file_paths)} files as processed...") - # Log chunks_info received + # Log chunks_info summary only (avoid massive logs) if chunks_info: - logger.info(f"RECEIVED CHUNKS INFO: {len(chunks_info)} documents") - for doc_hash, info in chunks_info.items(): - logger.info( - f" {doc_hash[:12]}... -> {info.get('chunk_count', 0)} chunks" - ) + total_chunks = sum( + info.get("chunk_count", 0) for info in chunks_info.values() + ) + logger.info( + f"RECEIVED CHUNKS INFO: {len(chunks_info)} documents, {total_chunks} total chunks" + ) else: logger.warning("No chunks_info provided to mark_files_processed") diff --git a/src/vector_indexer/diff_identifier/s3_ferry_client.py b/src/vector_indexer/diff_identifier/s3_ferry_client.py index bebb7464..f423bdd4 100644 --- a/src/vector_indexer/diff_identifier/s3_ferry_client.py +++ b/src/vector_indexer/diff_identifier/s3_ferry_client.py @@ -5,12 +5,16 @@ import time from typing import Any, Callable, Dict, Optional import requests -from loguru import logger +from loki_logger import LokiLogger + from typing_extensions import Self from diff_identifier.diff_models import DiffConfig, DiffError from constants import get_s3_ferry_payload +# Initialize Loki logger +logger = LokiLogger(service_name="s3-ferry-client") + class S3Ferry: """Client for interacting with S3Ferry service.""" diff --git a/src/vector_indexer/diff_identifier/version_manager.py b/src/vector_indexer/diff_identifier/version_manager.py index d7df5a83..a0592093 100644 --- a/src/vector_indexer/diff_identifier/version_manager.py +++ b/src/vector_indexer/diff_identifier/version_manager.py @@ -5,7 +5,8 @@ from datetime import datetime from pathlib import Path from typing import Dict, List, Optional, Set, Any -from loguru import logger +from loki_logger import LokiLogger + from typing_extensions import Self from diff_identifier.diff_models import ( @@ -16,6 +17,9 @@ ) from diff_identifier.s3_ferry_client import S3FerryClient +# Initialize Loki logger +logger = LokiLogger(service_name="version-manager") + class VersionManager: """Manages DVC operations and version tracking.""" diff --git a/src/vector_indexer/document_loader.py b/src/vector_indexer/document_loader.py index 9e03b290..6c572834 100644 --- a/src/vector_indexer/document_loader.py +++ b/src/vector_indexer/document_loader.py @@ -6,12 +6,15 @@ from typing import List from urllib.parse import urlparse -from loguru import logger +from loki_logger import LokiLogger from vector_indexer.config.config_loader import VectorIndexerConfig from vector_indexer.models import DocumentInfo, ProcessingDocument from vector_indexer.constants import DocumentConstants +# Initialize Loki logger +logger = LokiLogger(service_name="document-loader") + class DocumentLoadError(Exception): """Custom exception for document loading failures.""" diff --git a/src/vector_indexer/error_logger.py b/src/vector_indexer/error_logger.py index 1d11cba1..4a1db2aa 100644 --- a/src/vector_indexer/error_logger.py +++ b/src/vector_indexer/error_logger.py @@ -1,13 +1,15 @@ """Enhanced error logging for vector indexer.""" import json -import sys from pathlib import Path -from loguru import logger +from loki_logger import LokiLogger from vector_indexer.config.config_loader import VectorIndexerConfig from vector_indexer.models import ProcessingError, ProcessingStats +# Initialize Loki logger +logger = LokiLogger(service_name="error-logger") + class ErrorLogger: """Enhanced error logging with file-based failure tracking.""" @@ -27,24 +29,10 @@ def _ensure_log_directories(self) -> None: Path(log_file).parent.mkdir(parents=True, exist_ok=True) def _setup_logging(self) -> None: - """Setup loguru logging with file output.""" - logger.remove() # Remove default handler - - # Console logging - logger.add( - sys.stdout, - level=self.config.log_level, - format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}", - ) - - # File logging - logger.add( - self.config.processing_log_file, - level=self.config.log_level, - format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}", - rotation="10 MB", - retention="7 days", - ) + """Setup logging - LokiLogger handles all logging (console + Loki service).""" + # LokiLogger is already configured and handles console output + Loki integration + # No additional configuration needed + pass def log_document_failure( self, document_hash: str, error: str, retry_count: int = 0 @@ -158,15 +146,17 @@ def log_processing_stats(self, stats: ProcessingStats) -> None: stats_dict["end_time"] = stats.end_time.isoformat() stats_dict["duration"] = stats.duration stats_dict["success_rate"] = stats.success_rate + stats_dict["chunk_success_rate"] = stats.chunk_success_rate with open(self.config.stats_log_file, "w", encoding="utf-8") as f: json.dump(stats_dict, f, indent=2) logger.info( f"Processing completed - Success rate: {stats.success_rate:.1%}, " + f"Chunk success rate: {stats.chunk_success_rate:.1%}, " f"Duration: {stats.duration}, " f"Processed: {stats.documents_processed}/{stats.total_documents} documents, " - f"Chunks: {stats.total_chunks_processed}" + f"Chunks: {stats.total_chunks_processed} ok / {stats.total_chunks_failed} failed" ) except Exception as e: logger.error(f"Failed to write stats log: {e}") diff --git a/src/vector_indexer/loki_logger.py b/src/vector_indexer/loki_logger.py index e69de29b..a86a969c 100644 --- a/src/vector_indexer/loki_logger.py +++ b/src/vector_indexer/loki_logger.py @@ -0,0 +1,176 @@ +#!/usr/bin/env python3 +""" +Loki Logger for RAG Module +Sends logs directly to Loki API for centralized logging +""" + +import json +import sys +import time +from datetime import datetime +from threading import Thread +from queue import Full, Queue + +import requests + + +class LokiLogger: + """Simple logger that sends logs directly to Loki API with async background thread""" + + _instances: dict[str, "LokiLogger"] = {} + + def __new__( + cls, loki_url: str = "http://loki:3100", service_name: str = "default" + ) -> "LokiLogger": + key = f"{loki_url}:{service_name}" + if key not in cls._instances: + cls._instances[key] = super().__new__(cls) + return cls._instances[key] + + def __init__( + self, loki_url: str = "http://loki:3100", service_name: str = "default" + ) -> None: + """ + Initialize LokiLogger + + Args: + loki_url: URL for Loki service (default: container URL in bykstack network) + service_name: Name of the service for labeling logs + """ + if hasattr(self, "_initialized"): + return + self._initialized = True + self.loki_url = loki_url + self.service_name = service_name + self.session = requests.Session() + # Set default timeout for all requests + self.timeout = 5 + + # Queue for async log processing (bounded to avoid unbounded memory growth under load) + self.log_queue: Queue[tuple[str, str]] = Queue(maxsize=10_000) + + # Start background worker thread + self.worker_thread = Thread(target=self._process_logs, daemon=True) + self.worker_thread.start() + + def _process_logs(self) -> None: + """Background worker that processes log queue""" + while True: + try: + # Get log entry from queue (blocking) + level, message = self.log_queue.get() + + # Send to Loki + self._send_to_loki_sync(level, message) + + # Mark task as done + self.log_queue.task_done() + except Exception: + # Silently ignore errors in background thread + pass + + def _send_to_loki_sync(self, level: str, message: str) -> None: + """Send log entry directly to Loki API (called from background thread)""" + try: + # Create timestamp in nanoseconds (Loki requirement) + timestamp_ns = str(int(time.time() * 1_000_000_000)) + + # Prepare labels for Loki + labels = { + "service": self.service_name, + "level": level, + } + + # Create log entry + log_entry = { + "level": level, + "message": message, + "service": self.service_name, + } + + # Prepare Loki payload + payload = { + "streams": [ + { + "stream": labels, + "values": [[timestamp_ns, json.dumps(log_entry)]], + } + ] + } + + # Send to Loki + self.session.post( + f"{self.loki_url}/loki/api/v1/push", + json=payload, + headers={"Content-Type": "application/json"}, + timeout=self.timeout, + ) + + except Exception: + # Silently ignore logging errors to not affect main application + pass + + def _log(self, level: str, message: str) -> None: + """Queue log entry for async processing (non-blocking)""" + # Print to console immediately for real-time feedback. Written to + # stderr (not stdout) so callers that capture a subprocess's stdout + # for its return value (e.g. decrypt_vault_secrets.py) never pick up + # log lines mixed in with the actual output. + timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + print(f"[{timestamp}] {level: <8} | {message}", file=sys.stderr) # noqa: T201 + + # Queue for async Loki sending (non-blocking) + try: + self.log_queue.put_nowait((level, message)) + except Full: + # Queue full (Loki may be slow/unreachable) - drop log to avoid blocking + pass + + def info(self, message: str, **kwargs: object) -> None: + """Log info message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("INFO", message) + + def error(self, message: str, **kwargs: object) -> None: + """Log error message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("ERROR", message) + + def warning(self, message: str, **kwargs: object) -> None: + """Log warning message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("WARNING", message) + + def debug(self, message: str, **kwargs: object) -> None: + """Log debug message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("DEBUG", message) + + def success(self, message: str, **kwargs: object) -> None: + """Log success message (loguru compatibility). Extra kwargs ignored.""" + self._log("SUCCESS", message) + + def critical(self, message: str, **kwargs: object) -> None: + """Log critical message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("CRITICAL", message) + + def exception(self, message: str, **kwargs: object) -> None: + """Log exception message. Extra kwargs (extra, exc_info) are ignored for compatibility.""" + self._log("EXCEPTION", message) + + def add(self, *args: object, **kwargs: object) -> None: + """ + No-op method for loguru compatibility. + + LokiLogger sends logs to Loki/console only, not to files. + This method exists for backward compatibility with loguru code. + """ + pass # Silently ignore - logs go to Loki instead of files + + def remove(self, *args: object, **kwargs: object) -> None: + """No-op method for loguru compatibility.""" + pass # Silently ignore + + def bind(self, **kwargs: object) -> "LokiLogger": + """No-op method for loguru compatibility. Returns self for chaining.""" + return self # Allow method chaining + + def opt(self, **kwargs: object) -> "LokiLogger": + """No-op method for loguru compatibility. Returns self for chaining.""" + return self # Allow method chaining diff --git a/src/vector_indexer/main_indexer.py b/src/vector_indexer/main_indexer.py index 45ce5ff6..63596e90 100644 --- a/src/vector_indexer/main_indexer.py +++ b/src/vector_indexer/main_indexer.py @@ -7,22 +7,24 @@ from pathlib import Path from datetime import datetime from typing import List, Optional, Dict, Any -from loguru import logger +from loki_logger import LokiLogger import hashlib - # Add src to path for imports sys.path.append(str(Path(__file__).parent.parent)) from vector_indexer.config.config_loader import ConfigLoader -from vector_indexer.document_loader import DocumentLoader +from vector_indexer.document_loader import DocumentLoader, DocumentLoadError from vector_indexer.contextual_processor import ContextualProcessor from vector_indexer.qdrant_manager import QdrantManager from vector_indexer.error_logger import ErrorLogger from vector_indexer.models import ProcessingStats, DocumentInfo from vector_indexer.diff_identifier import DiffDetector, create_diff_config, DiffError from vector_indexer.diff_identifier.diff_models import DiffResult -from src.vector_indexer.dataset_download import download_and_extract_dataset +from vector_indexer.dataset_download import download_and_extract_dataset + +# Initialize Loki logger +logger = LokiLogger(service_name="main-indexer") class VectorIndexer: @@ -169,7 +171,7 @@ async def process_all_documents(self) -> ProcessingStats: # Process documents with controlled concurrency semaphore = asyncio.Semaphore(self.config.max_concurrent_documents) - tasks: List[asyncio.Task[tuple[int, str]]] = [] + tasks: List[asyncio.Task[tuple[int, str, int]]] = [] for doc_info in documents: task = asyncio.create_task( @@ -189,6 +191,9 @@ async def process_all_documents(self) -> ProcessingStats: chunks_info: Dict[ str, Dict[str, Any] ] = {} # Track chunk counts for metadata update + # Only documents that processed successfully are marked as + # processed in DVC tracking, so failures are retried next run. + processed_documents: List[DocumentInfo] = [] for i, result in enumerate(results): if isinstance(result, Exception): doc_info = documents[i] @@ -200,26 +205,25 @@ async def process_all_documents(self) -> ProcessingStats: doc_info.document_hash, str(result) ) else: - # Result should be tuple of (chunk_count, content_hash) + # Result should be tuple of (chunk_count, content_hash, failed_chunks) doc_info = documents[i] self.stats.documents_processed += 1 - if isinstance(result, tuple) and len(result) == 2: - chunk_count, content_hash = result + processed_documents.append(doc_info) + if isinstance(result, tuple) and len(result) == 3: + chunk_count, content_hash, failed_chunks = result self.stats.total_chunks_processed += chunk_count + self.stats.total_chunks_failed += failed_chunks # Track chunk count using content_hash (not directory hash) chunks_info[content_hash] = {"chunk_count": chunk_count} logger.info( - f"CHUNK COUNT: Document {doc_info.document_hash[:12]}... (content: {content_hash[:12]}...) -> {chunk_count} chunks" + f"CHUNK COUNT: Document {doc_info.document_hash[:12]}... (content: {content_hash[:12]}...) -> {chunk_count} chunks ({failed_chunks} failed)" ) - # Log the complete chunks_info dictionary + # Log summary only (avoid massive logs for large datasets) + total_chunks = sum(info["chunk_count"] for info in chunks_info.values()) logger.info( - f"CHUNKS INFO SUMMARY: {len(chunks_info)} documents tracked" + f"CHUNKS INFO SUMMARY: {len(chunks_info)} documents tracked, {total_chunks} total chunks" ) - for doc_hash, info in chunks_info.items(): - logger.info( - f" {doc_hash[:12]}... -> {info['chunk_count']} chunks" - ) # Calculate final statistics self.stats.end_time = datetime.now() @@ -227,10 +231,10 @@ async def process_all_documents(self) -> ProcessingStats: # Step 4: Update processed files tracking (even if no new documents processed) if diff_detector: try: - # Update metadata for newly processed files - if documents: + # Update metadata for newly processed files (successful only) + if processed_documents: processed_paths = [ - doc.cleaned_txt_path for doc in documents + doc.cleaned_txt_path for doc in processed_documents ] if processed_paths: logger.debug( @@ -290,7 +294,7 @@ async def _process_single_document( doc_info: DocumentInfo, qdrant_manager: QdrantManager, semaphore: asyncio.Semaphore, - ) -> tuple[int, str]: + ) -> tuple[int, str, int]: """ Process a single document with contextual retrieval. @@ -300,7 +304,9 @@ async def _process_single_document( semaphore: Concurrency control semaphore Returns: - tuple: (chunk_count: int, content_hash: str) or Exception on error + tuple: (chunk_count: int, content_hash: str, failed_chunks: int). + Raises on any failure (including load failure or zero usable chunks), + so the document is counted as failed rather than as success. """ async with semaphore: logger.info(f"Processing document: {doc_info.document_hash}") @@ -310,29 +316,31 @@ async def _process_single_document( document = self.document_loader.load_document(doc_info) if not document: - logger.warning(f"Could not load document: {doc_info.document_hash}") - return (0, doc_info.document_hash) + raise DocumentLoadError( + f"Could not load document: {doc_info.document_hash}" + ) # Process document with contextual retrieval - contextual_chunks = await self.contextual_processor.process_document( - document - ) + ( + contextual_chunks, + failed_chunks, + ) = await self.contextual_processor.process_document(document) if not contextual_chunks: - logger.warning( - f"No chunks created for document: {doc_info.document_hash}" + raise RuntimeError( + f"No chunks created for document: {doc_info.document_hash} " + f"({failed_chunks} chunks failed context generation)" ) - return (0, document.document_hash) # Store chunks in Qdrant await qdrant_manager.store_chunks(contextual_chunks) logger.info( f"Successfully processed document {doc_info.document_hash}: " - f"{len(contextual_chunks)} chunks" + f"{len(contextual_chunks)} chunks ({failed_chunks} dropped)" ) - return (len(contextual_chunks), document.document_hash) + return (len(contextual_chunks), document.document_hash, failed_chunks) except Exception as e: logger.error(f"Error processing document {doc_info.document_hash}: {e}") @@ -352,10 +360,12 @@ def _log_final_summary(self) -> None: logger.info(f" • Failed Chunks: {self.stats.total_chunks_failed}") if self.stats.total_documents > 0: - success_rate = ( - self.stats.documents_processed / self.stats.total_documents - ) * 100 - logger.info(f"Success Rate: {success_rate:.1f}%") + logger.info(f"Success Rate: {self.stats.success_rate * 100:.1f}%") + + if self.stats.total_chunks_processed + self.stats.total_chunks_failed > 0: + logger.info( + f"Chunk Success Rate: {self.stats.chunk_success_rate * 100:.1f}%" + ) logger.info(f"Processing Duration: {self.stats.duration}") @@ -365,6 +375,11 @@ def _log_final_summary(self) -> None: ) logger.info("Check failure logs for details") + if self.stats.total_chunks_failed > 0: + logger.warning( + f" {self.stats.total_chunks_failed} chunks failed processing" + ) + async def run_health_check(self) -> bool: """ Run health check on all components. @@ -617,12 +632,20 @@ async def _execute_cleanup_operations( return total_deleted def _cleanup_datasets(self) -> None: - """Remove datasets folder after processing.""" + """Remove datasets folder contents after processing. + + Only the folder's contents are removed, not the folder itself, since + the datasets path is a mounted volume in the container. + """ try: datasets_path = Path(self.config.dataset_base_path) if datasets_path.exists(): - shutil.rmtree(str(datasets_path)) - logger.info(f"Datasets folder cleaned up: {datasets_path}") + for child in datasets_path.iterdir(): + if child.is_dir(): + shutil.rmtree(str(child)) + else: + child.unlink() + logger.info(f"Datasets folder contents cleaned up: {datasets_path}") else: logger.debug(f"Datasets folder does not exist: {datasets_path}") except Exception as e: @@ -640,22 +663,8 @@ async def main() -> int: parser.add_argument("--signed-url", help="Signed URL for dataset download") args = parser.parse_args() - # Configure logging - logger.remove() # Remove default handler - logger.add( - sys.stdout, - format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}", - level="INFO", - ) - - # Add file logging - logger.add( - "vector_indexer.log", - format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}", - level="DEBUG", - rotation="10 MB", - retention="7 days", - ) + # LokiLogger handles all logging (console + Loki service) + # No additional configuration needed indexer = None try: diff --git a/src/vector_indexer/models.py b/src/vector_indexer/models.py index 752ea02a..41ae1ce1 100644 --- a/src/vector_indexer/models.py +++ b/src/vector_indexer/models.py @@ -96,6 +96,14 @@ def success_rate(self) -> float: return self.documents_processed / self.total_documents return 0.0 + @property + def chunk_success_rate(self) -> float: + """Calculate chunk success rate (processed vs processed + failed).""" + total_chunks = self.total_chunks_processed + self.total_chunks_failed + if total_chunks > 0: + return self.total_chunks_processed / total_chunks + return 0.0 + class ProcessingError(BaseModel): """Error information for failed processing.""" diff --git a/src/vector_indexer/qdrant_manager.py b/src/vector_indexer/qdrant_manager.py index 08664652..d465bcd2 100644 --- a/src/vector_indexer/qdrant_manager.py +++ b/src/vector_indexer/qdrant_manager.py @@ -1,7 +1,7 @@ """Qdrant vector database manager for storing contextual chunks.""" from typing import List, Dict, Any, Optional -from loguru import logger +from loki_logger import LokiLogger import httpx import uuid from typing_extensions import Self @@ -9,6 +9,8 @@ from vector_indexer.config.config_loader import VectorIndexerConfig from vector_indexer.models import ContextualChunk +logger = LokiLogger(service_name="qdrant-manager") + class QdrantOperationError(Exception): """Custom exception for Qdrant operations.""" @@ -171,22 +173,10 @@ async def _store_chunks_in_collection( upsert_payload = {"points": batch} - # DEBUG: Log the actual HTTP request payload being sent to Qdrant - logger.info("=== QDRANT HTTP REQUEST PAYLOAD DEBUG ===") - logger.info( - f"URL: {self.qdrant_url}/collections/{collection_name}/points" + # Log batch summary only (avoid massive logs) + logger.debug( + f"Storing batch {i // batch_size + 1}: {len(batch)} points to {collection_name}" ) - logger.info("Method: PUT") - logger.info(f"Batch size: {len(batch)} points") - for idx, point in enumerate(batch): - logger.info(f"Point {idx + 1}:") - logger.info(f" ID: {point['id']} (type: {type(point['id'])})") - logger.info( - f" Vector length: {len(point['vector'])} (type: {type(point['vector'])})" - ) - logger.info(f" Vector sample: {point['vector'][:3]}...") - logger.info(f" Payload keys: {list(point['payload'].keys())}") - logger.info("=== END QDRANT REQUEST DEBUG ===") response = await self.client.put( f"{self.qdrant_url}/collections/{collection_name}/points", diff --git a/src/vector_indexer/vector_indexer_integration.md b/src/vector_indexer/vector_indexer_integration.md index d6b10b22..a160c51e 100644 --- a/src/vector_indexer/vector_indexer_integration.md +++ b/src/vector_indexer/vector_indexer_integration.md @@ -161,15 +161,18 @@ chunking: async def generate_context_batch(self, document_content: str, chunks: List[str]): # Level 1: Batch processing (context_batch_size = 5) for i in range(0, len(chunks), self.config.context_batch_size): - batch = chunks[i:i + self.config.context_batch_size] - + batch = chunks[i : i + self.config.context_batch_size] + # Level 2: Semaphore limiting (max_concurrent_chunks_per_doc = 5) semaphore = asyncio.Semaphore(self.config.max_concurrent_chunks_per_doc) - + # Process batch concurrently with controlled limits batch_contexts = await asyncio.gather( - *[self._generate_context_with_retry(document_content, chunk) for chunk in batch], - return_exceptions=True + *[ + self._generate_context_with_retry(document_content, chunk) + for chunk in batch + ], + return_exceptions=True, ) ``` @@ -220,15 +223,15 @@ graph LR # Configuration-Driven Batch Optimization async def _create_embeddings_in_batches(self, contextual_contents: List[str]): all_embeddings = [] - + # Process in configurable batches (embedding_batch_size = 10) for i in range(0, len(contextual_contents), self.config.embedding_batch_size): - batch = contextual_contents[i:i + self.config.embedding_batch_size] - + batch = contextual_contents[i : i + self.config.embedding_batch_size] + # API call with comprehensive error handling batch_response = await self.api_client.create_embeddings_batch(batch) all_embeddings.extend(batch_response["embeddings"]) - + # Configurable delay between batches if i + self.config.embedding_batch_size < len(contextual_contents): delay = self.config.processing.batch_delay_seconds # 0.1s @@ -298,9 +301,9 @@ graph TD ```python # Step 5: Add embeddings to chunks with full traceability for chunk, embedding in zip(contextual_chunks, embeddings_response["embeddings"]): - chunk.embedding = embedding # Vector data - chunk.embedding_model = embeddings_response["model_used"] # Model traceability - chunk.vector_dimensions = len(embedding) # Dimension validation + chunk.embedding = embedding # Vector data + chunk.embedding_model = embeddings_response["model_used"] # Model traceability + chunk.vector_dimensions = len(embedding) # Dimension validation # Provider automatically detected from model name ``` @@ -315,18 +318,18 @@ self.collections_config = { "contextual_chunks_azure": { "vector_size": 3072, # text-embedding-3-large (Azure) "distance": "Cosine", - "models": ["text-embedding-3-large", "text-embedding-ada-002"] + "models": ["text-embedding-3-large", "text-embedding-ada-002"], }, "contextual_chunks_aws": { "vector_size": 1024, # amazon.titan-embed-text-v2:0 - "distance": "Cosine", - "models": ["amazon.titan-embed-text-v2:0", "amazon.titan-embed-text-v1"] + "distance": "Cosine", + "models": ["amazon.titan-embed-text-v2:0", "amazon.titan-embed-text-v1"], }, "contextual_chunks_openai": { "vector_size": 1536, # text-embedding-3-small (Direct OpenAI) "distance": "Cosine", - "models": ["text-embedding-3-small", "text-embedding-ada-002"] - } + "models": ["text-embedding-3-small", "text-embedding-ada-002"], + }, } ``` @@ -336,9 +339,9 @@ self.collections_config = { point_id = str(uuid.uuid5(uuid.NAMESPACE_DNS, chunk.chunk_id)) point = { - "id": point_id, # Deterministic UUID - "vector": chunk.embedding, # Provider-specific dimensions - "payload": self._create_chunk_payload(chunk) # Rich metadata + "id": point_id, # Deterministic UUID + "vector": chunk.embedding, # Provider-specific dimensions + "payload": self._create_chunk_payload(chunk), # Rich metadata } ``` @@ -347,15 +350,15 @@ point = { # Production-Grade Batch Processing batch_size = 100 # Prevents request timeout issues for i in range(0, len(points), batch_size): - batch = points[i:i + batch_size] - + batch = points[i : i + batch_size] + # Comprehensive request logging for debugging logger.info(f"=== QDRANT HTTP REQUEST PAYLOAD DEBUG ===") logger.info(f"Batch size: {len(batch)} points") - + response = await self.client.put( f"{self.qdrant_url}/collections/{collection_name}/points", - json={"points": batch} + json={"points": batch}, ) ``` @@ -367,22 +370,19 @@ for i in range(0, len(points), batch_size): "document_hash": "2e9493512b7f01aecdc66bbca60b5b6b75d966f8", "chunk_index": 0, "total_chunks": 25, - # Anthropic Contextual Retrieval Content "original_content": "FAQ about supporting children and families...", "contextual_content": "Estonian family support policies context. FAQ about...", "context_only": "Estonian family support policies context.", - # Model & Processing Metadata - "embedding_model": "text-embedding-3-large", + "embedding_model": "text-embedding-3-large", "vector_dimensions": 3072, "processing_timestamp": "2025-10-09T12:00:00Z", "tokens_count": 150, - # Document Source Information "document_url": "https://sm.ee/en/faq-about-supporting-children-and-families", "dataset_collection": "sm_someuuid", - "file_type": "html_cleaned" + "file_type": "html_cleaned", } ``` @@ -486,9 +486,9 @@ The Vector Indexer leverages existing LLM configuration through API calls: ```python # Process chunks in batches of 5 with concurrent API calls for batch in chunks_batches(5): - contexts = await asyncio.gather(*[ - api_client.generate_context(document, chunk) for chunk in batch - ]) + contexts = await asyncio.gather( + *[api_client.generate_context(document, chunk) for chunk in batch] + ) ``` 5. **Contextual Chunk Creation** @@ -547,12 +547,12 @@ logs/ collections = { "contextual_chunks_azure": { "vectors": {"size": 1536, "distance": "Cosine"}, # text-embedding-3-large - "model": "text-embedding-3-large" + "model": "text-embedding-3-large", }, "contextual_chunks_aws": { "vectors": {"size": 1024, "distance": "Cosine"}, # amazon.titan-embed-text-v2:0 - "model": "amazon.titan-embed-text-v2:0" - } + "model": "amazon.titan-embed-text-v2:0", + }, } ``` @@ -572,7 +572,7 @@ collections = { "embedding_model": "text-embedding-3-large", "vector_dimensions": 1536, "processing_timestamp": "2025-10-08T12:00:00Z", - "tokens_count": 150 + "tokens_count": 150, } ``` @@ -692,9 +692,9 @@ vector_indexer: class ResourceOptimizedProcessor: def __init__(self): # Process in streaming fashion - never load all documents - self.max_memory_chunks = 100 # Chunk buffer limit - self.gc_frequency = 50 # Garbage collection interval - + self.max_memory_chunks = 100 # Chunk buffer limit + self.gc_frequency = 50 # Garbage collection interval + async def process_documents_streaming(self): """Memory-efficient document processing""" async for document_batch in self.stream_documents(): @@ -720,22 +720,22 @@ class ResourceOptimizedProcessor: "embeddings_created": 26834, "qdrant_points_stored": 26834, "processing_duration_minutes": 186.5, - "average_chunks_per_document": 21.6 + "average_chunks_per_document": 21.6, }, "performance_metrics": { "context_generation_rate_per_minute": 14.4, "embedding_creation_rate_per_minute": 187.3, "end_to_end_documents_per_hour": 10.1, "api_success_rate": 99.7, - "average_response_time_ms": 850 + "average_response_time_ms": 850, }, "error_analysis": { "api_timeouts": 2, "rate_limit_hits": 1, "embedding_dimension_mismatches": 0, "qdrant_storage_failures": 0, - "context_generation_failures": 2 - } + "context_generation_failures": 2, + }, } ``` @@ -773,7 +773,7 @@ logger.info( document_hash="2e9493512b7f01aecdc66bbca60b5b6b75d966f8", document_path="datasets/sm_someuuid/2e9493.../cleaned.txt", chunk_count=23, - processing_id="proc_20241009_120034_789" + processing_id="proc_20241009_120034_789", ) logger.info( @@ -782,7 +782,7 @@ logger.info( model_used="claude-3-haiku-20240307", context_tokens=75, generation_time_ms=1247, - cached_response=False + cached_response=False, ) ``` diff --git a/store-langfuse-secrets.sh b/store-langfuse-secrets.sh index 234457e2..886986e2 100644 --- a/store-langfuse-secrets.sh +++ b/store-langfuse-secrets.sh @@ -92,13 +92,16 @@ echo " Host: $LANGFUSE_HOST" # Update Vault policy to include Langfuse secrets access echo "" -echo "Updating llm-orchestration policy to include Langfuse secrets..." -POLICY='path "secret/metadata/llm/*" { capabilities = ["list", "delete"] } -path "secret/data/llm/*" { capabilities = ["create", "read", "update", "delete"] } -path "secret/metadata/embeddings/*" { capabilities = ["list", "delete"] } -path "secret/data/embeddings/*" { capabilities = ["create", "read", "update", "delete"] } -path "secret/metadata/langfuse/*" { capabilities = ["list", "delete"] } -path "secret/data/langfuse/*" { capabilities = ["create", "read", "update", "delete"] } +echo "Updating llm-orchestration-policy to include Langfuse secrets..." +# Preserve the production policy paths (see vault-init.sh) and add Langfuse read. +# This is a full overwrite of the policy, so the existing grants must be repeated. +POLICY='path "secret/data/llm/connections/*" { capabilities = ["read", "list"] } +path "secret/metadata/llm/connections/*" { capabilities = ["read", "list"] } +path "secret/data/embeddings/connections/*" { capabilities = ["read", "list"] } +path "secret/metadata/embeddings/connections/*" { capabilities = ["read", "list"] } +path "secret/data/encryption/*" { capabilities = ["deny"] } +path "secret/data/langfuse/*" { capabilities = ["read"] } +path "secret/metadata/langfuse/*" { capabilities = ["read", "list"] } path "auth/token/lookup-self" { capabilities = ["read"] }' # Create JSON without jq (using printf for proper escaping) @@ -108,10 +111,12 @@ POLICY_JSON='{"policy":"'"$POLICY_ESCAPED"'"}' if wget -q -O- --post-data="$POLICY_JSON" \ --header="X-Vault-Token: $ROOT_TOKEN" \ --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/sys/policies/acl/llm-orchestration" >/dev/null 2>&1; then + "$VAULT_ADDR/v1/sys/policies/acl/llm-orchestration-policy" >/dev/null 2>&1; then echo "Policy updated successfully" else - echo "Warning: Policy update failed (may already be updated)" + echo "Error: Failed to update llm-orchestration-policy" + echo " Langfuse secrets would be stored but the agent would be denied access." + exit 1 fi # Store Langfuse secrets in Vault diff --git a/test-vault/agents/llm/agent.hcl b/test-vault/agents/llm/agent.hcl index 9883bfe4..69ffde4d 100644 --- a/test-vault/agents/llm/agent.hcl +++ b/test-vault/agents/llm/agent.hcl @@ -3,7 +3,7 @@ vault { address = "http://vault:8200" } -pid_file = "/agent/out/pidfile" +pid_file = "/agent/llm-token/pidfile" auto_auth { method "approle" { @@ -17,7 +17,7 @@ auto_auth { sink "file" { config = { - path = "/agent/out/token" + path = "/agent/llm-token/token" } } } @@ -36,7 +36,7 @@ listener "tcp" { # dummy template so cache is “active” (some versions require this) template { source = "/dev/null" - destination = "/agent/out/dummy" + destination = "/agent/llm-token/dummy" } # Disable API proxy; not needed here diff --git a/tests/api_tool_eval/integration_test_multi_intent.py b/tests/api_tool_eval/integration_test_multi_intent.py new file mode 100644 index 00000000..a896fc5c --- /dev/null +++ b/tests/api_tool_eval/integration_test_multi_intent.py @@ -0,0 +1,915 @@ +""" +Integration Test — Multi-Intent Parallel API Tool Classification +================================================================ + +Tests Phase 1 of the multi-intent ATC feature end-to-end via /orchestrate. + +The current Phase 1 implementation uses a temporary single-endpoint fallback: +IntentDecomposer detects parallel intent → both endpoints stored in context → +first matched endpoint used for param collection. True parallel execution +(collecting params for all endpoints simultaneously) arrives in Phase 6. + +What these tests verify +----------------------- +- Queries with multiple independent intents still route to ATC (not RAG/OOD) +- The temporary first-endpoint fallback correctly collects params and completes +- Single-intent queries in the ambiguous cosine band are unaffected (regression) +- Estonian multi-intent queries decompose and route correctly +- Sub-queries that both resolve to the same endpoint are deduplicated → single + path fallback (no crash, no duplicate calls) +- When one sub-query finds no matching endpoint (OOD intent mixed in), the + system gracefully falls back to the single matched endpoint +- A query spanning 3 distinct domains hits the MULTI_API_MAX_ENDPOINTS=3 cap + and still completes correctly + +Scenarios +--------- + MI-1 Parallel — vehicle tax + initiative details (2 distinct endpoints) + MI-2 Parallel — public holidays + electricity prices (distinct domains) + MI-3 Parallel — parliament votings + participation stats (same domain) + MI-4 Estonian multi-intent — vehicle tax + initiatives list + MI-5 Single-intent regression — address search (should stay single path) + MI-6 Deduplication fallback — both sub-queries hit the same endpoint + MI-7 Mixed ATC + OOD intent — one sub-query finds no endpoint → graceful fallback + MI-8 Three-domain query — address + holidays + electricity (3-way parallel cap) + +Usage +----- + # Service running locally on port 8100 + uv run python tests/api_tool_eval/integration_test_multi_intent.py + + # Against a different host/port + uv run python tests/api_tool_eval/integration_test_multi_intent.py --url http://localhost:8100 + + # Keep going after failures + uv run python tests/api_tool_eval/integration_test_multi_intent.py --no-fail-fast + + # Save results to JSON + uv run python tests/api_tool_eval/integration_test_multi_intent.py --output results-multi-intent.json +""" + +import argparse +import json +import sys +import time +import uuid +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional, Tuple + +import requests + +# --------------------------------------------------------------------------- +# Config +# --------------------------------------------------------------------------- +DEFAULT_URL = "http://localhost:8100" +ORCHESTRATE_ENDPOINT = "/orchestrate" +ENVIRONMENT = "production" +AUTHOR_ID = "multi-intent-test-user" +REQUEST_TIMEOUT = 45 # seconds + + +# --------------------------------------------------------------------------- +# Helpers (same interface as integration_test_agentic_loop.py) +# --------------------------------------------------------------------------- + + +def make_chat_id(label: str) -> str: + """Unique chatId per test run so Redis sessions never collide across runs.""" + return f"mi-test-{label}-{uuid.uuid4().hex[:8]}" + + +def send_turn( + base_url: str, + chat_id: str, + message: str, + history: List[Dict[str, str]], + connection_id: Optional[str] = None, +) -> Dict[str, Any]: + """POST one turn to /orchestrate and return the parsed JSON response.""" + payload: Dict[str, Any] = { + "chatId": chat_id, + "message": message, + "authorId": AUTHOR_ID, + "conversationHistory": history, + "url": "integration-test", + "environment": ENVIRONMENT, + } + if connection_id: + payload["connection_id"] = connection_id + + resp = requests.post( + f"{base_url}{ORCHESTRATE_ENDPOINT}", + json=payload, + timeout=REQUEST_TIMEOUT, + ) + resp.raise_for_status() + return resp.json() + + +def append_to_history( + history: List[Dict[str, str]], + user_message: str, + bot_response: str, +) -> List[Dict[str, str]]: + """Return an updated conversation history list.""" + ts = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()) + return history + [ + {"authorRole": "user", "message": user_message, "timestamp": ts}, + {"authorRole": "bot", "message": bot_response, "timestamp": ts}, + ] + + +def is_completed(content: str) -> bool: + """Return True if the response is a params-collected JSON payload.""" + try: + data = json.loads(content) + return "collected_params" in data and "endpoint" in data + except (json.JSONDecodeError, TypeError): + return False + + +def is_clarifying_question(content: str) -> bool: + """Return True if the response is a non-JSON natural-language question.""" + try: + json.loads(content) + return False + except (json.JSONDecodeError, TypeError): + return bool(content.strip()) + + +def endpoint_name_from_completed(content: str) -> Optional[str]: + """Extract the endpoint name from a completed JSON response.""" + try: + data = json.loads(content) + ep = data.get("endpoint", {}) + return ep.get("name") if isinstance(ep, dict) else None + except (json.JSONDecodeError, TypeError): + return None + + +def collected_params_from_completed(content: str) -> Dict[str, Any]: + """Extract collected_params dict from a completed JSON response.""" + try: + data = json.loads(content) + return data.get("collected_params", {}) + except (json.JSONDecodeError, TypeError): + return {} + + +# --------------------------------------------------------------------------- +# Result tracking (same dataclasses as integration_test_agentic_loop.py) +# --------------------------------------------------------------------------- + + +@dataclass +class TurnResult: + turn: int + message_sent: str + response_content: str + passed: bool + note: str = "" + + +@dataclass +class ScenarioResult: + name: str + passed: bool + turns: List[TurnResult] = field(default_factory=list) + error: str = "" + + +# --------------------------------------------------------------------------- +# Test scenarios +# --------------------------------------------------------------------------- + + +def scenario_mi1_parallel_vehicle_tax_and_initiative( + base_url: str, +) -> ScenarioResult: + """ + MI-1: Parallel — vehicle tax + initiative details + ------------------------------------------------- + Query combines two independent intents: + - Calculate vehicle tax (requires regNr + calculationYear) + - Get initiative details (requires initiative id) + + Expected flow: + Turn 1: IntentDecomposer fires (cosine in ambiguous band [0.40, 0.60)). + Parallel detected. Temporary fallback → get_vehicle_tax_info. + Bot asks for registration number and calculation year. + Turn 2: User provides vehicle params. + Bot returns completed JSON for get_vehicle_tax_info. + + Asserts: + - T1 is a clarifying question (not RAG/OOD — system routed to ATC) + - T1 question references vehicle/registration/tax + - T2 is completed JSON with regNr + calculationYear collected + - Completed endpoint is get_vehicle_tax_info + """ + name = "MI-1 — Parallel: vehicle tax + initiative details (first-endpoint fallback)" + chat_id = make_chat_id("mi1") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + # Turn 1 — multi-intent query + msg1 = ( + "Can you calculate the tax for my vehicle and also get the details " + "of the initiative?" + ) + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + t1_pass = is_clarifying_question(content1) and not is_completed(content1) + # Check that the clarifying question is about vehicle tax params, not initiative params. + # The temporary fallback must pick get_vehicle_tax_info as the first endpoint. + vehicle_hint = any( + kw in content1.lower() + for kw in ("registration", "reg", "vehicle", "tax", "year", "registreeri") + ) + t1_note = ( + f"Correctly asked for vehicle params (vehicle_hint={vehicle_hint})" + if t1_pass + else f"Expected clarifying question about vehicle tax, got: {content1[:150]}" + ) + if t1_pass and not vehicle_hint: + t1_note = ( + f"WARNING: clarifying question does not mention vehicle/registration. " + f"Fallback endpoint may differ. Content: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult(name, False, turns, "Turn 1 did not route to ATC") + + # Turn 2 — provide vehicle tax params + history = append_to_history(history, msg1, content1) + msg2 = "Registration number 123ABC, year 2026" + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + collected = collected_params_from_completed(content2) + ep_name = endpoint_name_from_completed(content2) + expected_keys = {"regNr", "calculationYear"} + missing = expected_keys - collected.keys() + t2_pass = not missing and ep_name == "get_vehicle_tax_info" + note2 = ( + f"endpoint={ep_name}, collected_params={collected}" + if t2_pass + else ( + f"missing keys={missing}" if missing else f"wrong endpoint={ep_name!r}" + ) + ) + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +def scenario_mi2_parallel_holidays_and_electricity( + base_url: str, +) -> ScenarioResult: + """ + MI-2: Parallel — public holidays + electricity prices (distinct domains) + ------------------------------------------------------------------------ + Two clearly unrelated API intents in one query. + + Expected flow: + Turn 1: IntentDecomposer detects parallel. + One endpoint selected as fallback → clarifying question. + Turn 2: User provides params → completed JSON. + + Asserts: + - T1 routes to ATC (clarifying question, not RAG) + - T2 completes with collected_params for the fallback endpoint + """ + name = "MI-2 — Parallel: public holidays + electricity prices (distinct domains)" + chat_id = make_chat_id("mi2") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + msg1 = ( + "Can you get the public holidays in Estonia for this year " + "and also show me the electricity market prices for this week?" + ) + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + t1_pass = is_clarifying_question(content1) and not is_completed(content1) + t1_note = ( + "Correctly routed to ATC and asked for params" + if t1_pass + else f"Expected clarifying question, got: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult(name, False, turns, "Turn 1 did not route to ATC") + + # Provide params that satisfy either fallback endpoint: + # - get_public_holidays: countryIsoCode=EE, validFrom/To + electricity start/end + # - get_electricity_prices: start/end datetime + # We supply all possible params so whichever endpoint was chosen can complete. + history = append_to_history(history, msg1, content1) + msg2 = ( + "Country EE, from 2026-01-01 to 2026-12-31, " + "electricity from 2026-05-12T00:00:00Z to 2026-05-18T23:59:59Z" + ) + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + ep_name = endpoint_name_from_completed(content2) + collected = collected_params_from_completed(content2) + # Accept either endpoint as the fallback + valid_endpoints = {"get_public_holidays", "get_electricity_prices"} + t2_pass = ep_name in valid_endpoints and bool(collected) + note2 = ( + f"endpoint={ep_name!r}, collected_params={collected}" + if t2_pass + else f"unexpected endpoint={ep_name!r} or empty params={collected}" + ) + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +def scenario_mi3_parallel_parliament_same_domain( + base_url: str, +) -> ScenarioResult: + """ + MI-3: Parallel — parliament votings + participation stats (same domain) + ----------------------------------------------------------------------- + Both intents relate to parliament but map to different endpoints. + Tests that same-domain multi-intent is still detected as parallel. + + Expected flow: + Turn 1: IntentDecomposer detects 2 parliament intents. + First parliament endpoint selected as fallback → clarifying question + asking for startDate + endDate. + Turn 2: User provides date range → completed JSON. + + Asserts: + - T1 routes to ATC + - T2 completes for a parliament endpoint with date params + """ + name = "MI-3 — Parallel: parliament votings + participation stats (same domain)" + chat_id = make_chat_id("mi3") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + msg1 = ( + "Show me the parliament voting records and also the participation " + "statistics of parliament members in Estonia" + ) + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + t1_pass = is_clarifying_question(content1) and not is_completed(content1) + t1_note = ( + "Routed to ATC — asking for date params" + if t1_pass + else f"Expected ATC clarifying question, got: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult(name, False, turns, "Turn 1 did not route to ATC") + + history = append_to_history(history, msg1, content1) + msg2 = "From 2026-01-01 to 2026-03-31" + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + ep_name = endpoint_name_from_completed(content2) + collected = collected_params_from_completed(content2) + valid_endpoints = { + "get_parliament_votings", + "get_parliament_participation_stats", + } + expected_keys = {"startDate", "endDate"} + missing = expected_keys - collected.keys() + t2_pass = ep_name in valid_endpoints and not missing + note2 = ( + f"endpoint={ep_name!r}, collected_params={collected}" + if t2_pass + else ( + f"missing keys={missing}" + if missing + else f"unexpected endpoint={ep_name!r}" + ) + ) + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +def scenario_mi4_parallel_estonian(base_url: str) -> ScenarioResult: + """ + MI-4: Estonian multi-intent — vehicle tax + initiatives list + ------------------------------------------------------------ + Verifies that the IntentDecomposer handles Estonian queries correctly. + DSPy signature explicitly lists ET/EN/RU as supported languages. + + Expected flow: + Turn 1: Estonian query → IntentDecomposer fires → parallel detected. + Fallback: get_vehicle_tax_info or get_initiatives → clarifying + question returned (in Estonian or English). + Turn 2: Provide the expected params. + Completed JSON. + + Asserts: + - T1 routes to ATC (clarifying question, not RAG/OOD) + - T2 completes with collected_params + """ + name = "MI-4 — Estonian multi-intent: vehicle tax + initiatives list" + chat_id = make_chat_id("mi4") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + msg1 = "Arvuta mu sõiduki maks ja näita mulle algatuste nimekiri" + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + t1_pass = is_clarifying_question(content1) and not is_completed(content1) + t1_note = ( + "Estonian multi-intent routed to ATC" + if t1_pass + else f"Expected ATC clarifying question for Estonian query, got: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult( + name, False, turns, "Turn 1 did not route Estonian multi-intent to ATC" + ) + + # Provide params covering both possible fallback endpoints + history = append_to_history(history, msg1, content1) + msg2 = "Registreerimisnumber 456DEF, aasta 2026" + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + ep_name = endpoint_name_from_completed(content2) + collected = collected_params_from_completed(content2) + note2 = f"endpoint={ep_name!r}, collected_params={collected}" + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +def scenario_mi5_single_intent_regression(base_url: str) -> ScenarioResult: + """ + MI-5: Single-intent regression — address search (not multi-intent) + ------------------------------------------------------------------ + Verifies that a clear single-intent query is NOT incorrectly split into + sub-queries by the IntentDecomposer. + + "Search for an address in Tallinn" has only one intent. + IntentDecomposer should return mode=single → single path used as before. + + Expected flow: + Turn 1: search_address matched → single path → bot asks for address param. + Turn 2: user provides address → completed JSON with address param. + + Asserts: + - T1 is a clarifying question asking for the address + - T2 completes for search_address with the address param present + """ + name = "MI-5 — Single-intent regression: address search stays on single path" + chat_id = make_chat_id("mi5") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + msg1 = "I need to search for an address in Tallinn" + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + t1_pass = is_clarifying_question(content1) and not is_completed(content1) + t1_note = ( + "Single-intent routed to ATC, asking for address" + if t1_pass + else f"Expected clarifying question for address, got: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult( + name, False, turns, "Turn 1 did not route single-intent to ATC" + ) + + history = append_to_history(history, msg1, content1) + msg2 = "Viru 4, Tallinn" + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + ep_name = endpoint_name_from_completed(content2) + collected = collected_params_from_completed(content2) + t2_pass = ep_name == "search_address" and "address" in collected + note2 = ( + f"endpoint={ep_name!r}, collected_params={collected}" + if t2_pass + else f"unexpected endpoint={ep_name!r} or missing 'address' in {collected}" + ) + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +def scenario_mi6_dedup_same_endpoint(base_url: str) -> ScenarioResult: + """ + MI-6: Deduplication — both sub-queries resolve to the same endpoint + ------------------------------------------------------------------- + When IntentDecomposer generates sub-queries that both match the same + endpoint, _try_parallel_api_tool_classification deduplicates them to + 1 unique endpoint → < 2 required → returns None → falls back to the + original single-path match. + + This verifies: + - No crash or duplicate call when dedup collapses to 1 endpoint + - The single-path fallback still works correctly + + Query: "Show me the list of initiatives and also display all initiatives" + Expected: get_initiatives matched via single path → completed on turn 1 + (get_initiatives has no required params → fast-path completion) + """ + name = "MI-6 — Dedup fallback: both sub-queries hit get_initiatives → single path" + chat_id = make_chat_id("mi6") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + msg1 = ( + "Show me the list of all citizen initiatives and also display all initiatives" + ) + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + # get_initiatives has only an optional 'page' param → may complete immediately + # or ask for the page number. Both are valid. + if is_completed(content1): + ep_name = endpoint_name_from_completed(content1) + passed = ep_name == "get_initiatives" + note = ( + f"Fast-path completion: endpoint={ep_name!r}" + if passed + else f"Completed for wrong endpoint: {ep_name!r}" + ) + turns.append(TurnResult(1, msg1, content1, passed, note)) + return ScenarioResult(name, passed, turns) + + # Bot asked a clarifying question (e.g., for page number) + t1_pass = is_clarifying_question(content1) + t1_note = ( + "ATC routed correctly after dedup, asking for optional page param" + if t1_pass + else f"Expected ATC response, got: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult(name, False, turns, "Turn 1 did not route to ATC") + + history = append_to_history(history, msg1, content1) + msg2 = "Page 1" + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + ep_name = endpoint_name_from_completed(content2) + t2_pass = ep_name == "get_initiatives" + note2 = ( + f"endpoint={ep_name!r}" if t2_pass else f"unexpected endpoint={ep_name!r}" + ) + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +def scenario_mi7_mixed_atc_ood_intent(base_url: str) -> ScenarioResult: + """ + MI-7: Mixed ATC + OOD intent — graceful fallback to single matched endpoint + --------------------------------------------------------------------------- + One intent maps cleanly to an endpoint; the other is out-of-domain (no + matching endpoint with cosine ≥ threshold). + + When _try_parallel_api_tool_classification collects results, the OOD + sub-query returns None → only 1 unique endpoint matched → < 2 required + → returns None → original single-path match used as fallback. + + Query: "Calculate my vehicle tax and explain climate change in Estonia" + Expected: + - "explain climate change" has no matching API endpoint + - Parallel collapses to 1 → single path → get_vehicle_tax_info + - Bot asks for regNr + calculationYear + - Turn 2 provides params → completes + + Asserts: + - T1 routes to ATC (not OOD/RAG) — vehicle tax intent rescued the query + - T2 completes for get_vehicle_tax_info + """ + name = "MI-7 — Mixed ATC+OOD: one sub-query OOD → graceful single-path fallback" + chat_id = make_chat_id("mi7") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + msg1 = "Calculate my vehicle tax and also explain climate change in Estonia" + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + t1_pass = is_clarifying_question(content1) and not is_completed(content1) + vehicle_hint = any( + kw in content1.lower() + for kw in ("registration", "reg", "vehicle", "tax", "year", "registreeri") + ) + t1_note = ( + f"ATC routed for vehicle tax intent (vehicle_hint={vehicle_hint})" + if t1_pass + else f"Expected ATC clarifying question, got: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult( + name, False, turns, "Turn 1 did not fall back to ATC for vehicle tax intent" + ) + + history = append_to_history(history, msg1, content1) + msg2 = "Registration number 789GHI, calculation year 2026" + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + ep_name = endpoint_name_from_completed(content2) + collected = collected_params_from_completed(content2) + expected_keys = {"regNr", "calculationYear"} + missing = expected_keys - collected.keys() + t2_pass = ep_name == "get_vehicle_tax_info" and not missing + note2 = ( + f"endpoint={ep_name!r}, collected_params={collected}" + if t2_pass + else ( + f"missing keys={missing}" + if missing + else f"unexpected endpoint={ep_name!r}" + ) + ) + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +def scenario_mi8_three_domain_parallel(base_url: str) -> ScenarioResult: + """ + MI-8: Three-domain query — address + public holidays + electricity prices + ------------------------------------------------------------------------- + Tests MULTI_API_MAX_ENDPOINTS=3 cap. The IntentDecomposer should produce + 3 sub-queries (capped at 3), each mapping to a distinct endpoint: + - "search for an address" → search_address + - "get public holidays" → get_public_holidays + - "check electricity prices" → get_electricity_prices + + _try_parallel_api_tool_classification deduplicates: 3 unique endpoints + → parallel with 3. Temporary fallback uses the first matched endpoint. + + Expected flow: + Turn 1: 3-way parallel detected → first endpoint's params asked. + Turn 2: Provide params covering the first endpoint → completes. + + Asserts: + - T1 routes to ATC + - T2 completes for one of the three expected endpoints + """ + name = "MI-8 — Three-domain: address + holidays + electricity (3-way parallel cap)" + chat_id = make_chat_id("mi8") + history: List[Dict[str, str]] = [] + turns: List[TurnResult] = [] + + msg1 = ( + "Search for an address in Tartu, get public holidays in Estonia this year, " + "and show me the electricity prices for this week" + ) + resp1 = send_turn(base_url, chat_id, msg1, history) + content1 = resp1.get("content", "") + + t1_pass = is_clarifying_question(content1) and not is_completed(content1) + t1_note = ( + "3-domain query routed to ATC, asking for first endpoint's params" + if t1_pass + else f"Expected ATC clarifying question, got: {content1[:150]}" + ) + turns.append(TurnResult(1, msg1, content1, t1_pass, t1_note)) + + if not t1_pass: + return ScenarioResult( + name, False, turns, "Turn 1 did not route to ATC for 3-domain query" + ) + + # Supply params covering all three possible first-endpoint fallbacks + history = append_to_history(history, msg1, content1) + msg2 = ( + "Address: Raekoja plats, Tartu. " + "Country EE, from 2026-01-01 to 2026-12-31. " + "Electricity from 2026-05-12T00:00:00Z to 2026-05-18T23:59:59Z" + ) + resp2 = send_turn(base_url, chat_id, msg2, history) + content2 = resp2.get("content", "") + + t2_pass = is_completed(content2) + note2 = "" + if t2_pass: + ep_name = endpoint_name_from_completed(content2) + collected = collected_params_from_completed(content2) + valid_endpoints = { + "search_address", + "get_public_holidays", + "get_electricity_prices", + } + t2_pass = ep_name in valid_endpoints and bool(collected) + note2 = ( + f"endpoint={ep_name!r}, collected_params={collected}" + if t2_pass + else f"unexpected endpoint={ep_name!r} or empty params={collected}" + ) + else: + note2 = f"Expected completed JSON, got: {content2[:150]}" + + turns.append(TurnResult(2, msg2, content2, t2_pass, note2)) + overall = all(t.passed for t in turns) + return ScenarioResult(name, overall, turns) + + +# --------------------------------------------------------------------------- +# Runner +# --------------------------------------------------------------------------- + +SCENARIOS = [ + scenario_mi1_parallel_vehicle_tax_and_initiative, + scenario_mi2_parallel_holidays_and_electricity, + scenario_mi3_parallel_parliament_same_domain, + scenario_mi4_parallel_estonian, + scenario_mi5_single_intent_regression, + scenario_mi6_dedup_same_endpoint, + scenario_mi7_mixed_atc_ood_intent, + scenario_mi8_three_domain_parallel, +] + + +def run_all( + base_url: str, fail_fast: bool = True +) -> Tuple[List[ScenarioResult], int, int]: + results: List[ScenarioResult] = [] + passed = 0 + failed = 0 + + print(f"\n{'=' * 70}") + print(f" Multi-Intent Integration Tests | {base_url}") + print(f"{'=' * 70}\n") + + for fn in SCENARIOS: + print(f"Running: {fn.__name__} ...", flush=True) + try: + result = fn(base_url) + except requests.exceptions.ConnectionError: + result = ScenarioResult( + fn.__name__, + False, + error="Connection refused — is the service running?", + ) + except requests.exceptions.Timeout: + result = ScenarioResult( + fn.__name__, + False, + error=f"Request timed out after {REQUEST_TIMEOUT}s", + ) + except Exception as exc: + result = ScenarioResult(fn.__name__, False, error=str(exc)) + + results.append(result) + status_icon = "✅" if result.passed else "❌" + print(f" {status_icon} {result.name}") + + for t in result.turns: + turn_icon = " ✓" if t.passed else " ✗" + print(f" {turn_icon} Turn {t.turn}: {t.message_sent[:70]!r}") + if t.note: + print(f" → {t.note}") + if not t.passed: + print(f" Response: {t.response_content[:250]}") + + if result.error: + print(f" ERROR: {result.error}") + + if result.passed: + passed += 1 + else: + failed += 1 + if fail_fast: + print("\n⚠ Stopping early (--no-fail-fast to continue)\n") + break + + print() + + print(f"{'=' * 70}") + print(f" Results: {passed} passed, {failed} failed / {len(results)} run") + print(f"{'=' * 70}\n") + + return results, passed, failed + + +# --------------------------------------------------------------------------- +# CLI +# --------------------------------------------------------------------------- + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Multi-intent ATC integration tests", + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=__doc__, + ) + parser.add_argument( + "--url", + default=DEFAULT_URL, + help=f"Base URL of the orchestration service (default: {DEFAULT_URL})", + ) + parser.add_argument( + "--no-fail-fast", + action="store_true", + help="Continue running all scenarios even after a failure", + ) + parser.add_argument( + "--output", + help="Optional path to save detailed JSON results", + ) + args = parser.parse_args() + + results, passed, failed = run_all( + base_url=args.url, + fail_fast=not args.no_fail_fast, + ) + + if args.output: + output_data = [ + { + "name": r.name, + "passed": r.passed, + "error": r.error, + "turns": [ + { + "turn": t.turn, + "message_sent": t.message_sent, + "response_content": t.response_content, + "passed": t.passed, + "note": t.note, + } + for t in r.turns + ], + } + for r in results + ] + with open(args.output, "w", encoding="utf-8") as f: + json.dump(output_data, f, ensure_ascii=False, indent=2) + print(f"Results saved to {args.output}\n") + + sys.exit(0 if failed == 0 else 1) + + +if __name__ == "__main__": + main() diff --git a/tests/api_tool_eval/results-multi-intent.json b/tests/api_tool_eval/results-multi-intent.json new file mode 100644 index 00000000..37d96078 --- /dev/null +++ b/tests/api_tool_eval/results-multi-intent.json @@ -0,0 +1,114 @@ +[ + { + "name": "MI-9 — EN: address search + vehicle tax (cross-domain)", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "Can you find an address for me and also calculate my vehicle tax?", + "response_content": "", + "passed": null, + "note": "" + } + ] + }, + { + "name": "MI-10 — ET: elektrihinnad + riigipühad (eri valdkonnad)", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "Näita mulle Eesti elektrihindu ja too välja Eesti riigipühad", + "response_content": "", + "passed": null, + "note": "" + } + ] + }, + { + "name": "MI-11 — EN: parliament votings + initiatives (civic domain overlap)", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "Show me the parliament voting results and also list the citizen initiatives", + "response_content": "", + "passed": null, + "note": "" + } + ] + }, + { + "name": "MI-12 — ET: aadressi otsing + parlamendi osavõtustatistika", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "Otsi mulle ühte aadressi ja näita ka riigikogu liikmete osalusstatistikat", + "response_content": "", + "passed": null, + "note": "" + } + ] + }, + { + "name": "MI-13 — EN: vehicle tax + public holidays (unrelated domains)", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "What are the public holidays in Estonia and can you also calculate my vehicle tax?", + "response_content": "", + "passed": null, + "note": "" + } + ] + }, + { + "name": "MI-14 — ET: algatused + parlamendi hääletused (kodanikuühiskond)", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "Kuva mulle kõik kodanikualgatused ja näita ka riigikogu hääletustulemusi", + "response_content": "", + "passed": null, + "note": "" + } + ] + }, + { + "name": "MI-15 — EN: electricity prices + parliament participation stats + initiatives (3-way)", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "Get me the electricity market prices for Estonia, show the parliament member participation stats, and also list the citizen initiatives", + "response_content": "", + "passed": null, + "note": "" + } + ] + }, + { + "name": "MI-16 — ET: sõidukimaks + elektrihinnad (eri valdkonnad, eesti keel)", + "passed": null, + "error": "", + "turns": [ + { + "turn": 1, + "message_sent": "Arvuta mu sõiduki maks ja näita mulle ka Eesti elektri turuhinda", + "response_content": "", + "passed": null, + "note": "" + } + ] + } +] \ No newline at end of file diff --git a/tests/api_tool_eval/test-endpoints.json b/tests/api_tool_eval/test-endpoints.json index a463de89..d5c729e1 100644 --- a/tests/api_tool_eval/test-endpoints.json +++ b/tests/api_tool_eval/test-endpoints.json @@ -31,8 +31,8 @@ "url": "https://dashboard.elering.ee/api/nps/price", "method": "GET", "params": [ - { "name": "start", "type": "datetime", "required": true, "description": "Alguskuup\u00e4ev ja -aeg UTC formaadis (YYYY-MM-DDTHH:MM:SSZ). N\u00e4ide: 2025-01-01T00:00:00Z. Kui kasutaja annab ainult kuup\u00e4eva, kasuta alguseks T00:00:00Z." }, - { "name": "end", "type": "datetime", "required": true, "description": "L\u00f5ppkuup\u00e4ev ja -aeg UTC formaadis (YYYY-MM-DDTHH:MM:SSZ). N\u00e4ide: 2025-01-31T23:59:59Z. Kui kasutaja annab ainult kuup\u00e4eva, kasuta l\u00f5puks T23:59:59Z." } + { "name": "start", "type": "datetime", "required": true, "description": "Alguskuupäev ja -aeg UTC formaadis (YYYY-MM-DDTHH:MM:SSZ). Näide: 2025-01-01T00:00:00Z. Kui kasutaja annab ainult kuupäeva, kasuta alguseks T00:00:00Z. Ei tohi olla tulevane kuupäev — kui kasutaja sisestab tulevase kuupäeva, keeldu ja palu kehtiv mineviku- või tänane kuupäev." }, + { "name": "end", "type": "datetime", "required": true, "description": "Lõppkuupäev ja -aeg UTC formaadis (YYYY-MM-DDTHH:MM:SSZ). Näide: 2025-01-31T23:59:59Z. Kui kasutaja annab ainult kuupäeva, kasuta lõpuks T23:59:59Z. Ei tohi olla tulevane kuupäev — kui kasutaja sisestab tulevase kuupäeva, keeldu ja palu kehtiv mineviku- või tänane kuupäev." } ] }, { @@ -54,8 +54,8 @@ "url": "https://api.riigikogu.ee/api/votings", "method": "GET", "params": [ - { "name": "startDate", "type": "date", "required": true, "description": "Alguskuupäev (YYYY-MM-DD)" }, - { "name": "endDate", "type": "date", "required": true, "description": "Lõppkuupäev (YYYY-MM-DD)" }, + { "name": "startDate", "type": "date", "required": true, "description": "Alguskuupäev (YYYY-MM-DD). Ei tohi olla tulevane kuupäev — kui kasutaja sisestab tulevase kuupäeva, keeldu ja palu kehtiv mineviku- või tänane kuupäev." }, + { "name": "endDate", "type": "date", "required": true, "description": "Lõppkuupäev (YYYY-MM-DD). Ei tohi olla tulevane kuupäev — kui kasutaja sisestab tulevase kuupäeva, keeldu ja palu kehtiv mineviku- või tänane kuupäev." }, { "name": "lang", "type": "string", "required": false, "description": "Vastuse keel (ET, EN, RU). Täidetakse automaatselt kasutaja keele põhjal." } ] }, @@ -66,8 +66,8 @@ "url": "https://api.riigikogu.ee/api/statistics/participations/plenary", "method": "GET", "params": [ - { "name": "startDate", "type": "date", "required": true, "description": "Alguskuupäev (YYYY-MM-DD)" }, - { "name": "endDate", "type": "date", "required": true, "description": "Lõppkuupäev (YYYY-MM-DD)" }, + { "name": "startDate", "type": "date", "required": true, "description": "Alguskuupäev (YYYY-MM-DD). Ei tohi olla tulevane kuupäev — kui kasutaja sisestab tulevase kuupäeva, keeldu ja palu kehtiv mineviku- või tänane kuupäev." }, + { "name": "endDate", "type": "date", "required": true, "description": "Lõppkuupäev (YYYY-MM-DD). Ei tohi olla tulevane kuupäev — kui kasutaja sisestab tulevase kuupäeva, keeldu ja palu kehtiv mineviku- või tänane kuupäev." }, { "name": "lang", "type": "string", "required": false, "description": "Vastuse keel (ET, EN, RU). Täidetakse automaatselt kasutaja keele põhjal." } ] }, @@ -156,17 +156,5 @@ "url": "https://ilmmicroservice.envir.ee/api/forecasts", "method": "GET", "params": [] - }, - { - "endpointId": "9b4d0e99-7a8b-4c9d-b033-9d1c5f990016", - "name": "check_legal_eligibility", - "description": "Kontrolli isiku õiguslikku sobivust teatud toimingute tegemiseks vanuse ja kodakondsuse alusel.", - "url": "https://bcd831cf-cd1b-40b9-af86-b760cd6143e8.mock.pstmn.io/api/legal-check", - "method": "POST", - "params": [ - { "name": "action", "type": "string", "required": true, "description": "Toiming, mille jaoks sobivust kontrollitakse (nt start_business, vote, drive)" }, - { "name": "citizenship", "type": "string", "required": true, "description": "Kahetäheline ISO kodakondsuskood (nt EE, LV, FI)" }, - { "name": "age", "type": "integer", "required": true, "description": "Isiku vanus täisaastates" } - ] } ] \ No newline at end of file diff --git a/tests/data/test_dataset.json b/tests/data/test_dataset.json index 259ba59f..431ab1e6 100644 --- a/tests/data/test_dataset.json +++ b/tests/data/test_dataset.json @@ -1,183 +1,62 @@ [ { - "input": "How flexible will pensions become in 2021?", - "expected_output": "In 2021, pensions will become more flexible allowing people to choose the most suitable time for retirement, partially withdraw their pension, or stop pension payments if they wish, effectively creating their own personal pension plan.", - "retrieval_context": [ - "In 2021, the pension will become more flexible. People will be able to choose the most suitable time for their retirement, partially withdraw their pension or stop payment of their pension if they wish, in effect creating their own personal pension plan." - ], - "category": "pension_information", - "language": "en" - }, - { - "input": "Когда изменятся расчеты пенсионного возраста?", - "expected_output": "Начиная с 2027 года расчеты пенсионного возраста будут основываться на ожидаемой продолжительности жизни 65-летних людей. Пенсионная система таким образом будет соответствовать демографическим изменениям.", - "retrieval_context": [ - "Starting in 2027, retirement age calculations will be based on the life expectancy of 65-year-olds. The pension system will thus be in line with demographic developments." - ], - "category": "pension_information", - "language": "ru" - }, - { - "input": "Kui palju raha maksti peredele 2021. aastal?", - "expected_output": "2021. aastal maksti peredele kokku umbes 653 miljonit eurot toetusi, sealhulgas umbes 310 miljonit eurot peretoetuste eest ja 280 miljonit eurot lapsetoetuste eest.", - "retrieval_context": [ - "In 2021, a total of approximately 653 million euros in benefits were paid to families. Approximately 310 million euros for family benefits; Approximately 280 million euros for parental benefit." - ], - "category": "family_benefits", + "input": "Mida teha kui mobiil-ID kasutamisel kinnituskood ei ilmu mobiilile?", + "expected_output": "Veendu, et sinu telefon on mobiilvõrgu levialas ja mobiilne andmeside on sisse lülitatud. Lisaks tuleb kontrollida, kas mobiilsidevõrgus pole hetkel katkestusi. Mõnikord võib abiks olla telefoni taaskäivitamine ja/või selle võrguseadete lähtestamine.", + "category": "mobile_id_usage", "language": "et" }, { - "input": "Сколько семей получает поддержку для многодетных семей?", - "expected_output": "23,687 семей и 78,296 детей получают поддержку для многодетных семей, включая 117 семей с семью или более детьми.", - "retrieval_context": [ - "23,687 families and 78,296 children receive support for families with many children, including 117 families with seven or more children." - ], - "category": "family_benefits", - "language": "ru" - }, - { - "input": "How many single parents receive support?", - "expected_output": "8,804 parents and 10,222 children receive single parent support.", - "retrieval_context": [ - "8,804 parents and 1,0222 children receive single parent support." - ], - "category": "single_parent_support", - "language": "en" - }, - { - "input": "Какие уровни бедности среди семей с одним родителем?", - "expected_output": "Семьи с одним родителем (в основном матери) находятся в группе наивысшего риска бедности: 5,3% живут в абсолютной бедности и 27,3% в относительной бедности.", - "retrieval_context": [ - "Single-parent (mostly mother) families are at the highest risk of poverty, of whom 5.3% live in absolute poverty and 27.3% in relative poverty." - ], - "category": "single_parent_support", - "language": "ru" - }, - { - "input": "Millal saab piletit tagastada?", - "expected_output": "Pileti tagastamine on võimalik ainult juhul, kui reisi väljumiseni on jäänud vähemalt 60 minutit.", - "retrieval_context": [ - "Pileti tagastamine on võimalik ainult juhul, kui reisi väljumiseni on jäänud vähemalt 60 minutit." - ], - "category": "train_services", + "input": "Mida teha, kui minu telefon Mobiil-ID-ga varastatakse?", + "expected_output": "Kui sinu telefon Mobiil-ID-ga on varastatud, tuleks kohe teavitada oma mobiilsideoperaatorit, et peatada mobiil-ID sertifikaadid. See aitab kaitsta sind identiteedivarguse eest. Kui leiad oma telefoni uuesti, saad jätkata vana SIM-kaardi kasutamist, kuid mobiil-ID funktsionaalsuse taastamiseks on vajalik tellida uus Mobiil-ID SIM-kaart operaatorilt.", + "category": "digital_identity_security", "language": "et" }, { - "input": "За сколько минут до отправления можно вернуть билет на поезд?", - "expected_output": "Возврат билета возможен только в том случае, если до отправления поездки остается не менее 60 минут.", - "retrieval_context": [ - "Pileti tagastamine on võimalik ainult juhul, kui reisi väljumiseni on jäänud vähemalt 60 minutit." - ], - "category": "train_services", - "language": "ru" - }, - { - "input": "Where is ticket refunded without charge?", - "expected_output": "The ticket cost is refunded to the Elron travel card without service charge only if the refund request is submitted through the Elron homepage refund form.", - "retrieval_context": [ - "Pileti maksumus tagastatakse Elroni sõidukaardile teenustasuta ainult juhul, kui tagastussoov esitatakse Elroni kodulehe tagastusvormi kaudu." - ], - "category": "train_services", - "language": "en" - }, - { - "input": "Что сказала министр Кармен Йоллер о дезинформации в области здравоохранения?", - "expected_output": "Министр социальных дел Эстонии Кармен Йоллер заявила, что Европа должна действовать более совместно и скоординированно, чтобы остановить распространение дезинформации, связанной со здоровьем.", - "retrieval_context": [ - "Europe must act more jointly and in a more coordinated way to stop the spread of health-related misinformation, said Estonia's Minister of Social Affairs, Karmen Joller." - ], - "category": "health_cooperation", - "language": "ru" + "input": "Mis on eIDAS määrus?", + "expected_output": "eIDAS määrus (electronic IDentification, Authentication and trust Services) on Euroopa Liidus kehtiv e-identimise ja e-tehingute määrus, mille eesmärk on lihtsustada piiriülest elektrooniliste teenuste tarbimist ühtsustatud standardite ja tegutsemispõhimõtete kaudu. Määrus võeti vastu 23. juulil 2014 ja alates 1. juulist 2016 peavad Euroopa Liidu riigid tunnustama teineteise e-allkirju ning Eesti kodanike digiallkirju peab aktsepteerima ka teiste EL liikmesriikide avaliku sektori asutused.", + "category": "digital_identity", + "language": "et" }, { - "input": "What agreement did Estonia sign with Ukraine?", - "expected_output": "Estonian Minister of Social Affairs Karmen Joller and Ukrainian Minister of Health Viktor Liashko signed the next stage of a health cooperation agreement to reinforce health collaboration and support Ukraine's healthcare system recovery.", - "retrieval_context": [ - "Estonian Minister of Social Affairs Karmen Joller and Ukrainian Minister of Health Viktor Liashko today signed the next stage of a health cooperation agreement.", - "The aim of the agreement is to reinforce health collaboration, support Ukraine's healthcare system recovery." - ], - "category": "health_cooperation", + "input": "Why am I getting an error when trying to sign documents in DigiDoc4?", + "expected_output": "This error occurs when your computer's clock time differs from the validation confirmation time. The system needs synchronized time to verify digital signatures properly.To fix this: Windows: Click the Start menu (Windows logo) in the bottom left corner, type \"Control Panel\" and open it. Then go to \"Date and Time\" settings and ensure your date, time, and time zone are correct. macOS: Open Spotlight Search and type \"System Settings\". Select \"General\", then adjust settings in \"Language & Region\" and \"Date & Time\". Ubuntu: Click \"Show applications\" icon in the bottom left corner, search for \"Date & time\" and open it. Before making changes, disable automatic configuration first. If the error persists: Disconnect any active VPN connection, restart your device, and disable antivirus/security software temporarily. If problems continue, contact ID support through their web form.", + "category": "digital_identity", "language": "en" }, { - "input": "Какая цель соглашения о сотрудничестве с Украиной?", - "expected_output": "Цель соглашения - укрепить сотрудничество в области здравоохранения и поддержать восстановление системы здравоохранения Украины.", - "retrieval_context": [ - "The aim of the agreement is to reinforce health collaboration, support Ukraine's healthcare system recovery." - ], - "category": "health_cooperation", - "language": "ru" + "input": "Kuidas aktiveerida Mobiil-ID?", + "expected_output": "Mobiil-ID aktiveerimine toimub operaatorite iseteeninduses (Telia, Elisa, Tele2).", + "category": "digital_identity", + "language": "et" }, { - "input": "Kuidas tellida sõiduplaani muudatuste teateid?", - "expected_output": "Saad tellida sõiduplaanimuudatuste ja hilinemiste teated e-postile.", - "retrieval_context": [ - "Telli sõiduplaanimuudatuste ja hilinemiste teated e-postile." - ], - "category": "train_services", + "input": "Mis on Eesti sotsiaaluuring ja miks ma peaksin osalema?", + "expected_output": "Eesti sotsiaaluuring aitab hinnata leibkondade ja isikute sissetulekute jaotust, elamistingimusi ning sotsiaalset tõrjutust. Samasugust uuringut korraldatakse kõikides Euroopa Liidu riikides. Teie osalemine aitab kujundada Eesti sotsiaalpoliitikat ja arendada teenuseid. Miks see on oluline: Need andmed aitavad riigijuhtidel, teadlastel ja arvamusliidridel mõista inimeste toimetulekut, vaesust, ilmajäetust, sissetulekut, elamistingimusi ja arstiabi kättesaadavust. Iga uuringusse sattunud inimene esindab umbes sadat Eesti elanikku, seega on teie vastused olulised. Kuidas osaleda: Kui teid valitakse, saate teavituskirja e-postile (rahvastikuregistris märgitud aadressile) või postiga. Saate küsimustikule vastata veebis aadressil https://estat.stat.ee, logides sisse ID-kaardi, mobiil-ID või Smart-ID-ga, või osaleda telefoniintervjuus. Uuring koosneb kahest osast: leibkonnaküsimustik (vastab põhiisik) ja isikuküsimustik (võib vastata iga vähemalt 15-aastane leibkonnaliige). Osalemine on vabatahtlik, kuid Statistikaamet kaitseb kõigi vastajate andmeid riikliku statistika seaduse ja isikuandmete kaitse seaduse alusel. Iga leibkonda küsitletakse neli aastat järjest 12 kuu tagant.", + "category": "statistics", "language": "et" }, { - "input": "Как подписаться на уведомления об изменениях в расписании поездов?", - "expected_output": "Вы можете подписаться на уведомления об изменениях расписания и опозданиях по электронной почте.", - "retrieval_context": [ - "Telli sõiduplaanimuudatuste ja hilinemiste teated e-postile." - ], - "category": "train_services", - "language": "ru" + "input": "Kas ma saan kodus elektritöid ise teha või vajan spetsialisti?", + "expected_output": "Enamik elektritöid nõuab ohutuse huvides kvalifitseeritud spetsialisti. Valesti tehtud elektritöö võib põhjustada elektrilöögi ning tulekahjuohtu. Siiski saate lihtsamaid töid ise teha, kui teil on vastavad teadmised. Mida VÕITE ise teha: Vahetada lüliteid, pistikupesi, lambipesi ja kaitsmeid (kuid MITTE paigaldada uusi) Parandada ja asendada juhtmelüliteid, lambipesi, pikendusjuhtmeid ja juhtmepistikuid Milleks PEATE palkama spetsialisti: Uute elektripaigaldiste ehitamine Uute pistikupesade ja lülitite paigaldamine Kohtkindlate kodumasinate ühendamine ja lahti ühendamine Kaitsekontaktita (maandamata) pistikupesade vahetamine kaitsekontaktiga (maandatud) pistikupesade vastu Elektritöö ettevõtjad peavad olema esitanud majandustegevuse registrisse majandustegevuseteatise ning neil peab olema tööde eest vastutav kompetentne elektritöö juht.", + "category": "ttja", + "language": "et" }, { - "input": "What are the contact details of the Ministry of Social Affairs?", - "expected_output": "Ministry of Social Affairs is located at Suur-Ameerika 1, 10122 Tallinn, phone +372 626 9301, email [email protected]. Open Monday-Thursday 8.30-17.15 and Friday 8.30-16.00.", - "retrieval_context": [ - "Ministry of Social Affairs Suur-Ameerika 1, 10122 Tallinn +372 626 9301 [email protected] Open Mon -Thu 8.30-17.15 and Fri 8.30-16.00" - ], - "category": "contact_information", + "input": "What is an electrical installation audit and when do I need one?", + "expected_output": "An electrical installation audit checks whether your electrical system meets safety requirements and is safe to use. During the audit, the auditor visually assesses the installation's condition, reviews documentation and test/measurement results, and performs additional control measurements if necessary. When you need an audit: Before commissioning: Required before putting a new or renovated building's electrical installation into use Periodic audits: Regular checks at intervals depending on the installation type and age. While not mandatory for residential spaces (private houses, apartments, summer cottages), you should still periodically check them to ensure safety and functionality How to get an audit: Only contractors with appropriate accreditation can perform audits Results and documents are digitally formatted in TTJA's information system at https://jvis.ttja.ee, where they're always accessible to the electrical installation owner For residential electrical system checks, contact a competent electrical professional or auditor who will perform necessary operations and provide feedback on the system's condition and safety.", + "category": "ttja", "language": "en" }, { - "input": "Каковы контактные данные Министерства социальных дел?", - "expected_output": "Министерство социальных дел находится по адресу Суур-Амеэрика 1, 10122 Таллинн, телефон +372 626 9301, электронная почта [email protected]. Открыто понедельник-четверг 8.30-17.15 и пятница 8.30-16.00.", - "retrieval_context": [ - "Ministry of Social Affairs Suur-Ameerika 1, 10122 Tallinn +372 626 9301 [email protected] Open Mon -Thu 8.30-17.15 and Fri 8.30-16.00" - ], - "category": "contact_information", - "language": "ru" - }, - { - "input": "Сколько родителей-одиночек получают поддержку в Эстонии?", - "expected_output": "8,804 родителя и 10,222 ребенка получают поддержку для родителей-одиночек.", - "retrieval_context": [ - "8,804 parents and 1,0222 children receive single parent support." - ], - "category": "single_parent_support", - "language": "ru" - }, - { - "input": "Когда Министерство социальных дел начало искать решения для поддержки семей с одним родителем?", - "expected_output": "С января 2022 года Министерство социальных дел ищет решения для поддержки семей с одним родителем.", - "retrieval_context": [ - "Since January 2022, the Ministry of Social Affairs has been looking for solutions to support single-parent families." - ], - "category": "single_parent_support", - "language": "ru" - }, - { - "input": "Какова была численность населения Эстонии согласно прогнозам?", - "expected_output": "Согласно прогнозам, население Эстонии сократится с 1,31 миллиона до 1,11 миллиона к 2060 году. Количество людей в возрасте 18-63 лет уменьшится на 256,000 человек, или на 32%.", - "retrieval_context": [ - "According to forecasts, the population of Estonia will decrease from 1.31 million to 1.11 million by 2060. The number of people aged 18-63 will decrease by 256,000, or 32%." - ], - "category": "pension_information", - "language": "ru" + "input": "How long is the e-residency digi-ID valid for?", + "expected_output": "The e-residency digi-ID is valid for 5 years", + "category": "digital_identity", + "language": "en" }, { - "input": "Какая была новая инновационная программа стоимостью 12 миллионов евро?", - "expected_output": "На Фестивале социальных технологий была представлена новая инновационная программа стоимостью 12 миллионов евро, направленная на поддержку самостоятельной жизни пожилых людей и людей с ограниченными возможностями с помощью технологических решений.", - "retrieval_context": [ - "New €12 million innovation programme unveiled at Welfare Technology Festival aimed at supporting independent living for older adults and people with disabilities through technology-driven solutions." - ], - "category": "health_cooperation", + "input": "Предоставляет ли электронное резидентство эстонское гражданство или налоговое резидентство?", + "expected_output": "Нет, электронное резидентство не предоставляет эстонское гражданство или налоговое резидентство.", + "category": "digital_identity", "language": "ru" } ] \ No newline at end of file diff --git a/tests/deepeval_tests/api_tool_report_generator.py b/tests/deepeval_tests/api_tool_report_generator.py new file mode 100644 index 00000000..18686f2f --- /dev/null +++ b/tests/deepeval_tests/api_tool_report_generator.py @@ -0,0 +1,257 @@ +""" +Render the API Tool Calling test results JSON as a Markdown report. + +Reads ``api_tool_test_results.json`` (written by the +``_save_api_tool_results`` autouse fixture in +``tests/deepeval_tests/api_tool_tests.py``) and writes +``api_tool_test_report.md``. The workflow uploads the markdown as an artifact +and posts it as a PR comment. + +This is the API-tool counterpart of ``report_generator.py``; the two are +deliberately independent so changes to the RAG metrics report don't risk +breaking the API-tool report or vice versa. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +RESULTS_FILE = Path("api_tool_test_results.json") +REPORT_FILE = Path("api_tool_test_report.md") + + +# --------------------------------------------------------------------------- +# Loading +# --------------------------------------------------------------------------- + + +def load_results(path: Path = RESULTS_FILE) -> Dict[str, Any]: + """Load the JSON written by ApiToolResultCollector.save().""" + if not path.exists(): + return {"error": f"Results file not found: {path}"} + try: + with open(path, "r", encoding="utf-8") as f: + return json.load(f) + except json.JSONDecodeError as e: + return {"error": f"Results file is not valid JSON: {e}"} + + +# --------------------------------------------------------------------------- +# Aggregation helpers +# --------------------------------------------------------------------------- + + +def _by_type(scenarios: List[Dict[str, Any]]) -> Dict[str, List[Dict[str, Any]]]: + grouped: Dict[str, List[Dict[str, Any]]] = {} + for s in scenarios: + grouped.setdefault(s.get("type", "unknown"), []).append(s) + return grouped + + +def _pass_rate(scenarios: List[Dict[str, Any]]) -> Tuple[int, int, float]: + total = len(scenarios) + passed = sum(1 for s in scenarios if s.get("passed")) + rate = (passed / total * 100.0) if total else 0.0 + return passed, total, rate + + +def _status_emoji(scenario: Dict[str, Any]) -> str: + if scenario.get("error"): + return "💥" + if scenario.get("passed"): + return "✅" + return "❌" + + +def _fmt_tool(tc: Optional[Dict[str, Any]]) -> str: + if not tc: + return "_(none)_" + name = tc.get("name", "?") + params = tc.get("input_parameters") or {} + if not params: + return f"`{name}()`" + param_str = ", ".join(f"{k}={v!r}" for k, v in params.items()) + return f"`{name}({param_str})`" + + +def _score_str(score: Optional[float], threshold: float) -> str: + if score is None: + return f"— / {threshold}" + return f"{score:.2f} / {threshold}" + + +# --------------------------------------------------------------------------- +# Report sections +# --------------------------------------------------------------------------- + + +def render_header(results: Dict[str, Any]) -> str: + total = results.get("total_tests", 0) + passed = results.get("passed_tests", 0) + failed = results.get("failed_tests", 0) + errored = results.get("errored_tests", 0) + rate = (passed / total * 100.0) if total else 0.0 + started = results.get("test_start_time", "") + + return ( + "## API Tool Calling Evaluation Report\n\n" + f"_Issue #447 — DeepEval coverage for the API Tool Calling feature._\n\n" + f"**Started:** `{started}`\n\n" + "| Metric | Value |\n" + "|---|---|\n" + f"| Total scenarios | {total} |\n" + f"| Passed | {passed} |\n" + f"| Failed | {failed} |\n" + f"| Errored | {errored} |\n" + f"| Pass rate | **{rate:.1f}%** |\n\n" + ) + + +def render_by_type(results: Dict[str, Any]) -> str: + scenarios = results.get("scenarios", []) + grouped = _by_type(scenarios) + out = "### Results by scenario type\n\n" + out += "| Type | Metric | Passed | Total | Pass rate |\n" + out += "|---|---|---|---|---|\n" + type_metric = { + "strict": "ToolCorrectnessMetric (deterministic, threshold=1.0)", + "loose": "ArgumentCorrectnessMetric (LLM judge, threshold=0.7)", + "multi_intent": "ArgumentCorrectnessMetric or routing-only", + } + type_label = { + "strict": "Strict single-intent (S1, S2a, S2b, S3)", + "loose": "Loose single-intent (S4)", + "multi_intent": "Multi-intent (MI-1..MI-8)", + } + for stype in ("strict", "loose", "multi_intent"): + items = grouped.get(stype, []) + if not items: + continue + p, t, r = _pass_rate(items) + out += ( + f"| {type_label.get(stype, stype)} | " + f"{type_metric.get(stype, '—')} | " + f"{p} | {t} | {r:.1f}% |\n" + ) + return out + "\n" + + +def render_scenario_table(results: Dict[str, Any]) -> str: + scenarios = results.get("scenarios", []) + if not scenarios: + return "### Detailed results\n\n_No scenarios recorded._\n\n" + + out = "### Detailed results\n\n" + out += "| Status | Scenario | Metric | Score | Expected → Actual |\n" + out += "|---|---|---|---|---|\n" + for s in scenarios: + status = _status_emoji(s) + sid = s.get("id", "?") + metric = s.get("metric", "—") + score_cell = _score_str(s.get("score"), s.get("threshold", 0.0)) + expected = _fmt_tool(s.get("expected_tool")) + actual = _fmt_tool(s.get("actual_tool")) + # For multi-intent / loose, expected_tool is None — show "→ {actual}" only + if s.get("expected_tool") is None: + arrow = actual + else: + arrow = f"{expected} → {actual}" + out += f"| {status} | `{sid}` | {metric} | {score_cell} | {arrow} |\n" + return out + "\n" + + +def render_failures(results: Dict[str, Any]) -> str: + failures = [ + s + for s in results.get("scenarios", []) + if not s.get("passed") and not s.get("error") + ] + if not failures: + return "" + out = "### Failed scenarios\n\n" + for s in failures: + sid = s.get("id", "?") + reason = s.get("reason") or "_(no reason captured)_" + preview = (s.get("final_response_preview") or "").replace("\n", " ") + if len(preview) > 200: + preview = preview[:200] + "…" + out += f"#### ❌ `{sid}` ({s.get('metric', '—')})\n\n" + out += f"- **Score:** {_score_str(s.get('score'), s.get('threshold', 0.0))}\n" + if s.get("expected_tool"): + out += f"- **Expected:** {_fmt_tool(s['expected_tool'])}\n" + out += f"- **Actual:** {_fmt_tool(s.get('actual_tool'))}\n" + out += f"- **Reason:** {reason}\n" + if preview: + out += f"- **Response preview:** `{preview}`\n" + out += "\n" + return out + + +def render_errors(results: Dict[str, Any]) -> str: + errored = [s for s in results.get("scenarios", []) if s.get("error")] + if not errored: + return "" + out = "### Errored scenarios (test harness errors, not assertion failures)\n\n" + for s in errored: + sid = s.get("id", "?") + err = s.get("error") or "_(no error captured)_" + out += f"- 💥 `{sid}` — {err}\n" + return out + "\n" + + +def render_methodology() -> str: + return ( + "### Methodology\n\n" + "- **Strict scenarios** (issue specifies the expected endpoint and " + "params) are scored with `ToolCorrectnessMetric`, " + "`evaluation_params=[ToolCallParams.INPUT_PARAMETERS]`, threshold 1.0. " + "Tool name must match exactly and every expected parameter must be " + "present with the expected value (extra parameters allowed).\n" + "- **Loose scenarios** (issue describes the flow but not the expected " + "resolution) are scored with `ArgumentCorrectnessMetric` (LLM-as-judge, " + "threshold 0.7).\n" + "- **Multi-intent scenarios** use `ArgumentCorrectnessMetric` when the " + "agent resolves a tool call; if the agent asks a clarifying question " + "instead, only routing-to-ATC is verified (rules out silent fall-through " + "to RAG/OOD).\n" + "- API tool endpoints are seeded into the testcontainers-backed Qdrant " + "from `tests/api_tool_eval/test-endpoints.json` via the " + "`api_tool_endpoints_indexed` fixture in `conftest.py`.\n\n" + ) + + +def render_report(results: Dict[str, Any]) -> str: + if results.get("error"): + return ( + f"## API Tool Calling Evaluation Report\n\n**ERROR:** {results['error']}\n" + ) + return ( + render_header(results) + + render_by_type(results) + + render_scenario_table(results) + + render_failures(results) + + render_errors(results) + + render_methodology() + ) + + +def main() -> int: + results = load_results() + markdown = render_report(results) + REPORT_FILE.write_text(markdown, encoding="utf-8") + print(f"Wrote {REPORT_FILE} ({len(markdown)} chars)") + if results.get("error"): + print(f" WARNING: {results['error']}", file=sys.stderr) + else: + print( + f" {results.get('passed_tests', 0)}/{results.get('total_tests', 0)} " + "scenarios passed" + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/deepeval_tests/api_tool_tests.py b/tests/deepeval_tests/api_tool_tests.py new file mode 100644 index 00000000..1edaae06 --- /dev/null +++ b/tests/deepeval_tests/api_tool_tests.py @@ -0,0 +1,779 @@ +""" +DeepEval tests for the API Tool Calling feature (issue #447). + +Covers the single-intent (Scenarios 1-4) and multi-intent (MI-1..MI-8) scenarios +listed in the issue against the running orchestration service. Each scenario +walks the ``/orchestrate`` endpoint turn-by-turn via the testcontainers-backed +``orchestration_client`` fixture, extracts the final agentic-loop tool call, +and scores it with a DeepEval agentic metric: + +* **Strict single-intent (S1, S2a, S2b, S3)** — the issue specifies an + "Expected endpoint" + URL/params per scenario. Scored with + ``ToolCorrectnessMetric`` (deterministic name + input-parameter comparison, + ``threshold=1.0``). + +* **Loose single-intent (S4)** — the issue documents the 5-turn flow but does + not specify an expected resolution. Scored with + ``ArgumentCorrectnessMetric`` (LLM-as-judge over the conversation input and + the resolved tool call). + +* **Multi-intent (MI-1..MI-8)** — issue lists only queries (EN+ET). If the + system resolves a tool call, scored with ``ArgumentCorrectnessMetric``; if + it asks a clarifying question instead, the test only verifies routing to + ATC (i.e. non-empty reply, not silent RAG/OOD). + +Endpoints are matched by ``name`` — the UUIDs in the issue do happen to match +those in ``tests/api_tool_eval/test-endpoints.json``, but name matching is the +stable contract. + +Depends on: +* ``orchestration_client`` — provides the testcontainers-mapped base URL. +* ``api_tool_endpoints_indexed`` — seeds the API tool fixture into Qdrant's + ``api_tool_collection`` so the agentic loop can find the endpoints. +""" + +import datetime +import json +import time +import uuid +from pathlib import Path +from typing import Any, Dict, List, Optional + +import pytest +import requests +from deepeval.metrics import ArgumentCorrectnessMetric, ToolCorrectnessMetric +from deepeval.test_case import LLMTestCase, ToolCall, ToolCallParams + +REQUEST_TIMEOUT = 60 +ENVIRONMENT = "development" +AUTHOR_ID = "api-tool-deepeval" + +# Where the result-collector writes the per-scenario record consumed by +# tests/deepeval_tests/api_tool_report_generator.py to render the PR +# comment / artifact markdown. +RESULTS_FILE = Path("api_tool_test_results.json") + +# Strict scenarios assert the exact expected tool was called with the exact +# expected params (extras allowed — see ToolCorrectnessMetric docs on +# should_exact_match). Threshold 1.0 because the deterministic comparison +# scores fractionally over expected_tools, and we want every expected param +# present and correct. +STRICT_TOOL_THRESHOLD = 1.0 + +# Loose scenarios are graded by an LLM judge — 0.7 matches the threshold used +# for the RAG metrics in standard_tests.py. +JUDGE_THRESHOLD = 0.7 + + +# --------------------------------------------------------------------------- +# HTTP helpers (mirror tests/api_tool_eval/integration_test_*.py) +# --------------------------------------------------------------------------- + + +def _make_chat_id(label: str) -> str: + return f"deepeval-api-tool-{label}-{uuid.uuid4().hex[:8]}" + + +def _send_turn( + base_url: str, + chat_id: str, + message: str, + history: List[Dict[str, str]], +) -> Dict[str, Any]: + payload: Dict[str, Any] = { + "chatId": chat_id, + "message": message, + "authorId": AUTHOR_ID, + "conversationHistory": history, + "url": "deepeval-test", + "environment": ENVIRONMENT, + } + resp = requests.post( + f"{base_url}/orchestrate", + json=payload, + timeout=REQUEST_TIMEOUT, + ) + resp.raise_for_status() + return resp.json() + + +def _append_history( + history: List[Dict[str, str]], + user_message: str, + bot_response: str, +) -> List[Dict[str, str]]: + ts = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()) + return history + [ + {"authorRole": "user", "message": user_message, "timestamp": ts}, + {"authorRole": "bot", "message": bot_response, "timestamp": ts}, + ] + + +def _parse_completed(content: str) -> Optional[Dict[str, Any]]: + """Return the parsed JSON if content is a completed agentic-loop payload, + else None. A completed payload carries both ``endpoint`` and + ``collected_params`` keys.""" + try: + data = json.loads(content) + except (json.JSONDecodeError, TypeError): + return None + if "collected_params" in data and "endpoint" in data: + return data + return None + + +def _to_tool_call(content: str) -> Optional[ToolCall]: + """Build a DeepEval ToolCall from a completed agentic-loop payload, or + None if the response wasn't a completed JSON.""" + data = _parse_completed(content) + if data is None: + return None + ep = data.get("endpoint", {}) + name = ep.get("name", "") if isinstance(ep, dict) else str(ep) + params = data.get("collected_params", {}) or {} + return ToolCall(name=name, input_parameters=params) + + +def _conversation_text(turns: List[Dict[str, Any]]) -> str: + """Flatten the user turns into a single string for LLMTestCase.input. + + DeepEval's agentic single-turn metrics take a single ``input`` string; + this approximation gives the LLM judge the full conversational context. + """ + return "\n".join(f"USER: {t['user']}" for t in turns) + + +def _walk_turns(base_url: str, label: str, turns: List[Dict[str, Any]]) -> str: + """POST each turn in sequence, maintaining a stable chatId + history. + Returns the final bot response content.""" + chat_id = _make_chat_id(label) + history: List[Dict[str, str]] = [] + final_content = "" + for turn in turns: + resp = _send_turn(base_url, chat_id, turn["user"], history) + final_content = resp.get("content", "") + history = _append_history(history, turn["user"], final_content) + return final_content + + +def _tool_call_to_dict(tc: Optional[ToolCall]) -> Optional[Dict[str, Any]]: + """Serialize a ToolCall for the results JSON (None-tolerant).""" + if tc is None: + return None + return {"name": tc.name, "input_parameters": dict(tc.input_parameters or {})} + + +# --------------------------------------------------------------------------- +# Result collector — every test pushes one record, autouse fixture flushes to +# disk at session end. tests/deepeval_tests/api_tool_report_generator.py +# reads the JSON and renders the markdown report consumed by the workflow. +# --------------------------------------------------------------------------- + + +class ApiToolResultCollector: + """Accumulates per-scenario results from the API tool tests.""" + + def __init__(self) -> None: + self.results: Dict[str, Any] = { + "total_tests": 0, + "passed_tests": 0, + "failed_tests": 0, + "errored_tests": 0, + "test_start_time": datetime.datetime.now().isoformat(), + "scenarios": [], + } + + def add( + self, + scenario_id: str, + scenario_type: str, + metric_name: str, + threshold: float, + score: Optional[float], + passed: bool, + reason: str = "", + error: str = "", + expected_tool: Optional[Dict[str, Any]] = None, + actual_tool: Optional[Dict[str, Any]] = None, + final_response_preview: str = "", + extra: Optional[Dict[str, Any]] = None, + ) -> None: + self.results["total_tests"] += 1 + if error: + self.results["errored_tests"] += 1 + elif passed: + self.results["passed_tests"] += 1 + else: + self.results["failed_tests"] += 1 + self.results["scenarios"].append( + { + "id": scenario_id, + "type": scenario_type, + "metric": metric_name, + "threshold": threshold, + "score": score, + "passed": passed, + "reason": reason, + "error": error, + "expected_tool": expected_tool, + "actual_tool": actual_tool, + "final_response_preview": final_response_preview, + "extra": extra or {}, + } + ) + + def save(self, path: Path = RESULTS_FILE) -> None: + with open(path, "w", encoding="utf-8") as f: + json.dump(self.results, f, indent=2, default=str, ensure_ascii=False) + print( + f"Saved API tool results to {path}: " + f"{self.results['passed_tests']}/{self.results['total_tests']} passed, " + f"{self.results['failed_tests']} failed, " + f"{self.results['errored_tests']} errored" + ) + + +_collector = ApiToolResultCollector() + + +@pytest.fixture(scope="session", autouse=True) +def _save_api_tool_results(): + """Flush collected results to RESULTS_FILE at end of session, even on + failure — mirrors save_results_fixture in standard_tests.py.""" + yield + _collector.save() + + +# --------------------------------------------------------------------------- +# Strict single-intent scenarios — issue specifies expected endpoint + params +# (S1, S2a, S2b, S3). Scored with ToolCorrectnessMetric. +# --------------------------------------------------------------------------- + + +STRICT_SINGLE_INTENT_SCENARIOS: List[Dict[str, Any]] = [ + # Scenario 1 — Normal Workflow (citizen initiative details) + { + "id": "S1-citizen-initiative-EN", + "label": "s1-en", + "turns": [ + {"user": "Can I see the details of a citizen initiative?"}, + {"user": "1790"}, + ], + "expected_tool": ToolCall( + name="get_initiative_details", + input_parameters={"id": "1790"}, + ), + }, + { + "id": "S1-citizen-initiative-ET", + "label": "s1-et", + "turns": [ + {"user": "Kas ma saan kodanikualgatuse üksikasju vaadata?"}, + {"user": "1790"}, + ], + "expected_tool": ToolCall( + name="get_initiative_details", + input_parameters={"id": "1790"}, + ), + }, + # Scenario 2a — Public Holidays (date range correction) + { + "id": "S2a-public-holidays-date-correction-EN", + "label": "s2a-en", + "turns": [ + {"user": "What are the public holidays in Estonia?"}, + {"user": "From 2026-01-01"}, + { + "user": ( + "My mistake — the correct period is April 1, 2026 " + "through December 31, 2026." + ) + }, + ], + "expected_tool": ToolCall( + name="get_public_holidays", + input_parameters={ + "countryIsoCode": "EE", + "validFrom": "2026-04-01", + "validTo": "2026-12-31", + }, + ), + }, + { + "id": "S2a-public-holidays-date-correction-ET", + "label": "s2a-et", + "turns": [ + {"user": "Millised on riigipühad Eestis?"}, + {"user": "Alates 1. jaanuarist 2026"}, + { + "user": ( + "Minu viga, õige periood on 01.04.2026 kuni 31.12.2026. " + "Tegelikult tahan 2026-04-01 kuni 2026-12-31." + ) + }, + ], + "expected_tool": ToolCall( + name="get_public_holidays", + input_parameters={ + "countryIsoCode": "EE", + "validFrom": "2026-04-01", + "validTo": "2026-12-31", + }, + ), + }, + # Scenario 2b — Parliament Votings (date range correction) + { + "id": "S2b-parliament-votings-date-correction-EN", + "label": "s2b-en", + "turns": [ + {"user": "What votes took place in the Estonian parliament?"}, + {"user": "2026-04-05"}, + { + "user": ( + "My mistake — the correct period is April 6, 2026 " + "through April 7, 2026." + ) + }, + ], + "expected_tool": ToolCall( + name="get_parliament_votings", + input_parameters={ + "startDate": "2026-04-06", + "endDate": "2026-04-07", + }, + ), + }, + { + "id": "S2b-parliament-votings-date-correction-DE", + "label": "s2b-de", + "turns": [ + {"user": ("Welche Abstimmungen fanden im estnischen Parlament statt?")}, + {"user": "2026-04-05"}, + { + "user": ( + "Mein Fehler — der richtige Zeitraum ist vom " + "6. April 2026 bis zum 7. April 2026." + ) + }, + ], + "expected_tool": ToolCall( + name="get_parliament_votings", + input_parameters={ + "startDate": "2026-04-06", + "endDate": "2026-04-07", + }, + ), + }, + # Scenario 3 — Intent Switch (electricity prices -> address search) + { + "id": "S3-intent-switch-EN", + "label": "s3-en", + "turns": [ + {"user": "Show last week's electricity prices in Estonia."}, + { + "user": ( + "Wait — could you check the following location instead: " + "Viru tn 4, Tallinn?" + ) + }, + ], + "expected_tool": ToolCall( + name="search_address", + input_parameters={"address": "Viru tn 4, Tallinn"}, + ), + }, + { + "id": "S3-intent-switch-ET", + "label": "s3-et", + "turns": [ + {"user": "Näita eelmise nädala elektrienergia hindu Eestis."}, + { + "user": ( + "Oota, kas saaksid hoopis järgmist asukohta kontrollida: " + "Viru tn 4, Tallinn?" + ) + }, + ], + "expected_tool": ToolCall( + name="search_address", + input_parameters={"address": "Viru tn 4, Tallinn"}, + ), + }, +] + + +@pytest.mark.parametrize( + "scenario", + STRICT_SINGLE_INTENT_SCENARIOS, + ids=[s["id"] for s in STRICT_SINGLE_INTENT_SCENARIOS], +) +def test_api_tool_strict_single_intent( + scenario: Dict[str, Any], + orchestration_client: Any, + api_tool_endpoints_indexed: None, +) -> None: + """Scenarios 1, 2a, 2b, 3 — issue specifies the expected tool call. + + Scored deterministically with ``ToolCorrectnessMetric``: + * the resolved tool's name must equal the expected name, and + * every expected input parameter must be present and equal in the + ``collected_params`` (extras allowed). + """ + del api_tool_endpoints_indexed # consumed for its setup side effect only + + expected_tool: ToolCall = scenario["expected_tool"] + final_content = "" + actual_tool: Optional[ToolCall] = None + score: Optional[float] = None + reason = "" + passed = False + error = "" + + try: + final_content = _walk_turns( + orchestration_client.base_url, scenario["label"], scenario["turns"] + ) + actual_tool = _to_tool_call(final_content) + test_case = LLMTestCase( + input=_conversation_text(scenario["turns"]), + actual_output=final_content, + tools_called=[actual_tool] if actual_tool is not None else [], + expected_tools=[expected_tool], + ) + metric = ToolCorrectnessMetric( + threshold=STRICT_TOOL_THRESHOLD, + evaluation_params=[ToolCallParams.INPUT_PARAMETERS], + ) + metric.measure(test_case) + score = metric.score + reason = metric.reason or "" + passed = score is not None and score >= STRICT_TOOL_THRESHOLD + assert passed, ( + f"[{scenario['id']}] tool correctness {score} < " + f"{STRICT_TOOL_THRESHOLD}: {reason}\n" + f"Expected tool: {expected_tool.name}({expected_tool.input_parameters})\n" + f"Final response: {final_content[:300]}" + ) + except AssertionError: + raise + except Exception as e: + error = f"{type(e).__name__}: {e}" + raise + finally: + _collector.add( + scenario_id=scenario["id"], + scenario_type="strict", + metric_name="ToolCorrectnessMetric", + threshold=STRICT_TOOL_THRESHOLD, + score=score, + passed=passed, + reason=reason, + error=error, + expected_tool=_tool_call_to_dict(expected_tool), + actual_tool=_tool_call_to_dict(actual_tool), + final_response_preview=final_content[:300], + ) + + +# --------------------------------------------------------------------------- +# Loose single-intent scenarios — issue documents the flow but does not +# specify an "Expected endpoint" (S4, both languages). Scored with +# ArgumentCorrectnessMetric (LLM judges arg correctness vs. the +# conversation input). +# --------------------------------------------------------------------------- + + +LOOSE_SINGLE_INTENT_SCENARIOS: List[Dict[str, Any]] = [ + { + "id": "S4-parliament-attendance-multi-turn-EN", + "label": "s4-en", + "turns": [ + { + "user": ( + "Can you show me the parliament attendance of former " + "Finance Minister Martin Helme?" + ) + }, + {"user": "Can you just check with what you have?"}, + {"user": "2026-04-01"}, + {"user": "Yes"}, + {"user": "2026-04-20"}, + ], + }, + { + "id": "S4-parliament-attendance-multi-turn-ET", + "label": "s4-et", + "turns": [ + { + "user": ( + "Kas saaksite mulle näidata endise rahandusministri " + "Martin Helme parlamendi kohaloleku andmeid?" + ) + }, + {"user": "Kas saaksite lihtsalt oma andmetest järele vaadata?"}, + {"user": "2026-04-01"}, + {"user": "Jah"}, + {"user": "2026-04-20"}, + ], + }, +] + + +@pytest.mark.parametrize( + "scenario", + LOOSE_SINGLE_INTENT_SCENARIOS, + ids=[s["id"] for s in LOOSE_SINGLE_INTENT_SCENARIOS], +) +def test_api_tool_loose_single_intent( + scenario: Dict[str, Any], + orchestration_client: Any, + api_tool_endpoints_indexed: None, +) -> None: + """Scenario 4 (EN/ET) — issue gives no expected resolution. + + Asserts: + 1. The final turn produced a completed JSON tool call (i.e. the agentic + loop actually resolved, didn't fall through to RAG). + 2. The LLM judge ``ArgumentCorrectnessMetric`` is satisfied that the + chosen tool's arguments fit the conversation input. + """ + del api_tool_endpoints_indexed # consumed for its setup side effect only + + final_content = "" + actual_tool: Optional[ToolCall] = None + score: Optional[float] = None + reason = "" + passed = False + error = "" + + try: + final_content = _walk_turns( + orchestration_client.base_url, scenario["label"], scenario["turns"] + ) + actual_tool = _to_tool_call(final_content) + + assert actual_tool is not None, ( + f"[{scenario['id']}] expected a completed JSON tool call on the " + f"final turn but got: {final_content[:300]}" + ) + + test_case = LLMTestCase( + input=_conversation_text(scenario["turns"]), + actual_output=final_content, + tools_called=[actual_tool], + ) + metric = ArgumentCorrectnessMetric(threshold=JUDGE_THRESHOLD) + metric.measure(test_case) + score = metric.score + reason = metric.reason or "" + passed = score is not None and score >= JUDGE_THRESHOLD + assert passed, ( + f"[{scenario['id']}] argument correctness {score} < " + f"{JUDGE_THRESHOLD}: {reason}\n" + f"Resolved tool: {actual_tool.name}({actual_tool.input_parameters})" + ) + except AssertionError: + raise + except Exception as e: + error = f"{type(e).__name__}: {e}" + raise + finally: + _collector.add( + scenario_id=scenario["id"], + scenario_type="loose", + metric_name="ArgumentCorrectnessMetric", + threshold=JUDGE_THRESHOLD, + score=score, + passed=passed, + reason=reason, + error=error, + actual_tool=_tool_call_to_dict(actual_tool), + final_response_preview=final_content[:300], + ) + + +# --------------------------------------------------------------------------- +# Multi-Intent scenarios (issue #447, MI-1..MI-8) +# +# Issue lists 8 queries (EN + ET) with no expected resolution. Under Phase 1 +# the orchestrator decomposes the query and falls back to a single endpoint +# (see tests/api_tool_eval/integration_test_multi_intent.py docstring). We: +# +# * If the system resolves a tool call → score with ArgumentCorrectnessMetric +# (LLM judges whether the chosen args fit the multi-intent query). +# * If the system asks a clarifying question → accept it (still ATC-routed), +# only assert the reply is non-empty (rules out a silent RAG/OOD fallthrough). +# --------------------------------------------------------------------------- + + +MULTI_INTENT_SCENARIOS: List[Dict[str, Any]] = [ + { + "id": "MI-1-address-and-vehicle-tax", + "label": "mi1", + "query_en": ( + "Can you find an address for me and also calculate my vehicle " + "tax? (Address: Viru tn 4, Tallinn / Plate: 123ABC / Year: 2026)" + ), + "query_et": ( + "Kas saaksite mulle aadressi leida ja arvutada ka mu sõiduki " + "maksu? (Aadress: Viru tn 4, Tallinn / Registreerimismärk: " + "123ABC / Aasta: 2026)" + ), + }, + { + "id": "MI-2-address-and-initiative-details", + "label": "mi2", + "query_en": ( + "I need to find an address and also check details of an " + "initiative. (Address: Viru tn 4, Tallinn / Initiative ID: 1790)" + ), + "query_et": ( + "Mul on vaja leida aadress ja vaadata ka kodanikualgatuse " + "üksikasju. (Aadress: Viru tn 4, Tallinn / Algatuse ID: 1790)" + ), + }, + { + "id": "MI-3-electricity-and-public-holidays", + "label": "mi3", + "query_en": ( + "Show me electricity prices in Estonia and list Estonia's public holidays." + ), + "query_et": ("Näita mulle Eesti elektrihindu ja too välja Eesti riigipühad."), + }, + { + "id": "MI-4-parliament-votings-and-initiatives", + "label": "mi4", + "query_en": ( + "Show me the parliament voting results and also list the " + "citizen initiatives." + ), + "query_et": ( + "Näita mulle parlamendi hääletustulemusi ja too välja ka kodanikualgatused." + ), + }, + { + "id": "MI-5-address-and-parliament-participation", + "label": "mi5", + "query_en": ( + "Could you find me an address and also show the attendance " + "statistics of Riigikogu members?" + ), + "query_et": ( + "Kas saaksite mulle leida aadressi ja näidata ka Riigikogu " + "liikmete osalusstatistikat?" + ), + }, + { + "id": "MI-6-public-holidays-and-vehicle-tax", + "label": "mi6", + "query_en": ( + "What are the public holidays in Estonia and can you also " + "calculate my vehicle tax?" + ), + "query_et": ( + "Millised on riigipühad Eestis ja kas saate arvutada ka mu sõiduki maksu?" + ), + }, + { + "id": "MI-7-initiatives-and-parliament-votings", + "label": "mi7", + "query_en": ( + "Show me all citizen initiatives and also display the parliament " + "voting results." + ), + "query_et": ( + "Kuva mulle kõik kodanikualgatused ja näita ka riigikogu hääletustulemusi." + ), + }, + { + "id": "MI-8-vehicle-tax-and-electricity", + "label": "mi8", + "query_en": ( + "Calculate my vehicle tax and also show me the electricity " + "market price in Estonia." + ), + "query_et": ( + "Arvuta mu sõiduki maks ja näita mulle ka Eesti elektri turuhinda." + ), + }, +] + + +@pytest.mark.parametrize( + "scenario", + MULTI_INTENT_SCENARIOS, + ids=[s["id"] for s in MULTI_INTENT_SCENARIOS], +) +@pytest.mark.parametrize("lang", ["en", "et"]) +def test_api_tool_multi_intent( + scenario: Dict[str, Any], + lang: str, + orchestration_client: Any, + api_tool_endpoints_indexed: None, +) -> None: + del api_tool_endpoints_indexed # consumed for its setup side effect only + + query = scenario[f"query_{lang}"] + content = "" + actual_tool: Optional[ToolCall] = None + score: Optional[float] = None + reason = "" + passed = False + error = "" + outcome = "unknown" # "tool_call" | "clarifying_question" + + try: + base_url = orchestration_client.base_url + chat_id = _make_chat_id(f"{scenario['label']}-{lang}") + resp = _send_turn(base_url, chat_id, query, []) + content = resp.get("content", "") + actual_tool = _to_tool_call(content) + + if actual_tool is not None: + outcome = "tool_call" + test_case = LLMTestCase( + input=query, + actual_output=content, + tools_called=[actual_tool], + ) + metric = ArgumentCorrectnessMetric(threshold=JUDGE_THRESHOLD) + metric.measure(test_case) + score = metric.score + reason = metric.reason or "" + passed = score is not None and score >= JUDGE_THRESHOLD + assert passed, ( + f"[{scenario['id']} {lang}] argument correctness {score} " + f"< {JUDGE_THRESHOLD}: {reason}\n" + f"Resolved tool: {actual_tool.name}({actual_tool.input_parameters})" + ) + else: + outcome = "clarifying_question" + # No tool call yet — must at least be a non-empty clarifying reply + # (rules out silent failure / RAG fallthrough returning no content). + passed = bool(content.strip()) + assert passed, ( + f"[{scenario['id']} {lang}] empty response — multi-intent query " + f"failed routing entirely. Got: {content!r}" + ) + reason = "Resolved to a clarifying question (no tool call yet)" + except AssertionError: + raise + except Exception as e: + error = f"{type(e).__name__}: {e}" + raise + finally: + _collector.add( + scenario_id=f"{scenario['id']}-{lang}", + scenario_type="multi_intent", + metric_name=( + "ArgumentCorrectnessMetric" if outcome == "tool_call" else "RoutingOnly" + ), + threshold=JUDGE_THRESHOLD if outcome == "tool_call" else 0.0, + score=score, + passed=passed, + reason=reason, + error=error, + actual_tool=_tool_call_to_dict(actual_tool), + final_response_preview=content[:300], + extra={"language": lang, "outcome": outcome, "query": query}, + ) diff --git a/tests/deepeval_tests/conftest.py b/tests/deepeval_tests/conftest.py new file mode 100644 index 00000000..37624c82 --- /dev/null +++ b/tests/deepeval_tests/conftest.py @@ -0,0 +1,901 @@ +import os +import time +import subprocess +from pathlib import Path +from typing import Dict, Any, Optional, Generator + +import pytest +import hvac +import requests +from loguru import logger +from testcontainers.compose import DockerCompose # type: ignore +from azure.storage.blob import BlobServiceClient + + +# ===================== Azure Blob Storage Helper ===================== + + +def download_embeddings_from_azure( + connection_string: str, container_name: str, blob_name: str, local_path: Path +) -> None: + """ + Download pre-computed embeddings from Azure Blob Storage. + + Args: + connection_string: Azure Storage connection string + container_name: Name of the blob container + blob_name: Name of the blob to download + local_path: Local path to save the downloaded file + """ + logger.info("Downloading embeddings from Azure Blob Storage...") + logger.info(f" Container: {container_name}") + logger.info(f" Blob: {blob_name}") + logger.info(f" Local path: {local_path}") + + try: + # Create BlobServiceClient + blob_service_client = BlobServiceClient.from_connection_string( + connection_string + ) + + # Get blob client + blob_client = blob_service_client.get_blob_client( + container=container_name, blob=blob_name + ) + + # Ensure parent directory exists + local_path.parent.mkdir(parents=True, exist_ok=True) + + # Download the blob + with open(local_path, "wb") as download_file: + download_stream = blob_client.download_blob() + download_file.write(download_stream.readall()) + + file_size_kb = local_path.stat().st_size / 1024 + logger.info(f"✓ Downloaded embeddings successfully ({file_size_kb:.2f} KB)") + + except Exception as e: + logger.error(f"Failed to download embeddings from Azure: {e}") + raise + + +# ===================== VaultAgentClient ===================== + + +class VaultAgentClient: + """Client for interacting with Vault using a token written by Vault Agent""" + + def __init__( + self, + vault_url: str, + token_path: Path = Path("test-vault/agent-out/token"), + mount_point: str = "secret", + timeout: int = 10, + ): + self.vault_url = vault_url + self.token_path = token_path + self.mount_point = mount_point + + self.client = hvac.Client(url=self.vault_url, timeout=timeout) + self._load_token() + + def _load_token(self) -> None: + """Load token from file written by Vault Agent""" + if not self.token_path.exists(): + raise FileNotFoundError(f"Vault token file missing: {self.token_path}") + token = self.token_path.read_text().strip() + if not token: + raise ValueError("Vault token file is empty") + self.client.token = token + + def is_authenticated(self) -> bool: + """Check if the current token is valid""" + try: + return self.client.is_authenticated() + except Exception as e: + logger.warning(f"Vault token is not valid: {e}") + return False + + def is_vault_available(self) -> bool: + """Check if Vault is initialized and unsealed""" + try: + status = self.client.sys.read_health_status(method="GET") + return ( + isinstance(status, dict) + and status.get("initialized", False) + and not status.get("sealed", True) + ) + except Exception as e: + logger.warning(f"Vault availability check failed: {e}") + return False + + def get_secret(self, path: str) -> dict: + """Read a secret from Vault KV v2""" + try: + result = self.client.secrets.kv.v2.read_secret_version( + path=path, mount_point=self.mount_point + ) + return result["data"]["data"] + except Exception as e: + logger.error(f"Failed to read Vault secret at {path}: {e}") + raise + + +# ===================== RAGStackTestContainers ===================== + + +class RAGStackTestContainers: + """Manages test containers for RAG stack including Vault, Qdrant, Langfuse, and LLM orchestration service""" + + def __init__(self, compose_file_name: str = "docker-compose-eval.yml"): + self.project_root = Path(__file__).parent.parent.parent + self.compose_file_path = self.project_root / compose_file_name + self.compose: Optional[DockerCompose] = None + self.services_info: Dict[str, Dict[str, Any]] = {} + + if not self.compose_file_path.exists(): + raise FileNotFoundError( + f"Docker compose file not found: {self.compose_file_path}" + ) + + def start(self) -> None: + """Start all test containers and bootstrap Vault""" + logger.info("Starting RAG Stack testcontainers...") + os.environ["EVAL_MODE"] = "true" + + # Download embeddings from Azure before starting containers + self._download_embeddings_from_azure() + + # Prepare Vault Agent directories + agent_in = self.project_root / "test-vault" / "agents" / "llm" + agent_out = self.project_root / "test-vault" / "agent-out" + agent_in.mkdir(parents=True, exist_ok=True) + agent_out.mkdir(parents=True, exist_ok=True) + + # Clean up any stale files from previous runs + for f in ["role_id", "secret_id", "token", "pidfile", "dummy"]: + (agent_in / f).unlink(missing_ok=True) + (agent_out / f).unlink(missing_ok=True) + + # Start all Docker Compose services + logger.info("Starting Docker Compose services...") + self.compose = DockerCompose( + str(self.project_root), + compose_file_name=self.compose_file_path.name, + pull=False, + ) + self.compose.start() + + # Get Vault connection details + vault_url = self._get_vault_url() + logger.info(f"Vault URL: {vault_url}") + + # Wait for Vault to be ready + self._wait_for_vault_ready(vault_url) + + # Configure Vault with AppRole, policies, and test secrets + self._bootstrap_vault_dev(agent_in, vault_url) + + # Verify credentials were written successfully + role_id = (agent_in / "role_id").read_text().strip() + secret_id = (agent_in / "secret_id").read_text().strip() + logger.info( + f"AppRole credentials written: role_id={role_id[:8]}..., secret_id={secret_id[:8]}..." + ) + + # Wait for Vault Agent to authenticate and write token + logger.info("Waiting for vault-agent to authenticate...") + self._wait_for_valid_token(agent_out / "token", vault_url, max_attempts=20) + + logger.info("Vault Agent authenticated successfully") + + # Wait for other services to be ready + self._wait_for_services() + self._collect_service_info() + + # Index test data into Qdrant + self._index_test_data() + + logger.info("RAG Stack testcontainers ready") + + def stop(self) -> None: + """Stop all test containers""" + if self.compose: + logger.info("Stopping RAG Stack testcontainers...") + self.compose.stop() + logger.info("Testcontainers stopped") + + def _download_embeddings_from_azure(self) -> None: + """Download embeddings from Azure Blob Storage if configured.""" + connection_string = os.getenv("AZURE_STORAGE_CONNECTION_STRING") + container_name = os.getenv("AZURE_STORAGE_CONTAINER_NAME", "test-embeddings") + blob_name = os.getenv("AZURE_STORAGE_BLOB_NAME", "test_embeddings.json") + + # Local path where embeddings should be saved + embeddings_file = self.project_root / "tests" / "data" / "test_embeddings.json" + + # Require Azure configuration for CI/CD + if not connection_string: + raise ValueError( + "AZURE_STORAGE_CONNECTION_STRING is required to download embeddings. " + "Either set this environment variable or ensure test_embeddings.json " + f"exists at {embeddings_file}" + ) + + logger.info("=" * 80) + logger.info("DOWNLOADING EMBEDDINGS FROM AZURE BLOB STORAGE") + logger.info("=" * 80) + + try: + download_embeddings_from_azure( + connection_string=connection_string, + container_name=container_name, + blob_name=blob_name, + local_path=embeddings_file, + ) + logger.info("Embeddings download complete") + except Exception as e: + logger.error(f"Failed to download embeddings from Azure: {e}") + raise + + def _get_vault_url(self) -> str: + """Get the mapped Vault URL accessible from the host""" + if not self.compose: + raise RuntimeError("Docker Compose not initialized") + host = self.compose.get_service_host("vault", 8200) + port = self.compose.get_service_port("vault", 8200) + return f"http://{host}:{port}" + + def _wait_for_vault_ready(self, vault_url: str, timeout: int = 60) -> None: + """Wait for Vault to be initialized and unsealed""" + logger.info("Waiting for Vault to be available...") + client = hvac.Client(url=vault_url, token="root", timeout=10) + + start = time.time() + while time.time() - start < timeout: + try: + status = client.sys.read_health_status(method="GET") + if status.get("initialized", False) and not status.get("sealed", True): + logger.info("Vault is available and unsealed") + return + except Exception: + pass + time.sleep(2) + + raise TimeoutError("Vault did not become available within 60s") + + def _bootstrap_vault_dev(self, agent_in: Path, vault_url: str) -> None: + """ + Bootstrap Vault dev instance with: + - AppRole auth method + - Policy for LLM orchestration service + - AppRole role and credentials + - Test secrets (LLM connections, Langfuse, embeddings, guardrails) + """ + logger.info("Bootstrapping Vault with AppRole and test secrets...") + client = hvac.Client(url=vault_url, token="root") + + # Enable AppRole authentication method + if "approle/" not in client.sys.list_auth_methods(): + client.sys.enable_auth_method("approle") + logger.info("AppRole enabled") + + # Create policy with permissions for all secret paths (updated with correct embedding paths) + policy = """ +path "secret/metadata/llm/*" { capabilities = ["list"] } +path "secret/data/llm/*" { capabilities = ["read"] } +path "secret/metadata/langfuse/*" { capabilities = ["list"] } +path "secret/data/langfuse/*" { capabilities = ["read"] } +path "secret/metadata/embeddings/*" { capabilities = ["list"] } +path "secret/data/embeddings/*" { capabilities = ["read"] } +path "secret/metadata/guardrails/*" { capabilities = ["list"] } +path "secret/data/guardrails/*" { capabilities = ["read"] } +path "auth/token/lookup-self" { capabilities = ["read"] } +path "auth/token/renew-self" { capabilities = ["update"] } +""" + client.sys.create_or_update_policy("llm-orchestration", policy) + logger.info("Policy 'llm-orchestration' created") + + # Create AppRole role with service token type + role_name = "llm-orchestration-service" + client.write( + f"auth/approle/role/{role_name}", + **{ + "token_policies": "llm-orchestration", + "secret_id_ttl": "24h", + "token_ttl": "1h", + "token_max_ttl": "24h", + "secret_id_num_uses": 0, + "bind_secret_id": True, + "token_no_default_policy": True, + "token_type": "service", + }, + ) + logger.info(f"AppRole '{role_name}' created") + + # Generate credentials for the AppRole + role_id = client.read(f"auth/approle/role/{role_name}/role-id")["data"][ + "role_id" + ] + secret_id = client.write(f"auth/approle/role/{role_name}/secret-id")["data"][ + "secret_id" + ] + + # Write credentials to files that Vault Agent will read + (agent_in / "role_id").write_text(role_id, encoding="utf-8") + (agent_in / "secret_id").write_text(secret_id, encoding="utf-8") + logger.info("AppRole credentials written to agent-in/") + + # Write test secrets + self._write_test_secrets(client) + + def _write_test_secrets(self, client: hvac.Client) -> None: + """Write all test secrets to Vault with correct path structure""" + + # ============================================================ + # CRITICAL DEBUG SECTION - Environment Variables + # ============================================================ + logger.info("=" * 80) + logger.info("VAULT SECRET BOOTSTRAP - ENVIRONMENT VARIABLES DEBUG") + logger.info("=" * 80) + + azure_endpoint = os.getenv("AZURE_OPENAI_ENDPOINT") + azure_api_key = os.getenv("AZURE_OPENAI_API_KEY") + azure_deployment = os.getenv("AZURE_OPENAI_DEPLOYMENT") + azure_embedding_deployment = os.getenv("AZURE_OPENAI_EMBEDDING_DEPLOYMENT") + + # Validate critical environment variables + missing_vars = [] + if not azure_endpoint: + missing_vars.append("AZURE_OPENAI_ENDPOINT") + if not azure_api_key: + missing_vars.append("AZURE_OPENAI_API_KEY") + if not azure_embedding_deployment: + missing_vars.append("AZURE_OPENAI_EMBEDDING_DEPLOYMENT") + + if missing_vars: + error_msg = f"CRITICAL: Missing required environment variables: {', '.join(missing_vars)}" + logger.error(error_msg) + raise ValueError(error_msg) + + logger.info("All required environment variables are set") + logger.info("=" * 80) + + # ============================================================ + # CHAT MODEL SECRET (LLM path) + # ============================================================ + logger.info("") + logger.info("Writing LLM connection secret (chat model)...") + llm_secret = { + "connection_id": "evalconnection-1", + "endpoint": azure_endpoint, + "api_key": azure_api_key, + "deployment_name": azure_deployment or "gpt-4o-mini", + "environment": "testing", + "model": "gpt-4o-mini", + "model_type": "chat", + "api_version": "2024-02-15-preview", + "tags": "azure,test,chat", + } + + logger.info(f" → chat deployment: {llm_secret['deployment_name']}") + logger.info(f" → endpoint: {llm_secret['endpoint']}") + logger.info(f" → connection_id: {llm_secret['connection_id']}") + + client.secrets.kv.v2.create_or_update_secret( + mount_point="secret", + path="llm/connections/azure_openai/evalconnection-1", + secret=llm_secret, + ) + logger.info( + "LLM connection secret written to llm/connections/azure_openai/evalconnection-1" + ) + + # ============================================================ + # EMBEDDING MODEL SECRET (Embeddings path) + # ============================================================ + logger.info("") + logger.info("Writing embedding model secret...") + embedding_secret = { + "connection_id": "evalconnection-1", + "endpoint": azure_endpoint, + "api_key": azure_api_key, + "deployment_name": azure_embedding_deployment, # This is the embedding deployment + "environment": "testing", + "model": "text-embedding-3-large", + "model_type": "embedding", + "api_version": "2024-02-15-preview", + "max_tokens": 2048, + "vector_size": 3072, + "tags": "azure,embedding,test", + } + + logger.info(f" → model: {embedding_secret['model']}") + logger.info(f" → connection_id: {embedding_secret['connection_id']}") + logger.info( + " → Vault path: embeddings/connections/azure_openai/evalconnection-1" + ) + + # Write to embeddings path with connection_id in the path + client.secrets.kv.v2.create_or_update_secret( + mount_point="secret", + path="embeddings/connections/azure_openai/evalconnection-1", + secret=embedding_secret, + ) + logger.info( + "Embedding secret written to embeddings/connections/azure_openai/evalconnection-1" + ) + + # ============================================================ + # VERIFY SECRETS WERE WRITTEN CORRECTLY + # ============================================================ + logger.info("") + logger.info("Verifying secrets in Vault...") + try: + # Verify LLM path + verify_llm = client.secrets.kv.v2.read_secret_version( + path="llm/connections/azure_openai/development/evalconnection-1", + mount_point="secret", + ) + llm_data = verify_llm["data"]["data"] + logger.info("LLM path verified:") + logger.info(f" • connection_id: {llm_data.get('connection_id')}") + + # Verify embeddings path + verify_embedding = client.secrets.kv.v2.read_secret_version( + path="embeddings/connections/azure_openai/development/evalconnection-1", + mount_point="secret", + ) + embedding_data = verify_embedding["data"]["data"] + logger.info("Embeddings path verified:") + logger.info(f" • model: {embedding_data.get('model')}") + logger.info(f" • connection_id: {embedding_data.get('connection_id')}") + + # Critical validation + if embedding_data.get("deployment_name") != azure_embedding_deployment: + error_msg = ( + "VAULT SECRET MISMATCH! " + f"Expected deployment_name='{azure_embedding_deployment}' " + f"but Vault has '{embedding_data.get('deployment_name')}'" + ) + logger.error(error_msg) + raise ValueError(error_msg) + + if embedding_data.get("connection_id") != "evalconnection-1": + error_msg = ( + "VAULT SECRET MISMATCH! " + "Expected connection_id='evalconnection-1' " + f"but Vault has '{embedding_data.get('connection_id')}'" + ) + logger.error(error_msg) + raise ValueError(error_msg) + + logger.info("Secret verification PASSED") + + except Exception as e: + logger.error(f"Failed to verify secrets: {e}") + raise + + # ============================================================ + # LANGFUSE CONFIGURATION + # ============================================================ + logger.info("") + logger.info("Writing Langfuse configuration secret...") + langfuse_secret = { + "public_key": "pk-lf-test", + "secret_key": "sk-lf-test", + "host": "http://langfuse-web:3000", + } + client.secrets.kv.v2.create_or_update_secret( + mount_point="secret", path="langfuse/config", secret=langfuse_secret + ) + logger.info("Langfuse configuration secret written") + + logger.info("=" * 80) + logger.info("ALL SECRETS WRITTEN SUCCESSFULLY") + logger.info("=" * 80) + + def _capture_service_logs(self) -> None: + """Capture logs from all services before cleanup.""" + services = ["llm-orchestration-service", "vault", "qdrant", "langfuse-web"] + + for service in services: + try: + logger.info(f"\n{'=' * 60}") + logger.info(f"LOGS: {service}") + logger.info("=" * 60) + + result = subprocess.run( + [ + "docker", + "compose", + "-f", + str(self.compose_file_path), + "logs", + "--tail", + "200", + service, + ], + capture_output=True, + text=True, + timeout=10, + cwd=str(self.project_root), + ) + + if result.stdout: + logger.info(result.stdout) + if result.stderr: + logger.error(result.stderr) + + except Exception as e: + logger.error(f"Failed to capture logs for {service}: {e}") + + def _wait_for_valid_token( + self, token_path: Path, vault_url: str, max_attempts: int = 20 + ) -> None: + """Wait for Vault Agent to write a valid token and verify it works""" + for attempt in range(max_attempts): + if token_path.exists() and token_path.stat().st_size > 0: + try: + # Fix permissions before reading + self._fix_token_file_permissions(token_path) + + token = token_path.read_text().strip() + + client = hvac.Client(url=vault_url, token=token) + try: + client.lookup_token() + + if client.is_authenticated(): + logger.info(f"Valid token obtained (attempt {attempt + 1})") + self._verify_token_permissions(client) + return + except Exception as e: + if attempt < max_attempts - 1: + logger.debug( + f"Token validation error (attempt {attempt + 1}): {type(e).__name__}" + ) + except PermissionError as e: + logger.warning( + f"Permission error reading token file (attempt {attempt + 1}): {e}" + ) + # Try to fix permissions again + self._fix_token_file_permissions(token_path, force=True) + + time.sleep(2) + + logger.error("Failed to obtain valid Vault token") + self._check_agent_logs() + raise TimeoutError( + f"Failed to obtain valid Vault token after {max_attempts} attempts" + ) + + def _fix_token_file_permissions( + self, token_path: Path, force: bool = False + ) -> None: + """Fix permissions on token file to make it readable by host user""" + try: + # Try to change permissions using subprocess (requires Docker to be accessible) + if force: + logger.info( + "Attempting to fix token file permissions using docker exec..." + ) + result = subprocess.run( + [ + "docker", + "exec", + "vault-agent-llm", + "chmod", + "644", + "/agent/llm-token/token", + ], + capture_output=True, + text=True, + timeout=5, + ) + if result.returncode == 0: + logger.info( + "Successfully fixed token file permissions via docker exec" + ) + else: + logger.warning( + f"Failed to fix permissions via docker exec: {result.stderr}" + ) + + # Also try direct chmod (may not work in all environments) + try: + os.chmod(token_path, 0o644) + except Exception as chmod_error: + logger.debug( + f"Direct chmod failed (expected in some environments): {chmod_error}" + ) + + except Exception as e: + logger.debug(f"Could not fix token file permissions: {e}") + + def _verify_token_permissions(self, client: hvac.Client) -> None: + """Verify the token has correct permissions to read secrets""" + try: + client.secrets.kv.v2.read_secret_version( + path="llm/connections/azure_openai/development/evalconnection-1", + mount_point="secret", + ) + logger.info("Token has correct permissions to read secrets") + except Exception as e: + logger.error(f"Token cannot read secrets: {e}") + raise + + def _check_agent_logs(self) -> None: + """Check vault-agent logs for debugging authentication issues""" + result = subprocess.run( + ["docker", "logs", "--tail", "50", "vault-agent-llm"], + capture_output=True, + text=True, + ) + logger.error(f"Vault Agent Logs:\n{result.stdout}\n{result.stderr}") + + def _wait_for_services(self, total_timeout: int = 300) -> None: + """Wait for all services to be healthy""" + services = [ + ("qdrant", 6333, self._check_qdrant, 60), + ("langfuse-web", 3000, self._check_langfuse, 120), + ("llm-orchestration-service", 8100, self._check_orchestration, 180), + ] + start = time.time() + for name, port, check, timeout in services: + self._wait_single(name, port, check, timeout, start, total_timeout) + + def _wait_single( + self, + name: str, + port: int, + check: Any, + timeout: int, + global_start: float, + total_timeout: int, + ) -> None: + """Wait for a single service to be ready""" + if self.compose is None: + return + + logger.info(f"Waiting for {name}...") + start = time.time() + while time.time() - start < timeout: + try: + host = self.compose.get_service_host(name, port) + mapped_port = self.compose.get_service_port(name, port) + if check(host, mapped_port): + logger.info(f"{name} ready at {host}:{mapped_port}") + self.services_info[name] = { + "host": host, + "port": mapped_port, + "url": f"http://{host}:{mapped_port}", + } + return + except Exception: + pass + time.sleep(3) + raise TimeoutError(f"Timeout waiting for {name}") + + def _check_qdrant(self, host: str, port: int) -> bool: + """Check if Qdrant is ready""" + try: + r = requests.get(f"http://{host}:{port}/collections", timeout=5) + return r.status_code == 200 + except Exception: + return False + + def _check_langfuse(self, host: str, port: int) -> bool: + """Check if Langfuse is ready""" + try: + r = requests.get(f"http://{host}:{port}/api/public/health", timeout=5) + return r.status_code == 200 + except Exception: + return False + + def _check_orchestration(self, host: str, port: int) -> bool: + """Check if LLM orchestration service is healthy""" + try: + r = requests.get(f"http://{host}:{port}/health", timeout=5) + return r.status_code == 200 and r.json().get("status") == "healthy" + except Exception: + return False + + def _collect_service_info(self) -> None: + """Collect service connection information""" + if self.compose: + self.services_info["vault"] = { + "host": self.compose.get_service_host("vault", 8200), + "port": self.compose.get_service_port("vault", 8200), + "url": self._get_vault_url(), + } + + def get_orchestration_service_url(self) -> str: + """Get the URL for the LLM orchestration service""" + return self.services_info["llm-orchestration-service"]["url"] + + def get_qdrant_url(self) -> str: + """Get the URL for Qdrant""" + return self.services_info["qdrant"]["url"] + + def get_vault_url(self) -> str: + """Get the URL for Vault""" + return self.services_info["vault"]["url"] + + def get_langfuse_url(self) -> str: + """Get the URL for Langfuse""" + return self.services_info.get("langfuse-web", {}).get( + "url", "http://localhost:3000" + ) + + def is_service_available(self, service_name: str) -> bool: + """Check if a service is available""" + return service_name in self.services_info + + def _index_test_data(self) -> None: + """Index test documents into Qdrant for retrieval testing.""" + logger.info("Indexing test data into Qdrant contextual collections...") + + try: + from tests.helpers.test_data_loader import load_test_data_into_qdrant + + load_test_data_into_qdrant( + orchestration_url=self.get_orchestration_service_url(), + qdrant_url=self.get_qdrant_url(), + ) + + logger.info("Test data indexing complete") + + except Exception as e: + logger.error(f"Failed to index test data: {e}") + raise + + +# ===================== Pytest Fixtures ===================== + + +@pytest.fixture(scope="session") +def rag_stack() -> Generator[RAGStackTestContainers, None, None]: + """ + Session-scoped fixture that starts all test containers once per test session. + Containers are automatically stopped after all tests complete. + """ + stack = RAGStackTestContainers() + try: + stack.start() + yield stack + except Exception as e: + # If startup fails, capture logs before cleanup + logger.error(f"RAG stack startup failed: {e}") + try: + stack._capture_service_logs() + except Exception as e: + logger.error(f"Could not capture logs after startup failure: {e}") + pass + raise + finally: + logger.info("=" * 80) + logger.info("CAPTURING SERVICE LOGS BEFORE CLEANUP") + logger.info("=" * 80) + try: + stack._capture_service_logs() + except Exception as e: + logger.error(f"Could not capture logs: {e}") + stack.stop() + + +@pytest.fixture(scope="function") +def orchestration_client(rag_stack: RAGStackTestContainers): + """ + function-scoped fixture that provides the orchestration service URL. + Tests can use either requests (sync) or httpx (async). + """ + + class OrchestrationClient: + def __init__(self, base_url: str): + self.base_url = base_url + + return OrchestrationClient(rag_stack.get_orchestration_service_url()) + + +@pytest.fixture(scope="session") +def api_tool_endpoints_indexed(rag_stack: RAGStackTestContainers): + """ + Session-scoped fixture that seeds the API tool endpoints from + ``tests/api_tool_eval/test-endpoints.json`` into Qdrant's + ``api_tool_collection`` so the agentic loop can find them. + + Required by tests in ``tests/deepeval_tests/api_tool_tests.py`` — without + it, semantic search over API tools returns nothing and the orchestrator + falls back to RAG. + + The indexer module hard-codes container-internal URLs in its constants + (``http://llm-orchestration-service:8100``, ``qdrant:6333``). Two of the + three are used in code paths that re-read the constant at call time + (monkey-patching the class attribute works), but ``ApiToolQdrantManager`` + is instantiated with no arguments inside ``index_endpoint`` and binds the + host/port at function-definition time as defaults. To override those, the + fixture also replaces the ``ApiToolQdrantManager`` name imported into + ``main_indexer`` with a factory that injects the testcontainers-mapped + host/port. + """ + import asyncio + import json as _json + from pathlib import Path as _Path + from unittest.mock import patch + from urllib.parse import urlparse + + from api_tool_indexer import main_indexer + from api_tool_indexer.constants import ApiToolIndexerConstants + from api_tool_indexer.main_indexer import index_endpoint + from api_tool_indexer.models import EndpointData + from api_tool_indexer.qdrant_manager import ApiToolQdrantManager + + orch_url = rag_stack.get_orchestration_service_url() + qdrant_url = rag_stack.get_qdrant_url() + parsed = urlparse(qdrant_url) + qdrant_host = parsed.hostname or "localhost" + qdrant_port = parsed.port or 6333 + + def _qdrant_factory(*args: Any, **kwargs: Any) -> ApiToolQdrantManager: + kwargs.setdefault("host", qdrant_host) + kwargs.setdefault("port", qdrant_port) + return ApiToolQdrantManager(*args, **kwargs) + + endpoints_file = ( + _Path(__file__).parent.parent / "api_tool_eval" / "test-endpoints.json" + ) + if not endpoints_file.exists(): + pytest.skip( + f"API tool endpoint fixture not found: {endpoints_file} — " + "cannot seed api_tool_collection" + ) + + with open(endpoints_file, encoding="utf-8") as f: + endpoints = _json.load(f) + + async def _seed_all() -> list: + results = [] + for ep in endpoints: + endpoint_data = EndpointData( + endpoint_id=ep["endpointId"], + name=ep["name"], + description=ep["description"], + url=ep["url"], + method=ep["method"], + params=ep.get("params", []), + ) + res = await index_endpoint(endpoint_data) + logger.info( + f"Seeded endpoint '{ep['name']}': success={res.success} ({res.message})" + ) + results.append((ep["name"], res)) + return results + + # Connection ID "evalconnection-1" matches the one used by standard_tests.py + # and the EVAL_MODE setup; vault is seeded for it. The default + # "gpt-4o-mini" used by the indexer in production has no test fixture. + with ( + patch.object(ApiToolIndexerConstants, "DEFAULT_API_BASE_URL", orch_url), + patch.object( + ApiToolIndexerConstants, + "DEFAULT_CONNECTION_ID", + "evalconnection-1", + ), + patch.object(main_indexer, "ApiToolQdrantManager", _qdrant_factory), + ): + logger.info( + f"Seeding {len(endpoints)} API tool endpoints " + f"(orch={orch_url}, qdrant={qdrant_host}:{qdrant_port})" + ) + results = asyncio.run(_seed_all()) + + failed = [name for name, r in results if not r.success] + if failed: + pytest.skip( + f"API tool seeding failed for {len(failed)}/{len(endpoints)} " + f"endpoints ({failed}) — skipping API tool tests" + ) + + logger.info(f"Indexed all {len(endpoints)} API tool endpoints successfully") + yield diff --git a/tests/deepeval_tests/report_generator.py b/tests/deepeval_tests/report_generator.py index 2321cbec..398d1847 100644 --- a/tests/deepeval_tests/report_generator.py +++ b/tests/deepeval_tests/report_generator.py @@ -155,23 +155,16 @@ def generate_failure_analysis(results: Dict[str, Any]) -> str: analysis += "| Test | Query | Metric | Score | Issue |\n" analysis += "|------|--------|--------|-------|-------|\n" - for failure in failed_results[:10]: # Limit to first 10 failures + for failure in failed_results: # Limit to first 10 failures query_preview = ( failure["input"][:50] + "..." if len(failure["input"]) > 50 else failure["input"] ) - reason_preview = ( - failure["reason"][:100] + "..." - if len(failure["reason"]) > 100 - else failure["reason"] - ) + reason_preview = failure["reason"] analysis += f"| {failure['test_case']} | {query_preview} | {failure['metric']} | {failure['score']:.2f} | {reason_preview} |\n" - if len(failed_results) > 10: - analysis += f"\n*({len(failed_results) - 10} additional failures not shown)*\n" - analysis += "\n" return analysis diff --git a/tests/deepeval_tests/standard_tests.py b/tests/deepeval_tests/standard_tests.py index 6d8c9bd3..d1cdde78 100644 --- a/tests/deepeval_tests/standard_tests.py +++ b/tests/deepeval_tests/standard_tests.py @@ -12,15 +12,17 @@ ContextualRelevancyMetric, FaithfulnessMetric, ) +import asyncio +import httpx + sys.path.insert(0, str(Path(__file__).parent.parent)) -from mocks.dummy_llm_orchestrator import process_query class StandardResultCollector: """Collects test results during execution for report generation.""" - def __init__(self) -> None: + def __init__(self): self.results = { "total_tests": 0, "passed_tests": 0, @@ -108,10 +110,10 @@ def save_results_fixture(): class TestRAGSystem: - """Test suite for RAG system evaluation using DeepEval metrics.""" + """Test suite for RAG system evaluation using DeepEval metrics via API.""" @classmethod - def setup_class(cls) -> None: + def setup_class(cls): """Setup test class with metrics and test data.""" print("Setting up TestRAGSystem...") @@ -129,23 +131,6 @@ def setup_class(cls) -> None: print(f"Loaded {len(cls.test_data)} test cases") - def create_test_case( - self, data_item: Dict[str, Any], provider: str = "anthropic" - ) -> LLMTestCase: - """Create a DeepEval test case from data item.""" - # Generate actual output using the dummy orchestrator - result = process_query( - question=data_item["input"], provider=provider, include_contexts=True - ) - - llm_test_case = LLMTestCase( - input=data_item["input"], - actual_output=result["response"], - expected_output=data_item["expected_output"], - retrieval_context=result["retrieval_context"], - ) - return llm_test_case - @pytest.mark.parametrize( "test_item", [ @@ -159,20 +144,82 @@ def create_test_case( ) ], ) - def test_all_metrics(self, test_item: Dict[str, Any]): - """Test all metrics for each test case and collect results.""" - test_case = self.create_test_case(test_item) + @pytest.mark.asyncio + async def test_all_metrics(self, test_item: Dict[str, Any], orchestration_client): + """Async version of DeepEval test with parallel metric execution.""" - # Get test case index for consistent numbering + orchestration_url = orchestration_client.base_url test_case_num = self.test_data.index(test_item) + 1 - print(f"\nTesting case {test_case_num}: {test_item['input'][:50]}...") - # Initialize metrics results - metrics_results = {} - failed_assertions = [] + # --- USE ASYNC HTTP CLIENT --- + result = None + async with httpx.AsyncClient(timeout=60.0) as client: + try: + response = await client.post( + f"{orchestration_url}/orchestrate-eval", + json={ + "chatId": f"test-{test_item.get('id', 'unknown')}", + "message": test_item["input"], + "authorId": "deepeval-tester", + "conversationHistory": [], + "url": "https://test.example.com", + "environment": "testing", + "connection_id": "evalconnection-1", + }, + ) + response.raise_for_status() + result = response.json() + except httpx.RequestError as e: + result = {"content": f"API Error: {str(e)}", "retrieval_context": []} + except Exception as e: + result = { + "content": f"Unexpected error: {str(e)}", + "retrieval_context": [], + } + if result is None: + result = {"content": "No response received", "retrieval_context": []} + # --- DEBUG LOGGING --- + print("=" * 80) + print(f"TEST CASE {test_case_num} API RESPONSE DEBUG") + print("=" * 80) + print(f"Response keys: {list(result.keys())}") + for key, value in result.items(): + print(key, value) + print(f"Content length: {len(result.get('content', ''))}") + print(f"Retrieval context: {len(result.get('retrieval_context') or [])} chunks") + + if result.get("retrieval_context"): + for chunk in result["retrieval_context"]: + print(chunk.keys()) + context = ( + chunk.get("content", "") if isinstance(chunk, dict) else str(chunk) + ) + meta = chunk.get("metadata", {}) if isinstance(chunk, dict) else {} + fused_score = meta.get("fused_score", "N/A") + bm25_score = meta.get("bm25_score", "N/A") + semantic_score = meta.get("semantic_score", "N/A") + print( + f"Chunk (fused: {fused_score}, bm25: {bm25_score}, semantic: {semantic_score}):\n {context}\n\n" + ) + else: + print("WARNING: No retrieval context returned!") + print("=" * 80) + + retrieval_context = result.get("retrieval_context") or [] + retrieval_context = [ + c.get("content", "") if isinstance(c, dict) else str(c) + for c in retrieval_context + ] + + llm_test_case = LLMTestCase( + input=test_item["input"], + actual_output=result.get("content", ""), + expected_output=test_item["expected_output"], + retrieval_context=retrieval_context, + ) - # Define all metrics to test + # --- Run metrics concurrently --- metrics = [ ("contextual_precision", self.contextual_precision), ("contextual_recall", self.contextual_recall), @@ -181,38 +228,33 @@ def test_all_metrics(self, test_item: Dict[str, Any]): ("faithfulness", self.faithfulness), ] - # Test each metric and collect results - for metric_name, metric in metrics: + async def run_metric(metric_name, metric): try: - metric.measure(test_case) + await asyncio.to_thread(metric.measure, llm_test_case) score = metric.score - passed = score >= 0.7 - reason = metric.reason - - metrics_results[metric_name] = { + return metric_name, { "score": score, - "passed": passed, - "reason": reason, + "passed": score >= 0.4, + "reason": metric.reason, } - - print(f" {metric_name}: {score:.3f} ({'PASS' if passed else 'FAIL'})") - - # Collect failed assertions but don't raise immediately - if not passed: - failed_assertions.append( - f"{metric_name} failed for query: '{test_item['input']}'. " - f"Score: {score}, Reason: {reason}" - ) - except Exception as e: - metrics_results[metric_name] = { + return metric_name, { "score": 0.0, "passed": False, "reason": f"Error: {str(e)}", } - failed_assertions.append(f"{metric_name} error: {str(e)}") - # Always add results to collector, regardless of pass/fail + # Run metrics sequentially with delays to avoid rate limiting + metrics_results = {} + for i, (name, metric) in enumerate(metrics): + print(f" Running {name} metric...") + result_name, result_data = await run_metric(name, metric) + metrics_results[result_name] = result_data + # Add delay between metrics to respect rate limits (except after last metric) + if i < len(metrics) - 1: + await asyncio.sleep(5) # 5 second delay for Azure gpt-4.1 50K TPM limit + + # --- Collect results --- try: standard_results_collector.add_test_result( test_case_num=test_case_num, @@ -224,7 +266,9 @@ def test_all_metrics(self, test_item: Dict[str, Any]): except Exception as e: print(f"Error adding test result: {e}") - # Now raise assertion if any metrics failed (for pytest reporting) - if failed_assertions: - # Just raise the first failure to keep pytest output clean - raise AssertionError(failed_assertions[0]) + # --- Assert --- + failed = [name for name, res in metrics_results.items() if not res["passed"]] + if failed: + pytest.fail( + f"Metrics failed: {', '.join(failed)} for input: {test_item['input'][:50]}" + ) diff --git a/tests/helpers/test_data_loader.py b/tests/helpers/test_data_loader.py new file mode 100644 index 00000000..f95c203a --- /dev/null +++ b/tests/helpers/test_data_loader.py @@ -0,0 +1,170 @@ +"""Helper module to load test data into Qdrant before running tests.""" + +import json +import uuid +from typing import List, Dict, Any, Tuple +from pathlib import Path +from loguru import logger +from datetime import datetime +import httpx + + +def load_test_data_into_qdrant( + orchestration_url: str, + qdrant_url: str, +) -> None: + """Load test documents into Qdrant contextual collections for retrieval testing.""" + logger.info("Loading test data into Qdrant contextual collections...") + + # Load pre-computed embeddings + embeddings_file = Path(__file__).parent.parent / "data" / "test_embeddings.json" + + if not embeddings_file.exists(): + raise FileNotFoundError( + f"Pre-computed embeddings not found at {embeddings_file}. " + "Run create_embeddings.py first!" + ) + + logger.info(f"Loading pre-computed embeddings from {embeddings_file}") + chunks_data, model_used = load_precomputed_embeddings(embeddings_file) + + # Index into Qdrant + index_embeddings_to_qdrant( + qdrant_url=qdrant_url, chunks_data=chunks_data, model_used=model_used + ) + + +def load_precomputed_embeddings( + embeddings_file: Path, +) -> Tuple[List[Dict[str, Any]], str]: + """Load pre-computed embeddings from file.""" + with open(embeddings_file, "r", encoding="utf-8") as f: + data = json.load(f) + + chunks_data = data["chunks"] + model_used = data["model_used"] + + logger.info(f"Loaded {len(chunks_data)} pre-computed chunks") + logger.info(f" Vector size: {data['vector_size']}") + logger.info(f" Model: {model_used}") + logger.info(f" Documents: {data['total_documents']}") + + return chunks_data, model_used + + +def index_embeddings_to_qdrant( + qdrant_url: str, chunks_data: List[Dict[str, Any]], model_used: str +) -> None: + """Index embeddings into Qdrant.""" + if not chunks_data: + logger.warning("No chunks to index") + return + + vector_size = chunks_data[0]["vector_dimensions"] + collection_name = _determine_collection_from_model(model_used) + + logger.info(f"Indexing into Qdrant collection: {collection_name}") + + client = httpx.Client(timeout=30.0) + + try: + # Check if collection exists + response = client.get(f"{qdrant_url}/collections/{collection_name}") + + if response.status_code == 404: + logger.info(f"Creating collection '{collection_name}'...") + create_payload = { + "vectors": { + "size": vector_size, + "distance": "Cosine", + }, + "optimizers_config": {"default_segment_number": 2}, + "replication_factor": 1, + } + + response = client.put( + f"{qdrant_url}/collections/{collection_name}", json=create_payload + ) + + if response.status_code not in [200, 201]: + raise RuntimeError(f"Failed to create collection: {response.text}") + + logger.info(f"Created collection '{collection_name}'") + else: + logger.info(f"Collection '{collection_name}' already exists") + + # Prepare points + points = [] + for chunk in chunks_data: + point_id = str(uuid.uuid5(uuid.NAMESPACE_DNS, chunk["chunk_id"])) + + payload = { + "chunk_id": chunk["chunk_id"], + "document_hash": chunk["document_hash"], + "chunk_index": chunk["chunk_index"], + "total_chunks": chunk["total_chunks"], + "original_content": chunk["original_content"], + "contextual_content": chunk["contextual_content"], + "context_only": chunk["context"], + "embedding_model": chunk["embedding_model"], + "vector_dimensions": chunk["vector_dimensions"], + "document_url": chunk["metadata"].get("source", "test_document"), + "dataset_collection": chunk["metadata"].get( + "dataset_collection", "test_collection" + ), + "processing_timestamp": datetime.now().isoformat(), + "tokens_count": chunk["tokens_count"], + **chunk["metadata"], + } + + points.append( + {"id": point_id, "vector": chunk["embedding"], "payload": payload} + ) + + # Upsert points in batches + batch_size = 100 + for i in range(0, len(points), batch_size): + batch = points[i : i + batch_size] + upsert_payload = {"points": batch} + + response = client.put( + f"{qdrant_url}/collections/{collection_name}/points", + json=upsert_payload, + ) + + if response.status_code not in [200, 201]: + raise RuntimeError(f"Failed to upsert points: {response.text}") + + logger.info(f"Indexed batch {i // batch_size + 1} ({len(batch)} points)") + + client.close() + + logger.info(f" Successfully indexed {len(points)} chunks into Qdrant") + + except Exception as e: + logger.error(f"Failed to index to Qdrant: {e}") + raise + + +def _determine_collection_from_model(model_name: str) -> str: + """Determine which Qdrant collection to use based on embedding model.""" + model_lower = model_name.lower() + + # Azure OpenAI models -> contextual_chunks_azure + if any( + keyword in model_lower for keyword in ["azure", "text-embedding", "ada-002"] + ): + return "contextual_chunks_azure" + + # AWS Bedrock models -> contextual_chunks_aws + elif any( + keyword in model_lower for keyword in ["titan", "amazon", "aws", "bedrock"] + ): + return "contextual_chunks_aws" + + # Default to Azure collection + else: + logger.warning( + f"Unknown model {model_name}, defaulting to contextual_chunks_azure" + ) + return "contextual_chunks_azure" diff --git a/tests/integration_tests/conftest.py b/tests/integration_tests/conftest.py index 1f6bd647..dad1af62 100644 --- a/tests/integration_tests/conftest.py +++ b/tests/integration_tests/conftest.py @@ -13,6 +13,14 @@ from loguru import logger +# Deterministic vault_uuid used for the integration-test LLM connection. +# Vault secret paths now terminate in the connection's vault_uuid (no environment +# segment) — see SecretResolver._build_vault_path. We pin a fixed, valid UUID so +# the secret we write, the DB connection row, and the value the DSL forwards to +# /orchestrate/test all line up without depending on a DB-generated UUID. +TEST_VAULT_UUID = "00000000-0000-0000-0000-000000000001" + + # ===================== VaultAgentClient ===================== class VaultAgentClient: """Client for interacting with Vault using a token written by Vault Agent""" @@ -298,17 +306,30 @@ def _write_test_secrets(self, client: hvac.Client) -> None: logger.info("All required environment variables are set") logger.info("=" * 80) + # ============================================================ + # SECRET PATHS + # ============================================================ + # Vault paths now terminate in the connection's vault_uuid (no environment + # segment) — see SecretResolver._build_vault_path / _build_embedding_vault_path. + # We write both the chat and embedding secrets under the SAME pinned + # vault_uuid (TEST_VAULT_UUID); the testing DB connection is forced to carry + # this UUID in ensure_testing_connection, and inference/test.yml forwards it + # to /orchestrate/test. The LLM, embedding and guardrails resolvers all read + # the secret using this UUID as the connection_id. + llm_vault_path = f"llm/connections/azure_openai/{TEST_VAULT_UUID}" + embedding_vault_path = f"embeddings/connections/azure_openai/{TEST_VAULT_UUID}" + # ============================================================ # CHAT MODEL SECRET (LLM path) # ============================================================ logger.info("") logger.info("Writing LLM connection secret (chat model)...") llm_secret = { - "connection_id": "gpt-4o-mini", + "connection_id": TEST_VAULT_UUID, "endpoint": azure_endpoint, "api_key": azure_api_key, "deployment_name": azure_deployment or "gpt-4o-mini", - "environment": "production", + "environment": "test", "model": "gpt-4o-mini", "model_type": "chat", "api_version": "2024-02-15-preview", @@ -317,16 +338,14 @@ def _write_test_secrets(self, client: hvac.Client) -> None: logger.info(f" chat deployment: {llm_secret['deployment_name']}") logger.info(f" endpoint: {llm_secret['endpoint']}") - logger.info(f" connection_id: {llm_secret['connection_id']}") + logger.info(f" vault path: {llm_vault_path}") client.secrets.kv.v2.create_or_update_secret( mount_point="secret", - path="llm/connections/azure_openai/production/gpt-4o-mini", + path=llm_vault_path, secret=llm_secret, ) - logger.info( - "LLM connection secret written to llm/connections/azure_openai/production/gpt-4o-mini" - ) + logger.info(f"LLM connection secret written to {llm_vault_path}") # ============================================================ # EMBEDDING MODEL SECRET (Embeddings path) @@ -334,31 +353,25 @@ def _write_test_secrets(self, client: hvac.Client) -> None: logger.info("") logger.info("Writing embedding model secret...") embedding_secret = { - "connection_id": "2", + "connection_id": TEST_VAULT_UUID, "endpoint": azure_embedding_endpoint, "api_key": azure_api_key, "deployment_name": azure_embedding_deployment, - "environment": "production", + "environment": "test", "model": "text-embedding-3-large", "api_version": "2024-12-01-preview", "tags": "azure,test,text-embedding-3-large", } logger.info(f" → model: {embedding_secret['model']}") - logger.info(f" → connection_id: {embedding_secret['connection_id']}") - logger.info( - " → Vault path: embeddings/connections/azure_openai/production/text-embedding-3-large" - ) + logger.info(f" → vault path: {embedding_vault_path}") - # Write to embeddings path with connection_id in the path client.secrets.kv.v2.create_or_update_secret( mount_point="secret", - path="embeddings/connections/azure_openai/production/text-embedding-3-large", + path=embedding_vault_path, secret=embedding_secret, ) - logger.info( - "Embedding secret written to embeddings/connections/azure_openai/production/text-embedding-3-large" - ) + logger.info(f"Embedding secret written to {embedding_vault_path}") # ============================================================ # VERIFY SECRETS WERE WRITTEN CORRECTLY @@ -368,7 +381,7 @@ def _write_test_secrets(self, client: hvac.Client) -> None: try: # Verify LLM path verify_llm = client.secrets.kv.v2.read_secret_version( - path="llm/connections/azure_openai/production/gpt-4o-mini", + path=llm_vault_path, mount_point="secret", ) llm_data = verify_llm["data"]["data"] @@ -377,7 +390,7 @@ def _write_test_secrets(self, client: hvac.Client) -> None: # Verify embeddings path verify_embedding = client.secrets.kv.v2.read_secret_version( - path="embeddings/connections/azure_openai/production/text-embedding-3-large", + path=embedding_vault_path, mount_point="secret", ) embedding_data = verify_embedding["data"]["data"] @@ -395,10 +408,10 @@ def _write_test_secrets(self, client: hvac.Client) -> None: logger.error(error_msg) raise ValueError(error_msg) - if embedding_data.get("connection_id") != "2": + if embedding_data.get("connection_id") != TEST_VAULT_UUID: error_msg = ( "VAULT SECRET MISMATCH! " - "Expected connection_id='2' " + f"Expected connection_id='{TEST_VAULT_UUID}' " f"but Vault has '{embedding_data.get('connection_id')}'" ) logger.error(error_msg) @@ -410,42 +423,6 @@ def _write_test_secrets(self, client: hvac.Client) -> None: logger.error(f"Failed to verify secrets: {e}") raise - # add the same secret configs to the 'testing' environment for test purposes - # connection_id is 1 (must match the database connection ID created by ensure_testing_connection) - llm_secret = { - "connection_id": 1, - "endpoint": azure_endpoint, - "api_key": azure_api_key, - "deployment_name": azure_deployment or "gpt-4o-mini", - "environment": "test", - "model": "gpt-4o-mini", - "model_type": "chat", - "api_version": "2024-02-15-preview", - "tags": "azure,test,chat", - } - client.secrets.kv.v2.create_or_update_secret( - mount_point="secret", - path="llm/connections/azure_openai/test/1", - secret=llm_secret, - ) - - embedding_secret = { - "connection_id": 1, - "endpoint": azure_embedding_endpoint, - "api_key": azure_api_key, - "deployment_name": azure_embedding_deployment, - "environment": "test", - "model": "text-embedding-3-large", - "api_version": "2024-12-01-preview", - "tags": "azure,test,text-embedding-3-large", - } - # Write to embeddings path with connection_id in the path - client.secrets.kv.v2.create_or_update_secret( - mount_point="secret", - path="embeddings/connections/azure_openai/test/1", - secret=embedding_secret, - ) - # ============================================================ # LANGFUSE CONFIGURATION # ============================================================ @@ -697,7 +674,7 @@ def _verify_token_permissions(self, client: hvac.Client) -> None: """Verify the token has correct permissions to read secrets""" try: client.secrets.kv.v2.read_secret_version( - path="llm/connections/azure_openai/production/gpt-4o-mini", + path=f"llm/connections/azure_openai/{TEST_VAULT_UUID}", mount_point="secret", ) logger.info("Token has correct permissions to read secrets") @@ -1205,6 +1182,10 @@ def postgres_client(rag_stack: RAGStackTestContainers): database="rag-search", user="postgres", password="dbadmin", + # RAG tables were moved to the rag_search schema (Liquibase v7 + # schema migration). Put it on the search_path so unqualified + # queries (e.g. FROM llm_connections) resolve. + options="-c search_path=rag_search,public", ) logger.info("PostgreSQL connection established") yield conn @@ -1372,10 +1353,9 @@ def ensure_testing_connection(postgres_client, ruuter_private_client, rag_stack) f"Found existing testing gpt-4o-mini connection: " f"ID={connection_id}, Name='{connection_name}'" ) - logger.warning( - f"IMPORTANT: Vault secret must exist at path: " - f"llm/connections/azure_openai/test/{connection_id}" - ) + # Pin the row's vault_uuid so it matches the vault secret path and the + # value inference/test.yml forwards to /orchestrate/test. + _pin_vault_uuid(postgres_client, connection_id) return connection_id # No testing gpt-4o-mini found - create one @@ -1414,28 +1394,46 @@ def ensure_testing_connection(postgres_client, ruuter_private_client, rag_stack) connection_id = response_data["id"] logger.info(f"Created testing gpt-4o-mini connection with ID: {connection_id}") - logger.warning( - f"IMPORTANT: Vault secret must exist at path: " - f"llm/connections/azure_openai/test/{connection_id}" - ) - logger.warning( - "Currently hardcoded vault path is: llm/connections/azure_openai/test/1" - ) - if connection_id != 1: - logger.error( - f"CONNECTION ID MISMATCH! Database assigned ID={connection_id}, " - f"but vault secret is at path .../test/1" - ) # Wait for database write time.sleep(2) + # Pin the row's vault_uuid so it matches the vault secret path and the + # value inference/test.yml forwards to /orchestrate/test. + _pin_vault_uuid(postgres_client, connection_id) + return connection_id finally: cursor.close() +def _pin_vault_uuid(postgres_client, connection_id: int) -> None: + """Force a testing connection's vault_uuid to the pinned TEST_VAULT_UUID. + + Vault secret paths terminate in the connection's vault_uuid (no environment + segment). We write the integration-test secrets under TEST_VAULT_UUID, so the + DB row that inference/test.yml looks up must carry the same UUID for the + forwarded value to resolve to the right vault path. + """ + cursor = postgres_client.cursor() + try: + cursor.execute( + "UPDATE llm_connections SET vault_uuid = %s WHERE id = %s", + (TEST_VAULT_UUID, connection_id), + ) + postgres_client.commit() + logger.info( + f"Pinned vault_uuid={TEST_VAULT_UUID} on testing connection id={connection_id}" + ) + except Exception as e: + postgres_client.rollback() + logger.error(f"Failed to pin vault_uuid on connection {connection_id}: {e}") + raise + finally: + cursor.close() + + @pytest.fixture(scope="session", autouse=True) def capture_container_logs_on_exit(rag_stack): """ diff --git a/tests/integration_tests/test_indexing.py b/tests/integration_tests/test_indexing.py index 08c14f5e..b9afac02 100644 --- a/tests/integration_tests/test_indexing.py +++ b/tests/integration_tests/test_indexing.py @@ -8,16 +8,11 @@ 4. Verify embeddings in Qdrant """ -import pytest -import zipfile import tempfile from pathlib import Path from datetime import timedelta import json -import requests import sys -import time -from loguru import logger from minio import Minio from qdrant_client import QdrantClient @@ -95,349 +90,6 @@ def test_document_structure(self, minio_client: Minio, test_document): assert meta["source"] == "integration_test" assert "title" in meta - @pytest.mark.asyncio - async def test_indexing_pipeline_e2e( - self, - rag_stack, - minio_client: Minio, - qdrant_client: QdrantClient, - test_bucket: str, - postgres_client, - setup_agency_sync_schema, - tmp_path: Path, - llm_orchestration_url: str, - ): - """ - End-to-end test of the indexing pipeline using Ruuter and Cron-Manager. - - This test: - 1. Creates test document and uploads to MinIO - 2. Generates presigned URL - 3. Prepares database (agency_sync + mock_ckb) - 4. Calls Ruuter endpoint to trigger indexing via Cron-Manager - 5. Waits for async indexing to complete (polls Qdrant) - 6. Verifies vectors stored in Qdrant - """ - # Step 0: Wait for LLM orchestration service to be healthy - max_retries = 30 - for i in range(max_retries): - try: - response = requests.get(f"{llm_orchestration_url}/health", timeout=5) - if response.status_code == 200: - health_data = response.json() - if health_data.get("orchestration_service") == "initialized": - break - except requests.exceptions.RequestException: - logger.debug( - f"LLM orchestration health check attempt {i + 1}/{max_retries} failed" - ) - time.sleep(2) - else: - pytest.fail("LLM orchestration service not healthy after 60 seconds") - - # Step 1: Create test document and upload to MinIO - # Create structure: test_agency//cleaned.txt - # so when extracted it becomes: extracted_datasets/test_agency//cleaned.txt - # The document loader expects: collection/hash_dir/cleaned.txt - source_dir = tmp_path / "source" - hash_dir = source_dir / "test_agency" / "doc_hash_001" - hash_dir.mkdir(parents=True) - dataset_dir = hash_dir - - cleaned_content = """This is an integration test document for the RAG Module. - -It tests the full vector indexing pipeline from end to end. - -The document will be chunked and embedded using the configured embedding model. - -Each chunk will be stored in Qdrant with contextual information generated by the LLM. - -The RAG (Retrieval-Augmented Generation) system uses semantic search to find relevant documents. - -Vector embeddings are numerical representations of text that capture semantic meaning. - -Qdrant is a vector database that enables fast similarity search across embeddings. - -The contextual retrieval process adds context to each chunk before embedding. - -This helps improve search accuracy by providing more context about each chunk's content. - -The LLM orchestration service manages connections to various language model providers. - -Supported providers include Azure OpenAI and AWS Bedrock for both LLM and embedding models. - -Integration testing ensures all components work together correctly in the pipeline. - -The MinIO object storage is used to store and retrieve dataset files for processing. - -Presigned URLs allow secure, temporary access to objects in MinIO buckets. - -The vector indexer downloads datasets, processes documents, and stores embeddings. - -Each document goes through chunking, contextual enrichment, and embedding stages. - -The final embeddings are upserted into Qdrant collections for later retrieval. - -This test verifies the complete flow from upload to storage in the vector database. -""" - (dataset_dir / "cleaned.txt").write_text(cleaned_content) - - meta = { - "source": "e2e_test", - "title": "E2E Test Document", - "agency_id": "test_agency", - } - (dataset_dir / "cleaned.meta.json").write_text(json.dumps(meta)) - - # Create ZIP without datasets/ prefix - just test_agency/files - zip_path = tmp_path / "test_dataset.zip" - with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zf: - for file in dataset_dir.rglob("*"): - if file.is_file(): - # Archive path: test_agency/cleaned.txt - arcname = file.relative_to(source_dir) - zf.write(file, arcname) - - object_name = "datasets/test_dataset.zip" - minio_client.fput_object(test_bucket, object_name, str(zip_path)) - - # Use simple direct URL instead of presigned URL - # Bucket is public, so no signature needed - dataset_url = f"http://minio:9000/{test_bucket}/{object_name}" - logger.info(f"Dataset URL for Docker network: {dataset_url}") - - # Step 1: Prepare database state for agency sync - cursor = postgres_client.cursor() - try: - # Insert agency_sync record with initial hash - cursor.execute( - """ - INSERT INTO public.agency_sync (id, agency_data_hash, data_url) - VALUES (%s, %s, %s) - ON CONFLICT (id) DO UPDATE - SET agency_data_hash = EXCLUDED.agency_data_hash - """, - ("test_agency", "initial_hash_000", ""), - ) - - # Insert mock CKB data with new hash and presigned URL - cursor.execute( - """ - INSERT INTO public.mock_ckb (client_id, client_data_hash, signed_s3_url) - VALUES (%s, %s, %s) - ON CONFLICT (client_id) DO UPDATE - SET client_data_hash = EXCLUDED.client_data_hash, - signed_s3_url = EXCLUDED.signed_s3_url - """, - ("test_agency", "new_hash_001", dataset_url), - ) - - postgres_client.commit() - logger.info( - "Database prepared: agency_sync (initial_hash_000) and mock_ckb (new_hash_001)" - ) - finally: - cursor.close() - - # Step 2: Call Ruuter Public endpoint to trigger indexing via Cron-Manager - logger.info("Calling /rag-search/data/update to trigger indexing...") - ruuter_public_url = "http://localhost:8086" - - response = requests.post( - f"{ruuter_public_url}/rag-search/data/update", - json={}, # No body required - timeout=60, - ) - - assert response.status_code == 200, ( - f"Expected 200, got {response.status_code}: {response.text}" - ) - data = response.json() - response_data = data.get("response", {}) - assert response_data.get("operationSuccessful") is True, ( - f"Operation failed: {data}" - ) - logger.info( - f"Indexing triggered successfully: {response_data.get('message', 'No message')}" - ) - - # Give Cron-Manager time to start the indexing process - logger.info("Waiting 5 seconds for Cron-Manager to start indexing...") - time.sleep(5) - - # Step 3: Wait for indexing to complete (poll Qdrant with verbose logging) - import asyncio - - max_wait = 120 # 2 minutes - poll_interval = 5 # seconds - start_time = time.time() - - logger.info(f"Waiting for indexing to complete (max {max_wait}s)...") - - # First, wait for collection to be created - collection_created = False - logger.info("Waiting for collection 'contextual_chunks_azure' to be created...") - - while time.time() - start_time < max_wait: - elapsed = time.time() - start_time - - try: - # Try to get collection info (will fail if doesn't exist) - collection_info = qdrant_client.get_collection( - "contextual_chunks_azure" - ) - if collection_info: - logger.info( - f"[{elapsed:.1f}s] Collection 'contextual_chunks_azure' created!" - ) - collection_created = True - break - except Exception as e: - logger.debug( - f"[{elapsed:.1f}s] Collection not yet created: {type(e).__name__}" - ) - - await asyncio.sleep(poll_interval) - - if not collection_created: - # Capture Cron-Manager logs for debugging - import subprocess - - try: - logger.error( - "Collection was not created - capturing Cron-Manager logs..." - ) - result = subprocess.run( - ["docker", "logs", "cron-manager", "--tail", "200"], - capture_output=True, - text=True, - timeout=10, - ) - logger.error("=" * 80) - logger.error("CRON-MANAGER LOGS:") - logger.error("=" * 80) - if result.stdout: - logger.error(result.stdout) - if result.stderr: - logger.error(f"STDERR: {result.stderr}") - except Exception as e: - logger.error(f"Failed to capture logs: {e}") - - pytest.fail( - f"Collection 'contextual_chunks_azure' was not created within {max_wait}s timeout" - ) - - # Now wait for documents to be indexed - indexing_completed = False - logger.info("Waiting for documents to be indexed in contextual_chunks_azure...") - poll_count = 0 - while time.time() - start_time < max_wait: - elapsed = time.time() - start_time - poll_count += 1 - - try: - azure_points = qdrant_client.count( - collection_name="contextual_chunks_azure" - ) - current_count = azure_points.count - - logger.info( - f"[{elapsed:.1f}s] Polling Qdrant: {current_count} documents in contextual_chunks_azure" - ) - - if current_count > 0: - logger.info( - f"✓ Indexing completed successfully in {elapsed:.1f}s with {current_count} documents" - ) - indexing_completed = True - break - - # After 30 seconds with no documents, check Cron-Manager logs once - if poll_count == 6 and current_count == 0: - import subprocess - - try: - logger.warning( - "No documents after 30s - checking Cron-Manager logs..." - ) - result = subprocess.run( - ["docker", "logs", "cron-manager", "--tail", "100"], - capture_output=True, - text=True, - timeout=5, - ) - if ( - "error" in result.stdout.lower() - or "failed" in result.stdout.lower() - ): - logger.error("Found errors in Cron-Manager logs:") - logger.error(result.stdout[-2000:]) # Last 2000 chars - except Exception as e: - logger.warning(f"Could not check logs: {e}") - - except Exception as e: - logger.warning(f"[{elapsed:.1f}s] Qdrant polling error: {e}") - - await asyncio.sleep(poll_interval) - - if not indexing_completed: - # Capture final state and Cron-Manager logs - try: - final_count = qdrant_client.count( - collection_name="contextual_chunks_azure" - ) - logger.error( - f"Final count after timeout: {final_count.count} documents" - ) - except Exception as e: - logger.error(f"Could not get final count: {e}") - - # Get Cron-Manager logs to see what happened - import subprocess - - try: - logger.error("=" * 80) - logger.error("CRON-MANAGER LOGS (indexing phase):") - logger.error("=" * 80) - result = subprocess.run( - ["docker", "logs", "cron-manager", "--tail", "300"], - capture_output=True, - text=True, - timeout=10, - ) - if result.stdout: - logger.error(result.stdout) - if result.stderr: - logger.error(f"STDERR: {result.stderr}") - except Exception as e: - logger.error(f"Failed to capture logs: {e}") - - pytest.fail( - f"Indexing did not complete within {max_wait}s timeout - no documents found in collection" - ) - - # Step 4: Verify vectors are stored in Qdrant - collections_to_check = ["contextual_chunks_azure", "contextual_chunks_aws"] - total_points = 0 - - for collection_name in collections_to_check: - try: - collection_info = qdrant_client.get_collection(collection_name) - if collection_info: - total_points += collection_info.points_count - except Exception: - # Collection might not exist - pass - - assert total_points > 0, ( - f"No vectors stored in Qdrant. Expected chunks but found {total_points} points." - ) - - logger.info( - f"E2E Test passed: Indexing completed via Ruuter/Cron-Manager, " - f"{total_points} points stored in Qdrant" - ) - class TestQdrantOperations: """Test Qdrant-specific operations.""" diff --git a/tests/integration_tests/test_inference.py b/tests/integration_tests/test_inference.py index 7529479c..5858d955 100644 --- a/tests/integration_tests/test_inference.py +++ b/tests/integration_tests/test_inference.py @@ -9,11 +9,55 @@ 5. Contextual retrieval integration """ +import subprocess import requests import json from loguru import logger +def _capture_orchestration_error_logs(tail: int = 300) -> str: + """Return recent llm-orchestration-service logs relevant to a failed inference. + + Used to surface the real server-side exception in the test's assertion message + when the orchestration pipeline returns the generic "technical issue" content + (llmServiceActive=False), so we don't depend on separately retrieving CI logs. + """ + try: + result = subprocess.run( + ["docker", "logs", "llm-orchestration-service", "--tail", str(tail)], + capture_output=True, + text=True, + timeout=15, + ) + except Exception as e: # pragma: no cover - diagnostics only + return f"" + + combined = f"{result.stdout}\n{result.stderr}" + markers = ( + "error", + "exception", + "traceback", + "rag_response_generation", + "refine", + "response generator", + "openai", + "azure", + "vault", + "deployment", + "not found", + "401", + "404", + "429", + ) + lines = combined.splitlines() + relevant = [line for line in lines if any(m in line.lower() for m in markers)] + if relevant: + # Keep the tail of the relevant lines (closest to the failure) + return "\n".join(relevant[-80:]) + # Fall back to the raw tail so we always surface something actionable + return "\n".join(lines[-60:]) or "" + + class TestInference: """Test LLM inference pipeline via Ruuter endpoints.""" @@ -72,7 +116,8 @@ def test_testing_inference_basic( logger.info(f"Testing inference with message: {test_case['question']}") logger.info( - f"Expected vault path: llm/connections/azure_openai/test/{connection_id}" + "Vault secret resolved via the connection's vault_uuid " + "(path: llm/connections/azure_openai/)" ) logger.info(f"Using payload: {json.dumps(payload)}") logger.info(f"Ruuter base URL: {ruuter_private_client.base_url}") @@ -96,7 +141,12 @@ def test_testing_inference_basic( assert "inputGuardFailed" in data assert "content" in data - assert data["llmServiceActive"] is True + assert data["llmServiceActive"] is True, ( + "Orchestration returned llmServiceActive=False (technical issue). " + f"Response content: {data.get('content')!r}\n" + "--- llm-orchestration-service error logs ---\n" + f"{_capture_orchestration_error_logs()}" + ) assert len(data["content"]) > 0 logger.info(f"Inference successful: {data['content'][:100]}...") diff --git a/tests/test_agentic_loop.py b/tests/test_agentic_loop.py index b8903788..d52fb3ce 100644 --- a/tests/test_agentic_loop.py +++ b/tests/test_agentic_loop.py @@ -5,9 +5,9 @@ import pytest -from src.tool_classifier.agentic_loop import AgenticLoop -from src.tool_classifier.enums import AgenticLoopStatus -from src.tool_classifier.param_extractor import ParamExtractionResult +from tool_classifier.agentic_loop import AgenticLoop +from tool_classifier.enums import AgenticLoopStatus +from tool_classifier.param_extractor import ParamExtractionResult # --------------------------------------------------------------------------- @@ -540,6 +540,7 @@ async def fake_to_thread(fn: Any, *args: Any, **kwargs: Any) -> Any: _HISTORY, {"validFrom": "2026-01-01"}, "en", + 1, ) @@ -963,3 +964,134 @@ async def test_user_exit_during_stream_returns_empty_tokens(self) -> None: assert tokens == [] # Collected params returned unchanged on exit assert result.collected_params == {"validFrom": "2026-01-01"} + + +# --------------------------------------------------------------------------- +# seeded_params — L2 param_update pre-population at turn 0 +# --------------------------------------------------------------------------- + + +class TestSeededParamsTurn0: + """Verify that seeded_params from L2 follow-up routing are merged into + collected_params at turn 0 only, with collected_params taking priority.""" + + @pytest.mark.asyncio + async def test_seeded_params_merged_at_turn_0(self) -> None: + """seeded_params are prepended to collected_params when turn_count=0.""" + # Extractor returns only validFrom as newly extracted; countryIsoCode comes + # from seeded_params. + extractor_mock = _make_extractor_mock( + _extraction( + {"validFrom": "2026-01-01"}, + [], # nothing missing — both params will be present after seed merge + "none", + ) + ) + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="January 2026", + conversation_history=[], + params_schema=_SCHEMA_TWO_REQUIRED, + collected_params={}, + turn_count=0, + seeded_params={"countryIsoCode": "EE"}, + ) + + # Both params present → COMPLETED + assert result.status == AgenticLoopStatus.COMPLETED + assert result.collected_params.get("countryIsoCode") == "EE" + assert result.collected_params.get("validFrom") == "2026-01-01" + + @pytest.mark.asyncio + async def test_collected_params_override_seeded_params(self) -> None: + """collected_params values beat seeded_params when the key overlaps.""" + extractor_mock = _make_extractor_mock( + _extraction( + {"validFrom": "2026-06-01"}, + [], + "none", + ) + ) + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="June 2026", + conversation_history=[], + params_schema=_SCHEMA_TWO_REQUIRED, + collected_params={"countryIsoCode": "LV"}, # explicit value takes priority + turn_count=0, + seeded_params={"countryIsoCode": "EE"}, # seeded value must be overridden + ) + + assert result.collected_params.get("countryIsoCode") == "LV" + + @pytest.mark.asyncio + async def test_seeded_params_not_applied_on_subsequent_turns(self) -> None: + """seeded_params are ignored when turn_count > 0.""" + extractor_mock = _make_extractor_mock( + _extraction({}, ["countryIsoCode", "validFrom"], "Which country and date?") + ) + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hello", + conversation_history=[], + params_schema=_SCHEMA_TWO_REQUIRED, + collected_params={}, + turn_count=1, # NOT turn 0 → seeded_params must be ignored + seeded_params={"countryIsoCode": "EE", "validFrom": "2026-01-01"}, + ) + + # Even though seeded_params would satisfy all required params, they should + # not be applied after turn 0 → still NEEDS_INPUT + assert result.status == AgenticLoopStatus.NEEDS_INPUT + # seeded values not present in collected_params + assert "countryIsoCode" not in result.collected_params + assert "validFrom" not in result.collected_params + + @pytest.mark.asyncio + async def test_seeded_params_none_does_not_raise(self) -> None: + """Passing seeded_params=None (default) at turn 0 behaves normally.""" + extractor_mock = _make_extractor_mock( + _extraction({}, ["countryIsoCode", "validFrom"], "Which country?") + ) + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hello", + conversation_history=[], + params_schema=_SCHEMA_TWO_REQUIRED, + collected_params={}, + turn_count=0, + seeded_params=None, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + @pytest.mark.asyncio + async def test_seeded_params_partial_fill_still_asks_for_missing(self) -> None: + """seeded_params satisfy only one of two required params → still NEEDS_INPUT.""" + extractor_mock = _make_extractor_mock( + _extraction({}, ["validFrom"], "From which date?") + ) + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="Estonia", + conversation_history=[], + params_schema=_SCHEMA_TWO_REQUIRED, + collected_params={}, + turn_count=0, + seeded_params={"countryIsoCode": "EE"}, # only one param seeded + ) + + # validFrom still missing → NEEDS_INPUT + assert result.status == AgenticLoopStatus.NEEDS_INPUT + # But seeded countryIsoCode should be present + assert result.collected_params.get("countryIsoCode") == "EE" diff --git a/tests/test_api_caller.py b/tests/test_api_caller.py index a6713fdb..de7a535c 100644 --- a/tests/test_api_caller.py +++ b/tests/test_api_caller.py @@ -7,8 +7,8 @@ import httpx import pytest -from src.tool_classifier.api_caller import APICaller, CircuitBreaker -from src.tool_classifier.constants import ( +from tool_classifier.api_caller import APICaller, CircuitBreaker +from tool_classifier.constants import ( CB_STATE_CLOSED, CB_STATE_HALF_OPEN, CB_STATE_OPEN, @@ -16,7 +16,7 @@ SERVICE_TIMEOUT_MESSAGES, SERVICE_UNAVAILABLE_MESSAGES, ) -from src.tool_classifier.models import APICallResult +from tool_classifier.models import APICallResult # --------------------------------------------------------------------------- diff --git a/tests/test_api_response_formatter.py b/tests/test_api_response_formatter.py index 752c57a2..05fe683f 100644 --- a/tests/test_api_response_formatter.py +++ b/tests/test_api_response_formatter.py @@ -9,7 +9,7 @@ import dspy.streaming import pytest -from src.tool_classifier.api_response_formatter import ( +from tool_classifier.api_response_formatter import ( APIResponseFormatterModule, _FORMATTER_ERROR_MESSAGES, ) diff --git a/tests/test_api_semantic_searcher.py b/tests/test_api_semantic_searcher.py index 15711269..e351d032 100644 --- a/tests/test_api_semantic_searcher.py +++ b/tests/test_api_semantic_searcher.py @@ -20,6 +20,7 @@ from tool_classifier.api_semantic_searcher import ( APISemanticSearcher, + DisambiguationResult, EndpointDisambiguatorModule, ) from tool_classifier.constants import ( @@ -160,7 +161,7 @@ def test_returns_winning_endpoint_id(self) -> None: }, ] result = module.forward("When is the next holiday?", candidates) - assert result == "ep-holidays" + assert result.winner_id == "ep-holidays" def test_returns_none_when_predictor_returns_none_string(self) -> None: module = self._make_module("none") @@ -173,7 +174,7 @@ def test_returns_none_when_predictor_returns_none_string(self) -> None: }, ] result = module.forward("Tell me a joke", candidates) - assert result is None + assert result.winner_id is None def test_returns_none_when_predictor_returns_none_uppercase(self) -> None: module = self._make_module("NONE") @@ -186,7 +187,7 @@ def test_returns_none_when_predictor_returns_none_uppercase(self) -> None: }, ] result = module.forward("Tell me a joke", candidates) - assert result is None + assert result.winner_id is None def test_returns_none_on_dspy_exception(self) -> None: module = EndpointDisambiguatorModule() @@ -200,7 +201,7 @@ def test_returns_none_on_dspy_exception(self) -> None: }, ] result = module.forward("Which holidays?", candidates) - assert result is None + assert result.winner_id is None def test_strips_whitespace_from_endpoint_id(self) -> None: module = self._make_module(" ep-holidays ") @@ -213,7 +214,7 @@ def test_strips_whitespace_from_endpoint_id(self) -> None: }, ] result = module.forward("Holidays?", candidates) - assert result == "ep-holidays" + assert result.winner_id == "ep-holidays" # --------------------------------------------------------------------------- @@ -448,17 +449,12 @@ async def test_multiple_medium_triggers_disambiguation(self) -> None: client.get = AsyncMock(return_value=count_resp) client.post = AsyncMock(side_effect=[dense_resp, hybrid_resp]) - mock_disambiguator = MagicMock() - mock_disambiguator.return_value = None # "forward" returns string or None - # Wrap in a module-like object that has a forward() callable via __call__ - mock_disambiguator_module = MagicMock() - mock_disambiguator_module.forward = MagicMock(return_value="ep-holidays") - mock_disambiguator_module.__call__ = MagicMock(return_value="ep-holidays") - # Inject our disambiguator — searcher calls self._disambiguator(query, candidates) # which in turn calls forward() via __call__ async_disambiguator = MagicMock() - async_disambiguator.forward = MagicMock(return_value="ep-holidays") + async_disambiguator.forward = MagicMock( + return_value=DisambiguationResult(winner_id="ep-holidays") + ) searcher = _make_searcher(client, disambiguator=async_disambiguator) @@ -466,7 +462,7 @@ async def test_multiple_medium_triggers_disambiguation(self) -> None: with patch( "tool_classifier.api_semantic_searcher.asyncio.to_thread", new_callable=AsyncMock, - return_value="ep-holidays", + return_value=DisambiguationResult(winner_id="ep-holidays"), ): results = await searcher.search("something ambiguous") @@ -507,6 +503,66 @@ async def test_disambiguation_rejects_all_returns_empty(self) -> None: assert results == [] + @pytest.mark.asyncio + async def test_disambiguation_rejects_all_multi_candidates_returns_top_with_hint( + self, + ) -> None: + """Disambiguator returns None for >1 medium candidates → top candidate returned + with multi_intent_hint=True and llm_validated=False so IntentDecomposer gate + can run in the classifier.""" + cos_a = API_TOOL_MIN_THRESHOLD + 0.08 # higher cosine → becomes 'top' + cos_b = API_TOOL_MIN_THRESHOLD + 0.02 + + dense_points = [ + _point({**_EP_HOLIDAYS}, cos_a), + _point({**_EP_WEATHER}, cos_b), + ] + hybrid_points = [ + _point({**_EP_HOLIDAYS}, 0.012), + _point({**_EP_WEATHER}, 0.009), + ] + + dense_resp = _make_qdrant_dense_response(dense_points) + hybrid_resp = _make_qdrant_hybrid_response(hybrid_points) + count_resp = _make_count_response(10) + + client = AsyncMock() + client.get = AsyncMock(return_value=count_resp) + client.post = AsyncMock(side_effect=[dense_resp, hybrid_resp]) + + searcher = _make_searcher(client) + + # asyncio.to_thread is called twice: + # 1st call → _get_query_embedding → must return a valid embedding vector + # 2nd call → disambiguator.forward → must return None ("none" response) + precomputed = [0.1] * 10 + _call_count = 0 + + async def _to_thread_side_effect(fn: Any, *args: Any, **kwargs: Any) -> Any: + nonlocal _call_count + _call_count += 1 + if _call_count == 1: + return precomputed # embedding call + return DisambiguationResult( + winner_id=None + ) # disambiguator call → rejects all candidates + + with patch( + "tool_classifier.api_semantic_searcher.asyncio.to_thread", + side_effect=_to_thread_side_effect, + ): + results = await searcher.search("holidays AND weather") + + # Must return exactly one result — the top cosine candidate + assert len(results) == 1 + top = results[0] + # Top candidate by cosine score is ep-holidays + assert top.endpoint_id == "ep-holidays" + # NOT llm_validated — disambiguator explicitly rejected it + assert top.llm_validated is False + # multi_intent_hint signals the classifier to try IntentDecomposer + assert top.multi_intent_hint is True + class TestSearchBelowThreshold: @pytest.mark.asyncio diff --git a/tests/test_api_tool_session_store.py b/tests/test_api_tool_session_store.py index eb54a2fe..9b877e0b 100644 --- a/tests/test_api_tool_session_store.py +++ b/tests/test_api_tool_session_store.py @@ -5,8 +5,8 @@ import pytest from pydantic import ValidationError -from src.models.session_models import APIToolSession -from src.utils.api_tool_session_store import ( +from models.session_models import APIToolSession, EndpointSessionState +from utils.api_tool_session_store import ( APIToolSessionStore, _key, require_session_store, @@ -46,6 +46,36 @@ def _make_redis_mock() -> AsyncMock: # --------------------------------------------------------------------------- +class TestEndpointSessionState: + def test_defaults(self): + ep = EndpointSessionState( + endpoint={"name": "get_holidays", "url": "https://example.com"} + ) + assert ep.collected_params == {} + assert ep.completed is False + + def test_mark_completed(self): + ep = EndpointSessionState( + endpoint={"name": "get_holidays"}, + collected_params={"year": "2026"}, + completed=True, + ) + assert ep.completed is True + assert ep.collected_params == {"year": "2026"} + + def test_serialization_roundtrip(self): + ep = EndpointSessionState( + endpoint={"name": "get_electricity_prices", "url": "https://example.com"}, + collected_params={"region": "EE"}, + ) + restored = EndpointSessionState.model_validate_json(ep.model_dump_json()) + assert restored == ep + + def test_endpoint_is_required(self): + with pytest.raises(ValidationError): + EndpointSessionState() # type: ignore[call-arg] + + class TestAPIToolSession: def test_defaults(self): session = APIToolSession(chat_id="abc", state="collecting_params") @@ -53,6 +83,10 @@ def test_defaults(self): assert session.turn_count == 0 assert session.max_turns == 5 assert session.selected_endpoint is None + # Phase 2 parallel-mode defaults + assert session.execution_mode == "single" + assert session.parallel_endpoints == [] + assert session.active_endpoint_index == 0 def test_serialization_roundtrip(self): session = _make_session() @@ -68,6 +102,75 @@ def test_max_turns_must_be_at_least_one(self): with pytest.raises(ValidationError): APIToolSession(chat_id="x", state="s", max_turns=0) + # ── Phase 2: parallel-mode fields ──────────────────────────────────── + + def test_parallel_session_stores_endpoint_states(self): + ep1 = EndpointSessionState(endpoint={"name": "get_holidays"}) + ep2 = EndpointSessionState(endpoint={"name": "get_electricity_prices"}) + session = APIToolSession( + chat_id="chat-p1", + state="collecting_params", + execution_mode="parallel", + parallel_endpoints=[ep1, ep2], + ) + assert session.execution_mode == "parallel" + assert len(session.parallel_endpoints) == 2 + assert session.parallel_endpoints[0].endpoint["name"] == "get_holidays" + assert ( + session.parallel_endpoints[1].endpoint["name"] == "get_electricity_prices" + ) + + def test_active_endpoint_index_defaults_to_zero(self): + session = APIToolSession( + chat_id="chat-p2", + state="collecting_params", + execution_mode="parallel", + parallel_endpoints=[EndpointSessionState(endpoint={"name": "ep1"})], + ) + assert session.active_endpoint_index == 0 + + def test_active_endpoint_index_must_be_non_negative(self): + with pytest.raises(ValidationError): + APIToolSession( + chat_id="x", + state="s", + active_endpoint_index=-1, + ) + + def test_parallel_session_serialization_roundtrip(self): + ep1 = EndpointSessionState( + endpoint={"name": "get_holidays", "url": "https://example.com"}, + collected_params={"year": "2026"}, + ) + ep2 = EndpointSessionState( + endpoint={"name": "get_electricity_prices", "url": "https://example.com"}, + completed=True, + ) + session = APIToolSession( + chat_id="chat-parallel", + state="collecting_params", + execution_mode="parallel", + parallel_endpoints=[ep1, ep2], + active_endpoint_index=1, + ) + restored = APIToolSession.model_validate_json(session.model_dump_json()) + assert restored == session + assert restored.parallel_endpoints[1].completed is True + assert restored.active_endpoint_index == 1 + + def test_backward_compat_session_without_parallel_fields(self): + """Sessions serialised before Phase 2 (no parallel fields) load cleanly.""" + legacy_json = ( + '{"chat_id":"legacy","state":"collecting_params",' + '"selected_endpoint":null,"collected_params":{},' + '"turn_count":0,"max_turns":5,"awaiting_continuation":false,' + '"detected_language":"en","original_query":""}' + ) + session = APIToolSession.model_validate_json(legacy_json) + assert session.execution_mode == "single" + assert session.parallel_endpoints == [] + assert session.active_endpoint_index == 0 + # --------------------------------------------------------------------------- # APIToolSessionStore.get @@ -82,7 +185,7 @@ async def test_get_returns_none_when_key_missing(self): redis_mock.get = AsyncMock(return_value=None) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.get("missing-chat") @@ -96,7 +199,7 @@ async def test_get_returns_session_when_key_exists(self): redis_mock.get = AsyncMock(return_value=session.model_dump_json()) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.get(session.chat_id) @@ -107,9 +210,7 @@ async def test_get_returns_session_when_key_exists(self): @pytest.mark.asyncio async def test_get_returns_none_when_redis_unavailable(self): store = APIToolSessionStore() - with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=None - ): + with patch("utils.api_tool_session_store.get_redis_client", return_value=None): result = await store.get("any-chat") assert result is None @@ -128,7 +229,7 @@ async def test_save_calls_redis_set_with_correct_key_and_ttl(self): redis_mock = _make_redis_mock() with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): await store.save(session) @@ -142,9 +243,7 @@ async def test_save_skips_when_redis_unavailable(self): store = APIToolSessionStore() session = _make_session() - with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=None - ): + with patch("utils.api_tool_session_store.get_redis_client", return_value=None): # Should not raise await store.save(session) @@ -181,7 +280,7 @@ async def test_update_merges_fields_and_resets_ttl(self): redis_mock.pipeline = MagicMock(return_value=pipe_mock) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.update( original.chat_id, @@ -210,7 +309,7 @@ async def test_update_returns_none_when_session_missing(self): redis_mock.pipeline = MagicMock(return_value=pipe_mock) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.update("ghost-chat", turn_count=3) @@ -235,7 +334,7 @@ async def test_update_resets_ttl(self): redis_mock.pipeline = MagicMock(return_value=pipe_mock) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): await store.update(session.chat_id, state="ready") @@ -257,7 +356,7 @@ async def test_delete_calls_redis_delete(self): redis_mock = _make_redis_mock() with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): await store.delete("chat-to-delete") @@ -266,9 +365,7 @@ async def test_delete_calls_redis_delete(self): @pytest.mark.asyncio async def test_delete_skips_when_redis_unavailable(self): store = APIToolSessionStore() - with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=None - ): + with patch("utils.api_tool_session_store.get_redis_client", return_value=None): await store.delete("any-chat") # Should not raise @@ -285,7 +382,7 @@ async def test_exists_returns_true_when_key_present(self): redis_mock.exists = AsyncMock(return_value=1) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.exists("chat-123") @@ -298,7 +395,7 @@ async def test_exists_returns_false_when_key_absent(self): redis_mock.exists = AsyncMock(return_value=0) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.exists("chat-123") @@ -307,9 +404,7 @@ async def test_exists_returns_false_when_key_absent(self): @pytest.mark.asyncio async def test_exists_returns_false_when_redis_unavailable(self): store = APIToolSessionStore() - with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=None - ): + with patch("utils.api_tool_session_store.get_redis_client", return_value=None): result = await store.exists("any-chat") assert result is False @@ -328,7 +423,7 @@ async def test_get_returns_none_on_redis_error(self): redis_mock.get = AsyncMock(side_effect=ConnectionError("timeout")) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.get("chat-xyz") @@ -342,7 +437,7 @@ async def test_save_does_not_raise_on_redis_error(self): redis_mock.set = AsyncMock(side_effect=ConnectionError("timeout")) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): await store.save(session) # Should not raise @@ -353,7 +448,7 @@ async def test_delete_does_not_raise_on_redis_error(self): redis_mock.delete = AsyncMock(side_effect=ConnectionError("timeout")) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): await store.delete("chat-xyz") # Should not raise @@ -364,7 +459,7 @@ async def test_exists_returns_false_on_redis_error(self): redis_mock.exists = AsyncMock(side_effect=ConnectionError("timeout")) with patch( - "src.utils.api_tool_session_store.get_redis_client", return_value=redis_mock + "utils.api_tool_session_store.get_redis_client", return_value=redis_mock ): result = await store.exists("chat-xyz") diff --git a/tests/test_api_tool_workflow.py b/tests/test_api_tool_workflow.py index 6fc640bb..e3275cdd 100644 --- a/tests/test_api_tool_workflow.py +++ b/tests/test_api_tool_workflow.py @@ -135,6 +135,10 @@ def _format_sse(chat_id: str, content: str) -> str: return f'data: {{"chatId":"{chat_id}","payload":{{"content":"{content}"}}}}\n\n' svc.format_sse = _format_sse + svc.store_streaming_inference = AsyncMock() + svc.handle_output_guardrails = AsyncMock( + side_effect=lambda _adapter, response, _req, _costs: response + ) return svc @@ -494,19 +498,25 @@ async def _fake_stream(**kwargs: Any) -> AsyncIterator[str]: for token in ["Holiday", " info", " here."]: yield token - executor._formatter.stream_forward = _fake_stream + mock_formatter = MagicMock() + mock_formatter.stream_forward = _fake_stream - frames = [ - frame - async for frame in executor._stream_api_and_format( - chat_id=_CHAT_ID, - endpoint=_ENDPOINT_HOLIDAYS, - collected_params={"countryIsoCode": "EE"}, - user_query="holidays", - detected_language="en", - orchestration_service=svc, - ) - ] + with patch( + "tool_classifier.workflows.api_tool_workflow.APIResponseFormatterModule", + return_value=mock_formatter, + ): + frames = [ + frame + async for frame in executor._stream_api_and_format( + chat_id=_CHAT_ID, + endpoint=_ENDPOINT_HOLIDAYS, + collected_params={"countryIsoCode": "EE"}, + user_query="holidays", + detected_language="en", + orchestration_service=svc, + request=_make_request(), + ) + ] # 3 token frames + 1 END frame assert len(frames) == 4 @@ -536,6 +546,7 @@ async def test_api_failure_streams_error_frame(self) -> None: user_query="holidays", detected_language="et", orchestration_service=svc, + request=_make_request(), ) ] diff --git a/tests/test_api_tool_workflow_integration.py b/tests/test_api_tool_workflow_integration.py index 4d97fd81..fda6d241 100644 --- a/tests/test_api_tool_workflow_integration.py +++ b/tests/test_api_tool_workflow_integration.py @@ -6,6 +6,8 @@ Covers: - Phase 2: Full multi-turn workflow, fast-path, streaming, cost tracking - Phase 4: Fallback chain regression tests +- Parallel execution mode (ExecutionMode.PARALLEL end-to-end) +- Test-endpoint session wipe guard """ import json @@ -17,13 +19,14 @@ import pytest from models.request_models import OrchestrationRequest, OrchestrationResponse -from models.session_models import APIToolSession +from models.session_models import APIToolSession, EndpointSessionState from tool_classifier.classifier import ToolClassifier -from tool_classifier.enums import AgenticLoopStatus, WorkflowType +from tool_classifier.enums import AgenticLoopStatus, ExecutionMode, WorkflowType from tool_classifier.models import ( AgenticLoopResult, APICallResult, ClassificationResult, + MultiAPICallResult, ) @@ -127,6 +130,10 @@ async def _mock_rag(**kwargs: Any) -> OrchestrationResponse: svc._execute_orchestration_pipeline = AsyncMock(side_effect=_mock_rag) svc._initialize_service_components = MagicMock(return_value={}) + svc.handle_output_guardrails = AsyncMock( + side_effect=lambda _adapter, response, _req, _costs: response + ) + svc.store_streaming_inference = AsyncMock() async def _mock_rag_stream(**kwargs: Any) -> AsyncGenerator[str, None]: yield 'data: {"chatId":"test","payload":{"content":"RAG stream answer"}}\n\n' @@ -484,6 +491,9 @@ async def _fake_stream_forward(**kwargs: Any) -> AsyncIterator[str]: for token in ["It is ", "15°C ", "in Tallinn."]: yield token + mock_formatter = MagicMock() + mock_formatter.stream_forward = _fake_stream_forward + with ( patch.object( classifier.api_tool_workflow._api_caller, @@ -491,10 +501,9 @@ async def _fake_stream_forward(**kwargs: Any) -> AsyncIterator[str]: new_callable=AsyncMock, return_value=api_call_result, ), - patch.object( - classifier.api_tool_workflow._formatter, - "stream_forward", - side_effect=_fake_stream_forward, + patch( + "tool_classifier.workflows.api_tool_workflow.APIResponseFormatterModule", + return_value=mock_formatter, ), ): stream = await classifier.route_to_workflow( @@ -896,3 +905,323 @@ def _make_mock_loop( loop = MagicMock() loop.stream_run_turn = AsyncMock(return_value=(result, question_tokens)) return loop + + +# --------------------------------------------------------------------------- +# TestParallelExecutionMode +# --------------------------------------------------------------------------- + + +class TestParallelExecutionMode: + """Full parallel path: classify → ExecutionMode.PARALLEL → MultiEndpointAgenticLoop + → MultiAPICaller → MultiResponseFormatterModule. + + Only DSPy (formatter/extractor), Qdrant HTTP, and Redis are mocked. + """ + + @pytest.mark.asyncio + async def test_parallel_fast_path_no_required_params_both_apis_called( + self, + classifier: ToolClassifier, + mock_session_store: AsyncMock, + ) -> None: + """Both endpoints have no required params → immediate parallel API calls, no session.""" + + # Two endpoints with no required params + ep_weather_no_params = {**_ENDPOINT_WEATHER, "params": []} + ep_holidays_no_params = {**_ENDPOINT_HOLIDAYS, "params": []} + + classification = ClassificationResult( + workflow=WorkflowType.API_TOOL_CALLING, + confidence=0.68, + metadata={ + "execution_mode": ExecutionMode.PARALLEL, + "matched_endpoints": [ep_weather_no_params, ep_holidays_no_params], + }, + ) + request = _make_request("holidays AND weather") + + weather_result = APICallResult( + success=True, status_code=200, response_data={"temp": 22}, error=None + ) + holidays_result = APICallResult( + success=True, + status_code=200, + response_data={"holidays": ["Jõulupüha"]}, + error=None, + ) + multi_result = MultiAPICallResult( + results=[weather_result, holidays_result], + endpoints=[ + {**ep_weather_no_params, "call_params": {}}, + {**ep_holidays_no_params, "call_params": {}}, + ], + ) + + with ( + patch.object( + classifier.api_tool_workflow._api_caller.__class__, + "__init__", + return_value=None, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.MultiAPICaller", + ) as mock_multi_caller_cls, + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value="It is 22°C and there are public holidays.", + ), + ): + mock_multi_caller_inst = AsyncMock() + mock_multi_caller_inst.call_all = AsyncMock(return_value=multi_result) + mock_multi_caller_cls.return_value = mock_multi_caller_inst + + response = await classifier.route_to_workflow( + classification=classification, + request=request, + is_streaming=False, + ) + + assert isinstance(response, OrchestrationResponse) + assert response.content != "" + # No session created (fast path) + assert await mock_session_store.get(_CHAT_ID) is None + + @pytest.mark.asyncio + async def test_parallel_session_created_when_params_needed( + self, + classifier: ToolClassifier, + mock_session_store: AsyncMock, + ) -> None: + """When endpoints have required params, a PARALLEL session is created and a + clarifying question is returned for the first turn.""" + + # Both endpoints have required params + classification = ClassificationResult( + workflow=WorkflowType.API_TOOL_CALLING, + confidence=0.68, + metadata={ + "execution_mode": ExecutionMode.PARALLEL, + "matched_endpoints": [_ENDPOINT_HOLIDAYS, _ENDPOINT_WEATHER], + }, + ) + request = _make_request("holidays AND weather please") + + loop_result = AgenticLoopResult( + status=AgenticLoopStatus.NEEDS_INPUT, + collected_params={}, + clarifying_question="Which country for holidays?", + turn_count=1, + ) + + with patch( + "tool_classifier.workflows.api_tool_workflow.MultiEndpointAgenticLoop", + ) as mock_multi_loop_cls: + mock_multi_loop_inst = MagicMock() + mock_multi_loop_inst.stream_run_turn = AsyncMock( + return_value=(loop_result, ["Which", " country", "?"]) + ) + mock_multi_loop_cls.return_value = mock_multi_loop_inst + + response = await classifier.route_to_workflow( + classification=classification, + request=request, + is_streaming=False, + ) + + assert isinstance(response, OrchestrationResponse) + assert response.content != "" + # Session created in Redis with parallel execution mode + session = await mock_session_store.get(_CHAT_ID) + assert session is not None + assert session.execution_mode == ExecutionMode.PARALLEL.value + + @pytest.mark.asyncio + async def test_parallel_max_turns_falls_back_to_rag( + self, + classifier: ToolClassifier, + mock_session_store: AsyncMock, + ) -> None: + """Parallel loop MAX_TURNS_REACHED → session deleted → RAG fallback.""" + # Seed a parallel session + session = APIToolSession( + chat_id=_CHAT_ID, + state="collecting_params", + selected_endpoint=_ENDPOINT_HOLIDAYS, + collected_params={}, + turn_count=6, + max_turns=6, + awaiting_continuation=False, + detected_language="en", + original_query="holidays AND weather", + execution_mode=ExecutionMode.PARALLEL.value, + parallel_endpoints=[ + EndpointSessionState(endpoint=_ENDPOINT_HOLIDAYS), + EndpointSessionState(endpoint=_ENDPOINT_WEATHER), + ], + ) + await mock_session_store.save(session) + + classification = ClassificationResult( + workflow=WorkflowType.API_TOOL_CALLING, + confidence=1.0, + metadata={"reason": "active_session_resume"}, + ) + request = _make_request("I give up") + + max_turns_result = AgenticLoopResult( + status=AgenticLoopStatus.MAX_TURNS_REACHED, + collected_params={}, + clarifying_question="", + turn_count=7, + ) + + with patch( + "tool_classifier.workflows.api_tool_workflow.MultiEndpointAgenticLoop", + ) as mock_multi_loop_cls: + mock_multi_loop_inst = MagicMock() + mock_multi_loop_inst.stream_run_turn = AsyncMock( + return_value=(max_turns_result, []) + ) + mock_multi_loop_cls.return_value = mock_multi_loop_inst + + response = await classifier.route_to_workflow( + classification=classification, + request=request, + is_streaming=False, + ) + + # Session deleted + assert await mock_session_store.get(_CHAT_ID) is None + # Falls back to RAG + assert isinstance(response, OrchestrationResponse) + + @pytest.mark.asyncio + async def test_parallel_streaming_question_yields_sse_frames( + self, + classifier: ToolClassifier, + mock_session_store: AsyncMock, + ) -> None: + """Streaming parallel path: first turn → clarifying question → SSE frames.""" + + classification = ClassificationResult( + workflow=WorkflowType.API_TOOL_CALLING, + confidence=0.68, + metadata={ + "execution_mode": ExecutionMode.PARALLEL, + "matched_endpoints": [_ENDPOINT_HOLIDAYS, _ENDPOINT_WEATHER], + }, + ) + request = _make_request("holidays AND weather") + + loop_result = AgenticLoopResult( + status=AgenticLoopStatus.NEEDS_INPUT, + collected_params={}, + clarifying_question="Which country for holidays?", + turn_count=1, + ) + + with patch( + "tool_classifier.workflows.api_tool_workflow.MultiEndpointAgenticLoop", + ) as mock_multi_loop_cls: + mock_multi_loop_inst = MagicMock() + mock_multi_loop_inst.stream_run_turn = AsyncMock( + return_value=(loop_result, ["Which", " country", "?"]) + ) + mock_multi_loop_cls.return_value = mock_multi_loop_inst + + stream = await classifier.route_to_workflow( + classification=classification, + request=request, + is_streaming=True, + ) + frames = [frame async for frame in stream] + + assert len(frames) >= 1 + for frame in frames: + assert frame.startswith("data: ") or frame.strip() == "" + + +# --------------------------------------------------------------------------- +# TestTestEndpointSessionWipe +# --------------------------------------------------------------------------- + + +class TestTestEndpointSessionWipe: + """Verify that the /orchestrate/test endpoint deletes any stale 'test-session' + in the API tool session store before each request so multi-turn state never + leaks between consecutive test API calls.""" + + @pytest.mark.asyncio + async def test_session_store_delete_called_with_test_session_key(self) -> None: + """The endpoint must call session_store.delete('test-session') on every request + regardless of whether a session exists.""" + from httpx import AsyncClient, ASGITransport + from llm_orchestration_service_api import app + + session_store_mock = AsyncMock() + session_store_mock.delete = AsyncMock(return_value=None) + + # Minimal orchestration service mock that returns a valid response + orch_mock = AsyncMock() + orch_mock.process_orchestration_request = AsyncMock( + return_value=OrchestrationResponse( + chatId="test-session", + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content="Test answer.", + ) + ) + + app.state.orchestration_service = orch_mock + app.state.session_store = session_store_mock + + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + await client.post( + "/orchestrate/test", + json={"message": "hello", "environment": "production"}, + ) + + # The endpoint must have called delete("test-session") before processing + session_store_mock.delete.assert_awaited_with("test-session") + + @pytest.mark.asyncio + async def test_stale_session_cleared_before_request_not_after(self) -> None: + """If session_store.delete raises, the endpoint propagates the error as HTTP 500 + because the wipe-guard is not wrapped in try/except.""" + from httpx import AsyncClient, ASGITransport + from llm_orchestration_service_api import app + + session_store_mock = AsyncMock() + session_store_mock.delete = AsyncMock( + side_effect=RuntimeError("Redis unavailable") + ) + + orch_mock = AsyncMock() + orch_mock.process_orchestration_request = AsyncMock( + return_value=OrchestrationResponse( + chatId="test-session", + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content="Fallback answer.", + ) + ) + + app.state.orchestration_service = orch_mock + app.state.session_store = session_store_mock + + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + resp = await client.post( + "/orchestrate/test", + json={"message": "hello", "environment": "production"}, + ) + + # delete() raised → the unguarded await bubbles up as HTTP 500 + assert resp.status_code == 500 diff --git a/tests/test_atc_cache.py b/tests/test_atc_cache.py new file mode 100644 index 00000000..50ae95b4 --- /dev/null +++ b/tests/test_atc_cache.py @@ -0,0 +1,877 @@ +"""Unit tests for the ATC Response Cache integration — T-6 (tests 6–13 of the spec). + +Tests 1–5 (L1 hit, L1 miss, param normalisation, L2 round-trip, L2 invalidate) are +covered by ``test_atc_cache_store.py`` and are not duplicated here. + +This file exercises the eight cross-cutting scenarios: + + A) Workflow write: ``cacheable=False`` endpoint → no L1 write + B) ``_compute_loop_step``: L1 hit → ``cached_response`` step + C) ``_compute_loop_step``: L2 routing — four FollowUpDetector outcomes + D) ``ToolClassifier``: intent switch → ``invalidate_l2`` + E) Multi-intent write: one ``LastCallContext`` per succeeded endpoint +""" + +import asyncio +from typing import Any, Dict, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from models.request_models import OrchestrationRequest +from models.session_models import ( + APIToolSession, + EndpointSessionState, + LastCallContext, +) +from tool_classifier.classifier import ToolClassifier +from tool_classifier.enums import AgenticLoopStatus, WorkflowType +from tool_classifier.models import ( + AgenticLoopResult, + APICallResult, + MultiAPICallResult, +) +from tool_classifier.workflows.api_tool_workflow import APIToolWorkflowExecutor +from utils.atc_cache_store import ATCCacheStore + +# --------------------------------------------------------------------------- +# Shared constants +# --------------------------------------------------------------------------- + +_CHAT_ID = "atc-cache-test-001" +_API_NAME = "get_national_holidays" +_RESPONSE: Dict[str, Any] = { + "holidays": [{"date": "2026-02-24", "name": "Independence Day"}] +} +_PARAMS: Dict[str, Any] = {"year": 2026, "country": "EE"} + +# Endpoint with two required params (year + country) +_ENDPOINT: Dict[str, Any] = { + "endpoint_id": "ep-holidays", + "name": _API_NAME, + "description": "Returns public holidays for a country and year", + "method": "GET", + "url": "https://openholidaysapi.org/PublicHolidays", + "params": [ + {"name": "year", "type": "integer", "required": True, "description": "Year"}, + { + "name": "country", + "type": "string", + "required": True, + "description": "Country ISO", + }, + ], + "cacheable": True, + "cache_ttl_seconds": None, +} + +_ENDPOINT_NON_CACHEABLE: Dict[str, Any] = { + **_ENDPOINT, + "name": "get_doc_status", + "cacheable": False, +} + +# --------------------------------------------------------------------------- +# Shared helpers +# --------------------------------------------------------------------------- + + +def _make_last_call_ctx( + api_name: str = _API_NAME, + collected_params: Optional[Dict[str, Any]] = None, +) -> LastCallContext: + return LastCallContext( + api_name=api_name, + endpoint=_ENDPOINT, + collected_params=collected_params if collected_params is not None else _PARAMS, + raw_response=_RESPONSE, + original_query="What are the public holidays in Estonia in 2026?", + timestamp=1748000000.0, + ) + + +def _make_request( + message: str = "public holidays", + chat_id: str = _CHAT_ID, +) -> OrchestrationRequest: + return OrchestrationRequest( + chatId=chat_id, + message=message, + authorId="test-user", + conversationHistory=[], + url="https://example.com", + environment="testing", + connection_id="test-conn", + ) + + +def _make_session_store( + session: Optional[APIToolSession] = None, +) -> AsyncMock: + store = AsyncMock() + store.get = AsyncMock(return_value=session) + store.save = AsyncMock() + store.delete = AsyncMock() + store.update = AsyncMock() + return store + + +def _make_orchestration_service( + session_store: Optional[AsyncMock] = None, +) -> MagicMock: + svc = MagicMock() + svc.session_store = session_store + # None prevents _get_custom_instructions from calling asyncio.to_thread, which + # would otherwise consume the mocked return value before the FUD call. + svc.prompt_config_loader = None + + def _format_sse(chat_id: str, content: str, **kwargs: Any) -> str: + return f'data: {{"chatId":"{chat_id}","payload":{{"content":"{content}"}}}}\n\n' + + svc.format_sse = _format_sse + return svc + + +def _make_executor( + session_store: Optional[AsyncMock] = None, +) -> APIToolWorkflowExecutor: + return APIToolWorkflowExecutor( + orchestration_service=_make_orchestration_service(session_store) + ) + + +def _make_loop_result( + status: AgenticLoopStatus, + collected_params: Optional[Dict[str, Any]] = None, +) -> AgenticLoopResult: + return AgenticLoopResult( + status=status, + collected_params=collected_params or {}, + clarifying_question="Which year?", + turn_count=1, + ) + + +def _make_loop_mock() -> AsyncMock: + """Return a mock AgenticLoop that reports NEEDS_INPUT with a question.""" + loop_mock = AsyncMock() + loop_mock.stream_run_turn = AsyncMock( + return_value=( + _make_loop_result(AgenticLoopStatus.NEEDS_INPUT), + ["Which", " year?"], + ) + ) + return loop_mock + + +# ───────────────────────────────────────────────────────────────────────────── +# Group A — Workflow write path +# ───────────────────────────────────────────────────────────────────────────── + + +class TestCacheableFlag: + @pytest.mark.asyncio + async def test_l1_skipped_when_not_cacheable(self) -> None: + """cacheable=False endpoint → no L1 write after a successful API call.""" + executor = _make_executor() + api_result = APICallResult( + success=True, + status_code=200, + response_data=_RESPONSE, + error=None, + ) + executor._api_caller.call = AsyncMock(return_value=api_result) + + mock_cache_store = AsyncMock() + mock_cache_store.set_l1 = AsyncMock() + mock_cache_store.set_l2 = AsyncMock() + + with ( + patch( + "tool_classifier.workflows.api_tool_workflow.ATCCacheStore", + return_value=mock_cache_store, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value="Formatted response.", + ), + ): + await executor._execute_api_and_format( + chat_id=_CHAT_ID, + endpoint=_ENDPOINT_NON_CACHEABLE, + collected_params=_PARAMS, + user_query="document status?", + detected_language="en", + ) + # Flush any background tasks that might have been scheduled + await asyncio.sleep(0) + + # Neither L1 nor L2 may be written for non-cacheable endpoints + mock_cache_store.set_l1.assert_not_called() + mock_cache_store.set_l2.assert_not_called() + + +# ───────────────────────────────────────────────────────────────────────────── +# Group B — _compute_loop_step: L1 cache hit +# ───────────────────────────────────────────────────────────────────────────── + + +class TestComputeLoopStepL1Hit: + @pytest.mark.asyncio + async def test_cache_hit_in_compute_loop_step(self) -> None: + """L1 hit on a new (session-less) query → _LoopStep(kind='cached_response').""" + session_store = _make_session_store(session=None) + executor = _make_executor(session_store) + + with ( + patch.object( + ATCCacheStore, + "get_l1", + new=AsyncMock(return_value=_RESPONSE), + ), + patch.object(ATCCacheStore, "get_l2", new=AsyncMock(return_value=None)), + patch.object(ATCCacheStore, "set_l1", new=AsyncMock()), + patch.object(ATCCacheStore, "set_l2", new=AsyncMock()), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + True, + ), + ): + step = await executor._compute_loop_step( + _make_request(), + context={"matched_endpoint": _ENDPOINT}, + ) + + assert step.kind == "cached_response" + assert step.cache_source == "L1" + assert step.cached_raw_response == _RESPONSE + + +# ───────────────────────────────────────────────────────────────────────────── +# Group C — _compute_loop_step: L2 follow-up routing +# ───────────────────────────────────────────────────────────────────────────── + + +class TestComputeLoopStepL2Routing: + """L1 miss + L2 hit: four outcomes based on FollowUpDetectorModule result. + + All four tests share: + - No active session (session_store.get returns None) + - L1 miss (get_l1 returns None) + - L2 hit (get_l2 returns [ctx] where ctx.api_name matches endpoint name) + - asyncio.to_thread is patched to return the desired FUD result dict directly + """ + + @pytest.mark.asyncio + async def test_follow_up_detector_response_question(self) -> None: + """FUD → 'response_question': returns cached_response from L2 raw_response.""" + session_store = _make_session_store(session=None) + executor = _make_executor(session_store) + ctx = _make_last_call_ctx(collected_params=_PARAMS) + + fud_result = {"follow_up_type": "response_question", "updated_params": {}} + + with ( + patch.object(ATCCacheStore, "get_l1", new=AsyncMock(return_value=None)), + patch.object(ATCCacheStore, "get_l2", new=AsyncMock(return_value=[ctx])), + patch.object(ATCCacheStore, "set_l1", new=AsyncMock()), + patch.object(ATCCacheStore, "set_l2", new=AsyncMock()), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + True, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value=fud_result, + ), + ): + step = await executor._compute_loop_step( + _make_request("Which of those holidays falls on a Monday?"), + context={"matched_endpoint": _ENDPOINT}, + ) + + assert step.kind == "cached_response" + assert step.cache_source == "L2" + assert step.cached_raw_response == ctx.raw_response + + @pytest.mark.asyncio + async def test_follow_up_detector_param_update_complete(self) -> None: + """FUD → 'param_update'; all required params in merged result → api_call step. + + Previous call had only 'country'; FUD supplies 'year' → merged is complete. + _param_hash is NOT patched so the real inequality check runs correctly. + """ + session_store = _make_session_store(session=None) + executor = _make_executor(session_store) + + # Prior call was missing 'year'; the new query supplies it + ctx = _make_last_call_ctx(collected_params={"country": "EE"}) + fud_result = { + "follow_up_type": "param_update", + "updated_params": {"year": 2025}, + } + + with ( + patch.object(ATCCacheStore, "get_l1", new=AsyncMock(return_value=None)), + patch.object(ATCCacheStore, "get_l2", new=AsyncMock(return_value=[ctx])), + patch.object(ATCCacheStore, "set_l1", new=AsyncMock()), + patch.object(ATCCacheStore, "set_l2", new=AsyncMock()), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + True, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value=fud_result, + ), + ): + step = await executor._compute_loop_step( + _make_request("What about 2025 instead?"), + context={"matched_endpoint": _ENDPOINT}, + ) + + assert step.kind == "api_call" + assert step.collected_params == {"country": "EE", "year": 2025} + + @pytest.mark.asyncio + async def test_follow_up_detector_param_update_partial(self) -> None: + """FUD → 'param_update'; merged params still incomplete → seeded_params set. + + Previous call only has 'year'; no updated params → 'country' still missing. + The agentic loop is mocked so the test can verify the step kind. + """ + session_store = _make_session_store(session=None) + executor = _make_executor(session_store) + + ctx = _make_last_call_ctx(collected_params={"year": 2026}) + # Empty update → merged = {"year": 2026}, still missing "country" + fud_result = {"follow_up_type": "param_update", "updated_params": {}} + + loop_mock = _make_loop_mock() + executor._build_agentic_loop = MagicMock(return_value=loop_mock) + + context: Dict[str, Any] = {"matched_endpoint": _ENDPOINT} + + with ( + patch.object(ATCCacheStore, "get_l1", new=AsyncMock(return_value=None)), + patch.object(ATCCacheStore, "get_l2", new=AsyncMock(return_value=[ctx])), + patch.object(ATCCacheStore, "set_l1", new=AsyncMock()), + patch.object(ATCCacheStore, "set_l2", new=AsyncMock()), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + True, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value=fud_result, + ), + ): + step = await executor._compute_loop_step( + _make_request("Show me the holidays"), + context=context, + ) + + # seeded_params must be injected into the mutable context dict + assert "seeded_params" in context + assert context["seeded_params"] == {"year": 2026} + # Normal agentic loop ran after seeding → returns a question + assert step.kind == "question" + + @pytest.mark.asyncio + async def test_follow_up_detector_new_intent(self) -> None: + """FUD → 'new_intent': falls through to normal loop; seeded_params NOT set.""" + session_store = _make_session_store(session=None) + executor = _make_executor(session_store) + + ctx = _make_last_call_ctx(collected_params=_PARAMS) + fud_result = {"follow_up_type": "new_intent", "updated_params": {}} + + loop_mock = _make_loop_mock() + executor._build_agentic_loop = MagicMock(return_value=loop_mock) + + context: Dict[str, Any] = {"matched_endpoint": _ENDPOINT} + + with ( + patch.object(ATCCacheStore, "get_l1", new=AsyncMock(return_value=None)), + patch.object(ATCCacheStore, "get_l2", new=AsyncMock(return_value=[ctx])), + patch.object(ATCCacheStore, "set_l1", new=AsyncMock()), + patch.object(ATCCacheStore, "set_l2", new=AsyncMock()), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + True, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value=fud_result, + ), + ): + step = await executor._compute_loop_step( + _make_request("How do I apply for a driving licence?"), + context=context, + ) + + # new_intent path must NOT seed params + assert "seeded_params" not in context + # Normal agentic loop took over → question returned + assert step.kind == "question" + + +# ───────────────────────────────────────────────────────────────────────────── +# Group D — ToolClassifier: intent switch → invalidate_l2 +# ───────────────────────────────────────────────────────────────────────────── + +_ENDPOINT_HOLIDAYS_CLS: Dict[str, Any] = { + "endpoint_id": "ep-holidays", + "name": "get_public_holidays", + "description": "Returns public holidays", + "method": "GET", + "url": "https://openholidaysapi.org/PublicHolidays", + "params": [ + { + "name": "countryIsoCode", + "type": "string", + "required": True, + "description": "ISO", + } + ], + "cosine_score": 0.68, + "rrf_score": 0.01, + "confidence": "high", +} + +_ENDPOINT_WEATHER_CLS: Dict[str, Any] = { + "endpoint_id": "ep-weather", + "name": "get_weather", + "description": "Current weather for a city", + "method": "GET", + "url": "https://publicapi.envir.ee/v1/weather", + "params": [], + "cosine_score": 0.65, + "rrf_score": 0.009, + "confidence": "high", +} + + +def _make_api_tool_search_result(endpoint: Dict[str, Any]) -> MagicMock: + """Build a mock APIToolSearchResult from an endpoint dict.""" + result = MagicMock() + result.endpoint_id = endpoint["endpoint_id"] + result.name = endpoint["name"] + result.description = endpoint["description"] + result.method = endpoint["method"] + result.url = endpoint["url"] + result.params = endpoint.get("params", []) + result.cosine_score = endpoint.get("cosine_score", 0.65) + result.rrf_score = endpoint.get("rrf_score", 0.01) + result.confidence = endpoint.get("confidence", "high") + result.to_dict.return_value = endpoint + return result + + +def _make_classifier_instance( + session_store: Optional[AsyncMock] = None, +) -> ToolClassifier: + """Instantiate ToolClassifier with all external I/O patched away.""" + svc = MagicMock() + svc.create_embeddings_for_indexer.return_value = {"embeddings": [[0.1] * 10]} + svc.session_store = session_store + + with patch("tool_classifier.classifier.httpx.AsyncClient"): + return ToolClassifier( + llm_manager=MagicMock(), + orchestration_service=svc, + ) + + +class TestIntentSwitchInvalidatesL2: + @pytest.mark.asyncio + async def test_intent_switch_invalidates_l2(self) -> None: + """Intent switch during session resume → ATCCacheStore.invalidate_l2 called.""" + # Active session for endpoint A (holidays) + active_session = APIToolSession( + chat_id=_CHAT_ID, + state="collecting_params", + selected_endpoint=_ENDPOINT_HOLIDAYS_CLS, + collected_params={}, + turn_count=1, + max_turns=5, + awaiting_continuation=False, + detected_language="en", + original_query="holiday query", + ) + session_store = _make_session_store(session=active_session) + + classifier = _make_classifier_instance(session_store) + # New query matches endpoint B (weather) — a different endpoint + classifier.api_tool_searcher.search = AsyncMock( + return_value=[_make_api_tool_search_result(_ENDPOINT_WEATHER_CLS)] + ) + + mock_cache_store = AsyncMock() + mock_cache_store.invalidate_l2 = AsyncMock() + + with ( + patch( + "tool_classifier.classifier.FeatureFlags" + ".API_TOOL_CALLING_WORKFLOW_ENABLED", + True, + ), + patch( + "tool_classifier.classifier.FeatureFlags.ATC_RESPONSE_CACHE_ENABLED", + True, + ), + patch( + "tool_classifier.classifier.ATCCacheStore", + return_value=mock_cache_store, + ), + ): + result = await classifier.classify( + query="What is the weather in Tallinn?", + conversation_history=[], + language="en", + request=_make_request("What is the weather in Tallinn?"), + ) + + # Both L2 invalidation and session delete must fire on intent switch + mock_cache_store.invalidate_l2.assert_called_once_with(_CHAT_ID) + session_store.delete.assert_called_once_with(_CHAT_ID) + assert result.workflow == WorkflowType.API_TOOL_CALLING + assert result.metadata.get("matched_endpoint", {}).get("name") == "get_weather" + + +# ───────────────────────────────────────────────────────────────────────────── +# Group E — Multi-intent write: one LastCallContext per succeeded endpoint +# ───────────────────────────────────────────────────────────────────────────── + + +class TestMultiIntentCacheWrite: + @pytest.mark.asyncio + async def test_multi_intent_writes_multiple_l2_entries(self) -> None: + """_execute_multi_api_and_format stores one LastCallContext per succeeded endpoint.""" + executor = _make_executor() + + ep1: Dict[str, Any] = { + "name": "get_national_holidays", + "description": "Returns holidays", + "method": "GET", + "url": "https://example.com/holidays", + "cacheable": True, + "cache_ttl_seconds": None, + } + ep2: Dict[str, Any] = { + "name": "get_electricity_prices", + "description": "Returns electricity prices", + "method": "GET", + "url": "https://example.com/electricity", + "cacheable": True, + "cache_ttl_seconds": None, + } + + parallel_endpoints = [ + EndpointSessionState( + endpoint=ep1, collected_params={"country": "EE"}, completed=True + ), + EndpointSessionState( + endpoint=ep2, collected_params={"date": "2026-05-28"}, completed=True + ), + ] + + multi_result = MultiAPICallResult( + results=[ + APICallResult( + success=True, status_code=200, response_data={"h": 1}, error=None + ), + APICallResult( + success=True, status_code=200, response_data={"p": 2}, error=None + ), + ], + endpoints=[ep1, ep2], + ) + + mock_multi_caller = AsyncMock() + mock_multi_caller.call_all = AsyncMock(return_value=multi_result) + + mock_cache_store = AsyncMock() + mock_cache_store.set_l1 = AsyncMock() + mock_cache_store.set_l2 = AsyncMock() + + with ( + patch( + "tool_classifier.workflows.api_tool_workflow.MultiAPICaller", + return_value=mock_multi_caller, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.ATCCacheStore", + return_value=mock_cache_store, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value="Multi-endpoint formatted response.", + ), + ): + await executor._execute_multi_api_and_format( + chat_id=_CHAT_ID, + parallel_endpoints=parallel_endpoints, + user_query="Show me holidays and electricity prices", + detected_language="en", + ) + # Flush the background _write_multi_cache() asyncio.create_task + await asyncio.sleep(0) + + # set_l2 must be called once with both endpoint contexts + mock_cache_store.set_l2.assert_called_once() + _call_args = mock_cache_store.set_l2.call_args + contexts: list[LastCallContext] = _call_args[0][1] + assert len(contexts) == 2 + api_names = {c.api_name for c in contexts} + assert api_names == {"get_national_holidays", "get_electricity_prices"} + + +# ───────────────────────────────────────────────────────────────────────────── +# Group F — Kill-switch: ATC_RESPONSE_CACHE_ENABLED=false disables all caching +# ───────────────────────────────────────────────────────────────────────────── + + +class TestCacheKillSwitch: + """Verify that ATC_RESPONSE_CACHE_ENABLED=false disables all cache operations.""" + + @pytest.mark.asyncio + async def test_kill_switch_disables_l1_check_in_execute_api(self) -> None: + """With cache disabled, _execute_api_and_format does NOT check L1.""" + executor = _make_executor() + api_result = APICallResult( + success=True, + status_code=200, + response_data=_RESPONSE, + error=None, + ) + executor._api_caller.call = AsyncMock(return_value=api_result) + + mock_cache_store = AsyncMock() + mock_cache_store.get_l1 = AsyncMock() # Should NOT be called + mock_cache_store.set_l1 = AsyncMock() # Should NOT be called + mock_cache_store.set_l2 = AsyncMock() # Should NOT be called + + with ( + patch( + "tool_classifier.workflows.api_tool_workflow.ATCCacheStore", + return_value=mock_cache_store, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value="Formatted response.", + ), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + False, # KILL SWITCH: disabled + ), + ): + await executor._execute_api_and_format( + chat_id=_CHAT_ID, + endpoint=_ENDPOINT, + collected_params=_PARAMS, + user_query="What are the public holidays?", + detected_language="en", + ) + # Flush any background tasks + await asyncio.sleep(0) + + # All cache methods must be skipped when kill-switch is off + mock_cache_store.get_l1.assert_not_called() + mock_cache_store.set_l1.assert_not_called() + mock_cache_store.set_l2.assert_not_called() + + @pytest.mark.asyncio + async def test_kill_switch_disables_multi_cache_write(self) -> None: + """With cache disabled, _execute_multi_api_and_format background write is skipped.""" + executor = _make_executor() + + ep1: Dict[str, Any] = { + "name": "get_national_holidays", + "description": "Returns holidays", + "method": "GET", + "url": "https://example.com/holidays", + "cacheable": True, + "cache_ttl_seconds": None, + } + ep2: Dict[str, Any] = { + "name": "get_electricity_prices", + "description": "Returns electricity prices", + "method": "GET", + "url": "https://example.com/electricity", + "cacheable": True, + "cache_ttl_seconds": None, + } + + parallel_endpoints = [ + EndpointSessionState( + endpoint=ep1, collected_params={"country": "EE"}, completed=True + ), + EndpointSessionState( + endpoint=ep2, collected_params={"date": "2026-05-28"}, completed=True + ), + ] + + multi_result = MultiAPICallResult( + results=[ + APICallResult( + success=True, status_code=200, response_data={"h": 1}, error=None + ), + APICallResult( + success=True, status_code=200, response_data={"p": 2}, error=None + ), + ], + endpoints=[ep1, ep2], + ) + + mock_multi_caller = AsyncMock() + mock_multi_caller.call_all = AsyncMock(return_value=multi_result) + + mock_cache_store = AsyncMock() + mock_cache_store.set_l1 = AsyncMock() # Should NOT be called + mock_cache_store.set_l2 = AsyncMock() # Should NOT be called + + with ( + patch( + "tool_classifier.workflows.api_tool_workflow.MultiAPICaller", + return_value=mock_multi_caller, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.ATCCacheStore", + return_value=mock_cache_store, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.asyncio.to_thread", + new_callable=AsyncMock, + return_value="Multi-endpoint formatted response.", + ), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + False, # KILL SWITCH: disabled + ), + ): + await executor._execute_multi_api_and_format( + chat_id=_CHAT_ID, + parallel_endpoints=parallel_endpoints, + user_query="Show me holidays and electricity prices", + detected_language="en", + ) + # Flush any background tasks + await asyncio.sleep(0) + + # No cache writes should occur when kill-switch is off + mock_cache_store.set_l1.assert_not_called() + mock_cache_store.set_l2.assert_not_called() + + @pytest.mark.asyncio + async def test_kill_switch_disables_cache_checks_in_compute_loop_step( + self, + ) -> None: + """With cache disabled, _compute_loop_step does NOT check L1 or L2.""" + session_store = _make_session_store(session=None) + executor = _make_executor(session_store) + + mock_cache_store = AsyncMock() + mock_cache_store.get_l1 = AsyncMock() # Should NOT be called + mock_cache_store.get_l2 = AsyncMock() # Should NOT be called + mock_cache_store.set_l1 = AsyncMock() # Should NOT be called + mock_cache_store.set_l2 = AsyncMock() # Should NOT be called + + # Create a loop mock to simulate normal agentic loop behavior + loop_mock = _make_loop_mock() + executor._build_agentic_loop = MagicMock(return_value=loop_mock) + + with ( + patch( + "tool_classifier.workflows.api_tool_workflow.ATCCacheStore", + return_value=mock_cache_store, + ), + patch( + "tool_classifier.workflows.api_tool_workflow.FeatureFlags" + ".ATC_RESPONSE_CACHE_ENABLED", + False, # KILL SWITCH: disabled + ), + ): + step = await executor._compute_loop_step( + _make_request(), + context={"matched_endpoint": _ENDPOINT}, + ) + + # With cache disabled, agentic loop should run normally (not cached_response) + assert step.kind == "question" # Normal loop behavior + + # No cache checks should have been performed + mock_cache_store.get_l1.assert_not_called() + mock_cache_store.get_l2.assert_not_called() + # set_l1 and set_l2 should also not be called during loop execution + mock_cache_store.set_l1.assert_not_called() + mock_cache_store.set_l2.assert_not_called() + + @pytest.mark.asyncio + async def test_kill_switch_disables_invalidate_l2_on_intent_switch(self) -> None: + """With cache disabled, intent switch does NOT call invalidate_l2.""" + # Active session for endpoint A (holidays) + active_session = APIToolSession( + chat_id=_CHAT_ID, + state="collecting_params", + selected_endpoint=_ENDPOINT_HOLIDAYS_CLS, + collected_params={}, + turn_count=1, + max_turns=5, + awaiting_continuation=False, + detected_language="en", + original_query="holiday query", + ) + session_store = _make_session_store(session=active_session) + + classifier = _make_classifier_instance(session_store) + # New query matches endpoint B (weather) — a different endpoint + classifier.api_tool_searcher.search = AsyncMock( + return_value=[_make_api_tool_search_result(_ENDPOINT_WEATHER_CLS)] + ) + + mock_cache_store = AsyncMock() + mock_cache_store.invalidate_l2 = AsyncMock() # Should NOT be called + + with ( + patch( + "tool_classifier.classifier.FeatureFlags" + ".API_TOOL_CALLING_WORKFLOW_ENABLED", + True, + ), + patch( + "tool_classifier.classifier.FeatureFlags.ATC_RESPONSE_CACHE_ENABLED", + False, # KILL SWITCH: disabled + ), + patch( + "tool_classifier.classifier.ATCCacheStore", + return_value=mock_cache_store, + ), + ): + result = await classifier.classify( + query="What is the weather in Tallinn?", + conversation_history=[], + language="en", + request=_make_request("What is the weather in Tallinn?"), + ) + + # With cache disabled, invalidate_l2 should NOT be called + mock_cache_store.invalidate_l2.assert_not_called() + # But session delete should still occur (independent of cache) + session_store.delete.assert_called_once_with(_CHAT_ID) + assert result.workflow == WorkflowType.API_TOOL_CALLING + assert result.metadata.get("matched_endpoint", {}).get("name") == "get_weather" diff --git a/tests/test_atc_cache_store.py b/tests/test_atc_cache_store.py new file mode 100644 index 00000000..38593f47 --- /dev/null +++ b/tests/test_atc_cache_store.py @@ -0,0 +1,497 @@ +"""Unit tests for ATCCacheStore — T-2.""" + +import json +from unittest.mock import AsyncMock, patch + +import pytest + +from models.session_models import LastCallContext +from tool_classifier.constants import ( + ATC_CACHE_DEFAULT_TTL_SECONDS, + ATC_CACHE_KEY_PREFIX, + ATC_LAST_CALL_KEY_PREFIX, + ATC_LAST_CALL_TTL_SECONDS, +) +from utils.atc_cache_store import ATCCacheStore + +# --------------------------------------------------------------------------- +# Shared fixtures +# --------------------------------------------------------------------------- + +CHAT_ID = "chat-test-001" +API_NAME = "get_national_holidays" +PARAMS = {"country": "EE", "year": 2026} +RESPONSE = {"holidays": [{"date": "2026-02-24", "name": "Independence Day"}]} + + +def _make_redis_mock() -> AsyncMock: + mock = AsyncMock() + mock.get = AsyncMock(return_value=None) + mock.set = AsyncMock() + mock.delete = AsyncMock() + return mock + + +def _make_last_call_ctx(api_name: str = API_NAME) -> LastCallContext: + return LastCallContext( + api_name=api_name, + endpoint={"name": api_name, "url": "https://example.com", "method": "GET"}, + collected_params=PARAMS, + raw_response=RESPONSE, + original_query="What are the public holidays in Estonia in 2026?", + timestamp=1748000000.0, + ) + + +def _fake_redis_store() -> tuple[dict, AsyncMock]: + """Return a (store_dict, redis_mock) pair backed by an in-memory dict.""" + store: dict = {} + + async def fake_set(key, value, ex=None): + store[key] = value + + async def fake_get(key): + return store.get(key) + + async def fake_delete(key): + store.pop(key, None) + + mock = AsyncMock() + mock.get = AsyncMock(side_effect=fake_get) + mock.set = AsyncMock(side_effect=fake_set) + mock.delete = AsyncMock(side_effect=fake_delete) + return store, mock + + +# --------------------------------------------------------------------------- +# _normalise_params +# --------------------------------------------------------------------------- + + +class TestNormaliseParams: + def test_strips_whitespace_from_strings(self): + result = ATCCacheStore._normalise_params({"country": " EE "}) + assert result["country"] == "ee" + + def test_casts_numeric_string_to_int(self): + result = ATCCacheStore._normalise_params({"year": "2026"}) + assert result["year"] == 2026 + assert isinstance(result["year"], int) + + def test_lowercases_alpha_only_strings(self): + result = ATCCacheStore._normalise_params({"method": "GET"}) + assert result["method"] == "get" + + def test_leaves_mixed_strings_as_stripped_only(self): + result = ATCCacheStore._normalise_params({"query": " hello world "}) + assert result["query"] == "hello world" + + def test_leaves_non_string_values_unchanged(self): + result = ATCCacheStore._normalise_params({"year": 2026, "active": True}) + assert result["year"] == 2026 + assert result["active"] is True + + def test_mixed_param_types(self): + result = ATCCacheStore._normalise_params( + {"country": " EE ", "year": "2026", "method": "GET", "count": 5} + ) + assert result == {"country": "ee", "year": 2026, "method": "get", "count": 5} + + +# --------------------------------------------------------------------------- +# _param_hash +# --------------------------------------------------------------------------- + + +class TestParamHash: + def test_string_and_int_year_produce_same_hash(self): + """Core normalisation requirement: "2026" and 2026 must hash identically.""" + h_str = ATCCacheStore._param_hash({"year": "2026", "country": "EE"}) + h_int = ATCCacheStore._param_hash({"year": 2026, "country": "EE"}) + assert h_str == h_int + + def test_hash_is_16_hex_chars(self): + h = ATCCacheStore._param_hash({"year": 2026}) + assert len(h) == 16 + assert all(c in "0123456789abcdef" for c in h) + + def test_different_params_produce_different_hash(self): + h1 = ATCCacheStore._param_hash({"year": 2026, "country": "EE"}) + h2 = ATCCacheStore._param_hash({"year": 2025, "country": "EE"}) + assert h1 != h2 + + def test_key_order_does_not_affect_hash(self): + h1 = ATCCacheStore._param_hash({"country": "EE", "year": 2026}) + h2 = ATCCacheStore._param_hash({"year": 2026, "country": "EE"}) + assert h1 == h2 + + def test_whitespace_is_normalised_before_hashing(self): + h1 = ATCCacheStore._param_hash({"country": "EE"}) + h2 = ATCCacheStore._param_hash({"country": " EE "}) + assert h1 == h2 + + def test_alpha_case_is_normalised_before_hashing(self): + h1 = ATCCacheStore._param_hash({"method": "get"}) + h2 = ATCCacheStore._param_hash({"method": "GET"}) + assert h1 == h2 + + +# --------------------------------------------------------------------------- +# Key format +# --------------------------------------------------------------------------- + + +class TestKeyFormat: + def test_l1_key_includes_all_components(self): + expected_hash = ATCCacheStore._param_hash(PARAMS) + key = ATCCacheStore._l1_key(CHAT_ID, API_NAME, PARAMS) + assert key == f"{ATC_CACHE_KEY_PREFIX}:{CHAT_ID}:{API_NAME}:{expected_hash}" + + def test_l2_key_format(self): + key = ATCCacheStore._l2_key(CHAT_ID) + assert key == f"{ATC_LAST_CALL_KEY_PREFIX}:{CHAT_ID}" + + def test_l1_and_l2_keys_have_different_prefixes(self): + l1 = ATCCacheStore._l1_key(CHAT_ID, API_NAME, PARAMS) + l2 = ATCCacheStore._l2_key(CHAT_ID) + assert not l1.startswith(ATC_LAST_CALL_KEY_PREFIX) + assert not l2.startswith(ATC_CACHE_KEY_PREFIX) + + +# --------------------------------------------------------------------------- +# get_l1 +# --------------------------------------------------------------------------- + + +class TestGetL1: + @pytest.mark.asyncio + async def test_returns_deserialised_value_on_hit(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(return_value=json.dumps(RESPONSE)) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + result = await store.get_l1(CHAT_ID, API_NAME, PARAMS) + + assert result == RESPONSE + + @pytest.mark.asyncio + async def test_returns_none_on_miss(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(return_value=None) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + result = await store.get_l1(CHAT_ID, API_NAME, PARAMS) + + assert result is None + + @pytest.mark.asyncio + async def test_returns_none_when_redis_unavailable(self): + store = ATCCacheStore() + with patch("utils.atc_cache_store.get_redis_client", return_value=None): + result = await store.get_l1(CHAT_ID, API_NAME, PARAMS) + + assert result is None + + @pytest.mark.asyncio + async def test_returns_none_on_redis_exception(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(side_effect=RuntimeError("connection lost")) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + result = await store.get_l1(CHAT_ID, API_NAME, PARAMS) + + assert result is None + + @pytest.mark.asyncio + async def test_different_params_return_none(self): + """Different params hash to a different key — must not return the stored entry.""" + store = ATCCacheStore() + stored_key = ATCCacheStore._l1_key(CHAT_ID, API_NAME, PARAMS) + different_params = {"country": "LV", "year": 2026} + + async def selective_get(key): + return json.dumps(RESPONSE) if key == stored_key else None + + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(side_effect=selective_get) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + result = await store.get_l1(CHAT_ID, API_NAME, different_params) + + assert result is None + + +# --------------------------------------------------------------------------- +# set_l1 +# --------------------------------------------------------------------------- + + +class TestSetL1: + @pytest.mark.asyncio + async def test_calls_redis_set_with_correct_key_and_value(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + expected_key = ATCCacheStore._l1_key(CHAT_ID, API_NAME, PARAMS) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l1(CHAT_ID, API_NAME, PARAMS, RESPONSE) + + redis_mock.set.assert_called_once_with( + expected_key, + json.dumps(RESPONSE), + ex=ATC_CACHE_DEFAULT_TTL_SECONDS, + ) + + @pytest.mark.asyncio + async def test_uses_default_ttl_when_not_specified(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l1(CHAT_ID, API_NAME, PARAMS, RESPONSE) + + assert redis_mock.set.call_args[1]["ex"] == ATC_CACHE_DEFAULT_TTL_SECONDS + + @pytest.mark.asyncio + async def test_uses_custom_ttl_when_provided(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l1(CHAT_ID, API_NAME, PARAMS, RESPONSE, ttl=120) + + assert redis_mock.set.call_args[1]["ex"] == 120 + + @pytest.mark.asyncio + async def test_no_op_when_redis_unavailable(self): + store = ATCCacheStore() + with patch("utils.atc_cache_store.get_redis_client", return_value=None): + await store.set_l1(CHAT_ID, API_NAME, PARAMS, RESPONSE) # must not raise + + @pytest.mark.asyncio + async def test_no_op_on_redis_exception(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + redis_mock.set = AsyncMock(side_effect=RuntimeError("write failed")) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l1(CHAT_ID, API_NAME, PARAMS, RESPONSE) # must not raise + + +# --------------------------------------------------------------------------- +# L1 round-trip +# --------------------------------------------------------------------------- + + +class TestL1RoundTrip: + @pytest.mark.asyncio + async def test_set_then_get_returns_same_response(self): + store = ATCCacheStore() + _, redis_mock = _fake_redis_store() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l1(CHAT_ID, API_NAME, PARAMS, RESPONSE) + result = await store.get_l1(CHAT_ID, API_NAME, PARAMS) + + assert result == RESPONSE + + @pytest.mark.asyncio + async def test_string_year_hits_entry_stored_with_int_year(self): + """Normalisation: set with int year, get with string year — must still hit.""" + store = ATCCacheStore() + _, redis_mock = _fake_redis_store() + params_int = {"country": "EE", "year": 2026} + params_str = {"country": "EE", "year": "2026"} + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l1(CHAT_ID, API_NAME, params_int, RESPONSE) + result = await store.get_l1(CHAT_ID, API_NAME, params_str) + + assert result == RESPONSE + + @pytest.mark.asyncio + async def test_list_response_survives_round_trip(self): + """raw_response may be a list — must serialise/deserialise cleanly.""" + store = ATCCacheStore() + _, redis_mock = _fake_redis_store() + list_response = [{"date": "2026-02-24"}, {"date": "2026-06-23"}] + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l1(CHAT_ID, API_NAME, PARAMS, list_response) + result = await store.get_l1(CHAT_ID, API_NAME, PARAMS) + + assert result == list_response + + +# --------------------------------------------------------------------------- +# get_l2 / set_l2 +# --------------------------------------------------------------------------- + + +class TestSetAndGetL2: + @pytest.mark.asyncio + async def test_round_trip_returns_correct_context_list(self): + store = ATCCacheStore() + ctx = _make_last_call_ctx() + _, redis_mock = _fake_redis_store() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l2(CHAT_ID, [ctx]) + result = await store.get_l2(CHAT_ID) + + assert result is not None + assert len(result) == 1 + assert result[0].api_name == API_NAME + assert result[0].collected_params == PARAMS + assert result[0].raw_response == RESPONSE + assert result[0].original_query == ctx.original_query + + @pytest.mark.asyncio + async def test_multi_intent_stores_all_entries(self): + """set_l2 with two contexts — get_l2 returns both.""" + store = ATCCacheStore() + ctx1 = _make_last_call_ctx("get_national_holidays") + ctx2 = _make_last_call_ctx("get_electricity_prices") + _, redis_mock = _fake_redis_store() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l2(CHAT_ID, [ctx1, ctx2]) + result = await store.get_l2(CHAT_ID) + + assert result is not None + assert len(result) == 2 + assert {r.api_name for r in result} == { + "get_national_holidays", + "get_electricity_prices", + } + + @pytest.mark.asyncio + async def test_set_l2_uses_correct_ttl(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + ctx = _make_last_call_ctx() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l2(CHAT_ID, [ctx]) + + assert redis_mock.set.call_args[1]["ex"] == ATC_LAST_CALL_TTL_SECONDS + + @pytest.mark.asyncio + async def test_set_l2_writes_to_correct_key(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + ctx = _make_last_call_ctx() + expected_key = ATCCacheStore._l2_key(CHAT_ID) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l2(CHAT_ID, [ctx]) + + assert redis_mock.set.call_args[0][0] == expected_key + + @pytest.mark.asyncio + async def test_get_l2_returns_none_on_miss(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + result = await store.get_l2(CHAT_ID) + + assert result is None + + @pytest.mark.asyncio + async def test_get_l2_returns_none_when_redis_unavailable(self): + store = ATCCacheStore() + with patch("utils.atc_cache_store.get_redis_client", return_value=None): + result = await store.get_l2(CHAT_ID) + + assert result is None + + @pytest.mark.asyncio + async def test_get_l2_returns_none_on_exception(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(side_effect=RuntimeError("connection reset")) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + result = await store.get_l2(CHAT_ID) + + assert result is None + + @pytest.mark.asyncio + async def test_set_l2_no_op_when_redis_unavailable(self): + store = ATCCacheStore() + ctx = _make_last_call_ctx() + with patch("utils.atc_cache_store.get_redis_client", return_value=None): + await store.set_l2(CHAT_ID, [ctx]) # must not raise + + @pytest.mark.asyncio + async def test_set_l2_no_op_on_redis_exception(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + redis_mock.set = AsyncMock(side_effect=RuntimeError("write error")) + ctx = _make_last_call_ctx() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l2(CHAT_ID, [ctx]) # must not raise + + +# --------------------------------------------------------------------------- +# invalidate_l2 +# --------------------------------------------------------------------------- + + +class TestInvalidateL2: + @pytest.mark.asyncio + async def test_calls_delete_with_correct_l2_key(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + expected_key = ATCCacheStore._l2_key(CHAT_ID) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.invalidate_l2(CHAT_ID) + + redis_mock.delete.assert_called_once_with(expected_key) + + @pytest.mark.asyncio + async def test_get_l2_returns_none_after_invalidate(self): + store = ATCCacheStore() + _, redis_mock = _fake_redis_store() + ctx = _make_last_call_ctx() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.set_l2(CHAT_ID, [ctx]) + await store.invalidate_l2(CHAT_ID) + result = await store.get_l2(CHAT_ID) + + assert result is None + + @pytest.mark.asyncio + async def test_invalidate_only_deletes_l2_key_not_l1(self): + """invalidate_l2 must call delete exactly once with the L2 prefix.""" + store = ATCCacheStore() + redis_mock = _make_redis_mock() + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.invalidate_l2(CHAT_ID) + + deleted_key: str = redis_mock.delete.call_args[0][0] + assert deleted_key.startswith(ATC_LAST_CALL_KEY_PREFIX) + assert not deleted_key.startswith(ATC_CACHE_KEY_PREFIX) + + @pytest.mark.asyncio + async def test_no_op_when_redis_unavailable(self): + store = ATCCacheStore() + with patch("utils.atc_cache_store.get_redis_client", return_value=None): + await store.invalidate_l2(CHAT_ID) # must not raise + + @pytest.mark.asyncio + async def test_no_op_on_redis_exception(self): + store = ATCCacheStore() + redis_mock = _make_redis_mock() + redis_mock.delete = AsyncMock(side_effect=RuntimeError("gone")) + + with patch("utils.atc_cache_store.get_redis_client", return_value=redis_mock): + await store.invalidate_l2(CHAT_ID) # must not raise diff --git a/tests/test_conversation_history_helpers.py b/tests/test_conversation_history_helpers.py new file mode 100644 index 00000000..a5bbd61a --- /dev/null +++ b/tests/test_conversation_history_helpers.py @@ -0,0 +1,249 @@ +"""Unit tests for src.utils.conversation_history_helpers.get_conversation_history.""" + +from unittest.mock import AsyncMock + +import pytest + +from src.models.conversation_history_models import ( + ConversationHistoryState, + ConversationRound, +) +from models.request_models import ConversationItem +from src.utils.conversation_history_helpers import get_conversation_history + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_item(role: str = "user", msg: str = "hello") -> ConversationItem: + return ConversationItem( + authorRole=role, message=msg, timestamp="2024-01-01T00:00:00" + ) # type: ignore[arg-type] + + +def _make_round( + user: str = "What is the rate?", + bot: str = "The rate is 20%.", + ts: float = 1_700_000_000.0, +) -> ConversationRound: + return ConversationRound(user_message=user, bot_message=bot, timestamp=ts) + + +# --------------------------------------------------------------------------- +# Redis available with rounds +# --------------------------------------------------------------------------- + + +class TestGetConversationHistoryRedisRounds: + @pytest.mark.asyncio + async def test_returns_redis_rounds_as_conversation_items(self) -> None: + """Two ConversationItems (user + bot) per stored round.""" + round_ = _make_round() + state = ConversationHistoryState(chat_id="c1", rounds=[round_], summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + fallback = [_make_item()] + history, summary = await get_conversation_history("c1", store, fallback) + + assert len(history) == 2 + assert history[0].authorRole == "user" + assert history[0].message == round_.user_message + assert history[1].authorRole == "bot" + assert history[1].message == round_.bot_message + assert summary is None + + @pytest.mark.asyncio + async def test_ignores_fallback_when_redis_has_rounds(self) -> None: + """Fallback list is not returned when Redis has valid rounds.""" + round_ = _make_round() + state = ConversationHistoryState(chat_id="c1", rounds=[round_], summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + fallback = [_make_item(msg="fallback-only message")] + history, _ = await get_conversation_history("c1", store, fallback) + + messages = [item.message for item in history] + assert "fallback-only message" not in messages + + @pytest.mark.asyncio + async def test_returns_summary_alongside_rounds(self) -> None: + """Redis summary is returned as the second tuple element.""" + round_ = _make_round() + state = ConversationHistoryState( + chat_id="c1", + rounds=[round_], + summary="Earlier we discussed tax rates.", + ) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + _, summary = await get_conversation_history("c1", store, []) + + assert summary == "Earlier we discussed tax rates." + + @pytest.mark.asyncio + async def test_multiple_rounds_expand_to_all_items(self) -> None: + """N rounds → 2*N ConversationItems in order.""" + rounds = [ + _make_round(user="q1", bot="a1"), + _make_round(user="q2", bot="a2"), + _make_round(user="q3", bot="a3"), + ] + state = ConversationHistoryState(chat_id="c1", rounds=rounds, summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + history, _ = await get_conversation_history("c1", store, []) + + assert len(history) == 6 + assert history[0].message == "q1" + assert history[1].message == "a1" + assert history[4].message == "q3" + assert history[5].message == "a3" + + @pytest.mark.asyncio + async def test_timestamp_is_str_of_round_timestamp(self) -> None: + """Timestamp on returned items is the string form of the round's float timestamp.""" + round_ = _make_round(ts=1_700_000_123.456) + state = ConversationHistoryState(chat_id="c1", rounds=[round_], summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + history, _ = await get_conversation_history("c1", store, []) + + assert history[0].timestamp == str(round_.timestamp) + assert history[1].timestamp == str(round_.timestamp) + + +# --------------------------------------------------------------------------- +# Redis empty → fallback +# --------------------------------------------------------------------------- + + +class TestGetConversationHistoryRedisEmpty: + @pytest.mark.asyncio + async def test_falls_back_when_redis_returns_no_rounds(self) -> None: + """Empty rounds list → fallback returned with summary=None.""" + state = ConversationHistoryState( + chat_id="c1", rounds=[], summary="stale summary" + ) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + fallback = [_make_item(msg="from request")] + history, summary = await get_conversation_history("c1", store, fallback) + + assert len(history) == 1 + assert history[0].message == "from request" + assert summary is None # summary not returned when rounds are empty + + +# --------------------------------------------------------------------------- +# Redis unavailable → graceful degradation +# --------------------------------------------------------------------------- + + +class TestGetConversationHistoryRedisUnavailable: + @pytest.mark.asyncio + async def test_falls_back_when_get_context_raises(self) -> None: + """Any exception from get_context → fallback with summary=None, no propagation.""" + store = AsyncMock() + store.get_context = AsyncMock(side_effect=RuntimeError("Redis down")) + + fallback = [_make_item(msg="fallback on error")] + history, summary = await get_conversation_history("c1", store, fallback) + + assert len(history) == 1 + assert history[0].message == "fallback on error" + assert summary is None + + @pytest.mark.asyncio + async def test_falls_back_when_store_is_none(self) -> None: + """When store=None, fallback is returned immediately without calling Redis.""" + fallback = [_make_item(role="bot", msg="bot message")] + history, summary = await get_conversation_history("c1", None, fallback) + + assert len(history) == 1 + assert history[0].authorRole == "bot" + assert summary is None + + @pytest.mark.asyncio + async def test_does_not_raise_on_connection_error(self) -> None: + """ConnectionError from Redis is caught and does not propagate.""" + store = AsyncMock() + store.get_context = AsyncMock(side_effect=ConnectionError("refused")) + + history, summary = await get_conversation_history("c1", store, []) + + assert history == [] + assert summary is None + + @pytest.mark.asyncio + async def test_get_context_called_with_correct_chat_id(self) -> None: + """The helper passes chat_id to store.get_context.""" + state = ConversationHistoryState(chat_id="my-chat", rounds=[], summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + await get_conversation_history("my-chat", store, []) + + store.get_context.assert_awaited_once_with("my-chat") + + +# --------------------------------------------------------------------------- +# Conversion correctness +# --------------------------------------------------------------------------- + + +class TestGetConversationHistoryConversionCorrectness: + @pytest.mark.asyncio + async def test_returned_items_are_conversation_item_instances(self) -> None: + """All returned history items must be ConversationItem Pydantic objects.""" + round_ = _make_round() + state = ConversationHistoryState(chat_id="c1", rounds=[round_], summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + history, _ = await get_conversation_history("c1", store, []) + + for item in history: + assert isinstance(item, ConversationItem) + + @pytest.mark.asyncio + async def test_user_role_is_literal_user(self) -> None: + """User item authorRole must be the literal string 'user'.""" + round_ = _make_round() + state = ConversationHistoryState(chat_id="c1", rounds=[round_], summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + history, _ = await get_conversation_history("c1", store, []) + + assert history[0].authorRole == "user" + + @pytest.mark.asyncio + async def test_bot_role_is_literal_bot(self) -> None: + """Bot item authorRole must be the literal string 'bot'.""" + round_ = _make_round() + state = ConversationHistoryState(chat_id="c1", rounds=[round_], summary=None) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + history, _ = await get_conversation_history("c1", store, []) + + assert history[1].authorRole == "bot" + + @pytest.mark.asyncio + async def test_fallback_items_returned_unchanged(self) -> None: + """Fallback items are returned as-is (same objects, same content).""" + fallback = [ + _make_item(role="user", msg="user says"), + _make_item(role="bot", msg="bot replies"), + ] + history, _ = await get_conversation_history("c1", None, fallback) + + assert history is fallback diff --git a/tests/test_conversation_history_store.py b/tests/test_conversation_history_store.py new file mode 100644 index 00000000..59ab9775 --- /dev/null +++ b/tests/test_conversation_history_store.py @@ -0,0 +1,805 @@ +"""Unit tests for ConversationHistoryStore and conversation history models.""" + +import asyncio +import json +import time +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from pydantic import ValidationError +from redis import WatchError + +from src.models.conversation_history_models import ( + ConversationHistoryState, + ConversationRound, +) +from src.utils.conversation_history_store import ( + ConversationHistoryStore, + _HISTORY_TTL_SECONDS, + _MAX_ROUNDS, + _history_key, + _summary_key, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_round(**kwargs) -> ConversationRound: + defaults = { + "user_message": "Hello", + "bot_message": "Hi there!", + } + defaults.update(kwargs) + return ConversationRound(**defaults) + + +def _make_redis_mock() -> AsyncMock: + mock = AsyncMock() + mock.get = AsyncMock(return_value=None) + mock.set = AsyncMock() + mock.expire = AsyncMock() + mock.ping = AsyncMock(return_value=True) + return mock + + +def _make_pipe_mock() -> AsyncMock: + pipe = AsyncMock() + pipe.watch = AsyncMock() + pipe.unwatch = AsyncMock() + pipe.get = AsyncMock(return_value=None) + pipe.multi = MagicMock() + pipe.set = MagicMock() + pipe.expire = MagicMock() + pipe.execute = AsyncMock(return_value=[True, True]) + pipe.__aenter__ = AsyncMock(return_value=pipe) + pipe.__aexit__ = AsyncMock(return_value=False) + return pipe + + +# --------------------------------------------------------------------------- +# ConversationRound model tests +# --------------------------------------------------------------------------- + + +class TestConversationRound: + def test_required_fields(self): + r = _make_round() + assert r.user_message == "Hello" + assert r.bot_message == "Hi there!" + assert isinstance(r.timestamp, float) + + def test_timestamp_defaults_to_current_time(self): + before = time.time() + r = _make_round() + after = time.time() + assert before <= r.timestamp <= after + + def test_explicit_timestamp(self): + r = ConversationRound(user_message="u", bot_message="b", timestamp=12345.0) + assert r.timestamp == 12345.0 + + def test_user_message_required(self): + with pytest.raises(ValidationError): + ConversationRound(bot_message="b") # type: ignore[call-arg] + + def test_bot_message_required(self): + with pytest.raises(ValidationError): + ConversationRound(user_message="u") # type: ignore[call-arg] + + def test_serialization_roundtrip(self): + r = _make_round(user_message="What is the weather?", bot_message="It is sunny.") + restored = ConversationRound.model_validate_json(r.model_dump_json()) + assert restored == r + + +# --------------------------------------------------------------------------- +# ConversationHistoryState model tests +# --------------------------------------------------------------------------- + + +class TestConversationHistoryState: + def test_defaults(self): + state = ConversationHistoryState(chat_id="chat-1") + assert state.rounds == [] + assert state.summary is None + + def test_chat_id_required(self): + with pytest.raises(ValidationError): + ConversationHistoryState() # type: ignore[call-arg] + + def test_serialization_roundtrip(self): + state = ConversationHistoryState( + chat_id="chat-2", + rounds=[_make_round(), _make_round(user_message="q2", bot_message="a2")], + summary="User asked about weather and holidays.", + ) + restored = ConversationHistoryState.model_validate_json(state.model_dump_json()) + assert restored == state + assert len(restored.rounds) == 2 + assert restored.summary == "User asked about weather and holidays." + + +# --------------------------------------------------------------------------- +# Key helper tests +# --------------------------------------------------------------------------- + + +class TestKeyHelpers: + def test_history_key(self): + assert _history_key("abc") == "conv:abc" + + def test_summary_key(self): + assert _summary_key("abc") == "conv:summary:abc" + + +# --------------------------------------------------------------------------- +# ConversationHistoryStore.save_round +# --------------------------------------------------------------------------- + + +class TestSaveRound: + @pytest.mark.asyncio + async def test_save_round_appends_and_sets_ttl(self): + store = ConversationHistoryStore() + round_ = _make_round() + + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=None) # no existing rounds + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-1", round_) + + pipe.watch.assert_awaited_once_with(_history_key("chat-1")) + pipe.multi.assert_called_once() + pipe.set.assert_called_once() + set_call = pipe.set.call_args + assert set_call[0][0] == _history_key("chat-1") + assert set_call[1]["ex"] == _HISTORY_TTL_SECONDS + # Verify stored JSON contains the round + stored = json.loads(set_call[0][1]) + assert len(stored) == 1 + assert stored[0]["user_message"] == round_.user_message + + @pytest.mark.asyncio + async def test_save_round_appends_to_existing_rounds(self): + store = ConversationHistoryStore() + existing = [ + _make_round(user_message=f"q{i}", bot_message=f"a{i}").model_dump() + for i in range(3) + ] + new_round = _make_round(user_message="q_new", bot_message="a_new") + + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=json.dumps(existing)) + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-2", new_round) + + stored_json = pipe.set.call_args[0][1] + stored = json.loads(stored_json) + assert len(stored) == 4 + assert stored[-1]["user_message"] == "q_new" + + @pytest.mark.asyncio + async def test_save_round_trims_to_max_rounds(self): + store = ConversationHistoryStore() + # Start with exactly MAX_ROUNDS rounds + existing = [ + _make_round(user_message=f"q{i}", bot_message=f"a{i}").model_dump() + for i in range(_MAX_ROUNDS) + ] + new_round = _make_round(user_message="q_overflow", bot_message="a_overflow") + + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=json.dumps(existing)) + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-3", new_round) + + stored_json = pipe.set.call_args[0][1] + stored = json.loads(stored_json) + assert len(stored) == _MAX_ROUNDS + # Oldest round was evicted; newest is last + assert stored[-1]["user_message"] == "q_overflow" + assert stored[0]["user_message"] == "q1" + + @pytest.mark.asyncio + async def test_save_round_resets_summary_ttl(self): + store = ConversationHistoryStore() + round_ = _make_round() + + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=None) + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-4", round_) + + pipe.expire.assert_called_once_with( + _summary_key("chat-4"), _HISTORY_TTL_SECONDS + ) + + @pytest.mark.asyncio + async def test_save_round_skips_when_redis_unavailable(self): + store = ConversationHistoryStore() + with patch( + "src.utils.conversation_history_store.get_redis_client", return_value=None + ): + # Should not raise + await store.save_round("chat-x", _make_round()) + + @pytest.mark.asyncio + async def test_save_round_retries_on_watch_error(self): + store = ConversationHistoryStore() + round_ = _make_round() + + call_count = 0 + + async def fake_execute(): + nonlocal call_count + call_count += 1 + if call_count < 3: + raise WatchError("conflict") + return [True, True] + + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=None) + pipe.execute = fake_execute + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-5", round_) + + assert call_count == 3 + + @pytest.mark.asyncio + async def test_save_round_exhausts_retries_gracefully(self): + store = ConversationHistoryStore() + round_ = _make_round() + + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=None) + pipe.execute = AsyncMock(side_effect=WatchError("always conflicts")) + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + # Should not raise after exhausting retries + await store.save_round("chat-6", round_) + + @pytest.mark.asyncio + async def test_save_round_graceful_on_unexpected_error(self): + store = ConversationHistoryStore() + + pipe = _make_pipe_mock() + pipe.watch = AsyncMock(side_effect=RuntimeError("boom")) + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-7", _make_round()) + + +# --------------------------------------------------------------------------- +# ConversationHistoryStore.get_history +# --------------------------------------------------------------------------- + + +class TestGetHistory: + @pytest.mark.asyncio + async def test_returns_empty_list_when_key_missing(self): + store = ConversationHistoryStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(return_value=None) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + result = await store.get_history("missing-chat") + + assert result == [] + + @pytest.mark.asyncio + async def test_returns_deserialized_rounds(self): + store = ConversationHistoryStore() + rounds = [ + _make_round(user_message="q1", bot_message="a1"), + _make_round(user_message="q2", bot_message="a2"), + ] + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock( + return_value=json.dumps([r.model_dump() for r in rounds]) + ) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + result = await store.get_history("chat-1") + + assert len(result) == 2 + assert result[0].user_message == "q1" + assert result[1].user_message == "q2" + + @pytest.mark.asyncio + async def test_returns_empty_list_when_redis_unavailable(self): + store = ConversationHistoryStore() + with patch( + "src.utils.conversation_history_store.get_redis_client", return_value=None + ): + result = await store.get_history("any-chat") + + assert result == [] + + @pytest.mark.asyncio + async def test_returns_empty_list_on_error(self): + store = ConversationHistoryStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(side_effect=RuntimeError("connection lost")) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + result = await store.get_history("chat-err") + + assert result == [] + + +# --------------------------------------------------------------------------- +# ConversationHistoryStore.get_summary +# --------------------------------------------------------------------------- + + +class TestGetSummary: + @pytest.mark.asyncio + async def test_returns_none_when_key_missing(self): + store = ConversationHistoryStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(return_value=None) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + result = await store.get_summary("chat-1") + + assert result is None + + @pytest.mark.asyncio + async def test_returns_summary_string(self): + store = ConversationHistoryStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(return_value="User asked about holidays.") + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + result = await store.get_summary("chat-1") + + assert result == "User asked about holidays." + + @pytest.mark.asyncio + async def test_returns_none_when_redis_unavailable(self): + store = ConversationHistoryStore() + with patch( + "src.utils.conversation_history_store.get_redis_client", return_value=None + ): + result = await store.get_summary("any-chat") + + assert result is None + + @pytest.mark.asyncio + async def test_returns_none_on_error(self): + store = ConversationHistoryStore() + redis_mock = _make_redis_mock() + redis_mock.get = AsyncMock(side_effect=RuntimeError("boom")) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + result = await store.get_summary("chat-err") + + assert result is None + + +# --------------------------------------------------------------------------- +# ConversationHistoryStore.save_summary +# --------------------------------------------------------------------------- + + +class TestSaveSummary: + @pytest.mark.asyncio + async def test_save_summary_sets_key_with_ttl(self): + store = ConversationHistoryStore() + + pipe = _make_pipe_mock() + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_summary("chat-1", "Some summary text.") + + pipe.set.assert_called_once() + set_call = pipe.set.call_args + assert set_call[0][0] == _summary_key("chat-1") + assert set_call[0][1] == "Some summary text." + assert set_call[1]["ex"] == _HISTORY_TTL_SECONDS + + @pytest.mark.asyncio + async def test_save_summary_resets_history_key_ttl(self): + store = ConversationHistoryStore() + + pipe = _make_pipe_mock() + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_summary("chat-2", "Summary.") + + pipe.expire.assert_called_once_with( + _history_key("chat-2"), _HISTORY_TTL_SECONDS + ) + + @pytest.mark.asyncio + async def test_save_summary_skips_when_redis_unavailable(self): + store = ConversationHistoryStore() + with patch( + "src.utils.conversation_history_store.get_redis_client", return_value=None + ): + await store.save_summary("chat-x", "text") + + @pytest.mark.asyncio + async def test_save_summary_graceful_on_error(self): + store = ConversationHistoryStore() + pipe = _make_pipe_mock() + pipe.execute = AsyncMock(side_effect=RuntimeError("io error")) + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_summary("chat-err", "text") + + +# --------------------------------------------------------------------------- +# ConversationHistoryStore.get_context +# --------------------------------------------------------------------------- + + +class TestGetContext: + @pytest.mark.asyncio + async def test_get_context_combines_history_and_summary(self): + store = ConversationHistoryStore() + rounds = [_make_round(user_message="hello", bot_message="hi")] + summary_text = "Earlier the user greeted the bot." + + async def _fake_get_history(chat_id: str): + return rounds + + async def _fake_get_summary(chat_id: str): + return summary_text + + with ( + patch.object(store, "get_history", side_effect=_fake_get_history), + patch.object(store, "get_summary", side_effect=_fake_get_summary), + ): + result = await store.get_context("chat-1") + + assert isinstance(result, ConversationHistoryState) + assert result.chat_id == "chat-1" + assert result.rounds == rounds + assert result.summary == summary_text + + @pytest.mark.asyncio + async def test_get_context_returns_empty_defaults_when_redis_unavailable(self): + store = ConversationHistoryStore() + with patch( + "src.utils.conversation_history_store.get_redis_client", return_value=None + ): + result = await store.get_context("chat-gone") + + assert isinstance(result, ConversationHistoryState) + assert result.chat_id == "chat-gone" + assert result.rounds == [] + assert result.summary is None + + @pytest.mark.asyncio + async def test_get_context_fetches_concurrently(self): + """Both sub-calls must run; verify gather behaviour by checking both are awaited.""" + store = ConversationHistoryStore() + history_called = False + summary_called = False + + async def _hist(chat_id: str): + nonlocal history_called + history_called = True + return [] + + async def _summ(chat_id: str): + nonlocal summary_called + summary_called = True + return None + + with ( + patch.object(store, "get_history", side_effect=_hist), + patch.object(store, "get_summary", side_effect=_summ), + ): + await store.get_context("chat-concurrent") + + assert history_called + assert summary_called + + +# --------------------------------------------------------------------------- +# Incremental summary: save_round eviction triggering +# --------------------------------------------------------------------------- + + +class TestSaveRoundIncrementalSummary: + """Tests for the fire-and-forget summarizer integration in save_round.""" + + @staticmethod + def _make_pipe_with_existing(rounds_data: list[dict]) -> AsyncMock: + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=json.dumps(rounds_data)) + return pipe + + @pytest.mark.asyncio + async def test_eviction_triggers_background_task(self): + """When len(rounds) exceeds _MAX_ROUNDS, summarizer is called with the + evicted round(s) and the existing summary.""" + received_summary: list[str | None] = [] + received_evicted: list[list] = [] + + async def mock_summarizer( + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + received_summary.append(existing_summary) + received_evicted.append(list(evicted_rounds)) + return "merged summary" + + store = ConversationHistoryStore(summarizer=mock_summarizer) + + # Pre-populate with exactly _MAX_ROUNDS rounds so the new one triggers eviction. + existing = [ + _make_round(user_message=f"q{i}", bot_message=f"a{i}").model_dump() + for i in range(_MAX_ROUNDS) + ] + pipe = self._make_pipe_with_existing(existing) + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + new_round = _make_round(user_message="q_new", bot_message="a_new") + + with ( + patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ), + patch.object(store, "_get_summary_lock", return_value=asyncio.Lock()), + patch.object(store, "get_summary", AsyncMock(return_value="old summary")), + patch.object(store, "save_summary", AsyncMock()), + ): + await store.save_round("chat-evict", new_round) + # Drain the event loop so the background task completes. + await asyncio.gather(*store._pending_tasks, return_exceptions=True) + + assert len(received_evicted) == 1 + assert len(received_evicted[0]) == 1 + assert received_evicted[0][0].user_message == "q0" + assert received_summary[0] == "old summary" + + @pytest.mark.asyncio + async def test_no_eviction_does_not_trigger_summarizer(self): + """When rounds stay within _MAX_ROUNDS, the summarizer is never called.""" + called = False + + async def mock_summarizer( + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + nonlocal called + called = True + return "should not be called" + + store = ConversationHistoryStore(summarizer=mock_summarizer) + + # Start with fewer than _MAX_ROUNDS rounds. + existing = [ + _make_round(user_message=f"q{i}", bot_message=f"a{i}").model_dump() + for i in range(_MAX_ROUNDS - 2) + ] + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=json.dumps(existing)) + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-no-evict", _make_round()) + + assert not called + assert len(store._pending_tasks) == 0 + + @pytest.mark.asyncio + async def test_no_summarizer_no_task_on_eviction(self): + """Store without a summarizer still trims correctly; no tasks scheduled.""" + store = ConversationHistoryStore(summarizer=None) + + existing = [ + _make_round(user_message=f"q{i}", bot_message=f"a{i}").model_dump() + for i in range(_MAX_ROUNDS) + ] + pipe = _make_pipe_mock() + pipe.get = AsyncMock(return_value=json.dumps(existing)) + + redis_mock = _make_redis_mock() + redis_mock.pipeline = MagicMock(return_value=pipe) + + with patch( + "src.utils.conversation_history_store.get_redis_client", + return_value=redis_mock, + ): + await store.save_round("chat-no-summarizer", _make_round()) + + assert len(store._pending_tasks) == 0 + + # Verify that trimming still happened. + stored_json = pipe.set.call_args[0][1] + stored = json.loads(stored_json) + assert len(stored) == _MAX_ROUNDS + + +# --------------------------------------------------------------------------- +# Incremental summary: _run_incremental_summary +# --------------------------------------------------------------------------- + + +class TestRunIncrementalSummary: + """Tests for the _run_incremental_summary private method.""" + + @pytest.mark.asyncio + async def test_calls_save_summary_with_merged_result(self): + """Happy path: summarizer returns non-empty string → save_summary called.""" + + async def mock_summarizer( + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + return "merged: " + (existing_summary or "") + " + new info" + + store = ConversationHistoryStore(summarizer=mock_summarizer) + evicted = [_make_round(user_message="old q", bot_message="old a")] + + with ( + patch.object(store, "_get_summary_lock", return_value=asyncio.Lock()), + patch.object(store, "get_summary", AsyncMock(return_value="prior summary")), + patch.object(store, "save_summary", AsyncMock()) as mock_save, + ): + await store._run_incremental_summary("chat-1", evicted) + + mock_save.assert_awaited_once_with("chat-1", "merged: prior summary + new info") + + @pytest.mark.asyncio + async def test_skips_save_when_summarizer_returns_empty(self): + """If summarizer returns empty string, save_summary must NOT be called.""" + + async def mock_summarizer( + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + return "" + + store = ConversationHistoryStore(summarizer=mock_summarizer) + evicted = [_make_round()] + + with ( + patch.object(store, "_get_summary_lock", return_value=asyncio.Lock()), + patch.object(store, "get_summary", AsyncMock(return_value=None)), + patch.object(store, "save_summary", AsyncMock()) as mock_save, + ): + await store._run_incremental_summary("chat-2", evicted) + + mock_save.assert_not_awaited() + + @pytest.mark.asyncio + async def test_summarizer_exception_does_not_propagate(self): + """A failing summarizer must be caught; the method returns cleanly.""" + + async def exploding_summarizer( + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + raise RuntimeError("LLM unavailable") + + store = ConversationHistoryStore(summarizer=exploding_summarizer) + evicted = [_make_round()] + + with ( + patch.object(store, "_get_summary_lock", return_value=asyncio.Lock()), + patch.object(store, "get_summary", AsyncMock(return_value=None)), + patch.object(store, "save_summary", AsyncMock()) as mock_save, + ): + # Must not raise. + await store._run_incremental_summary("chat-3", evicted) + + mock_save.assert_not_awaited() + + @pytest.mark.asyncio + async def test_passes_none_existing_summary_when_no_summary_stored(self): + """When get_summary returns None, summarizer receives None as first arg.""" + received: list[str | None] = [] + + async def capture_summarizer( + existing_summary: str | None, + evicted_rounds: list[ConversationRound], + ) -> str: + received.append(existing_summary) + return "new summary" + + store = ConversationHistoryStore(summarizer=capture_summarizer) + evicted = [_make_round()] + + with ( + patch.object(store, "_get_summary_lock", return_value=asyncio.Lock()), + patch.object(store, "get_summary", AsyncMock(return_value=None)), + patch.object(store, "save_summary", AsyncMock()), + ): + await store._run_incremental_summary("chat-4", evicted) + + assert received == [None] diff --git a/tests/test_direct_step_executor.py b/tests/test_direct_step_executor.py index 95bc582e..f10cca0b 100644 --- a/tests/test_direct_step_executor.py +++ b/tests/test_direct_step_executor.py @@ -10,8 +10,8 @@ import pytest -from src.models.request_models import OrchestrationRequest -from src.tool_classifier.workflows.service_workflow import ServiceWorkflowExecutor +from models.request_models import OrchestrationRequest +from tool_classifier.workflows.service_workflow import ServiceWorkflowExecutor def _make_request( @@ -171,6 +171,7 @@ class TestExecuteDirectStepStreaming: async def test_yields_content_and_end(self) -> None: """Valid prefix → yields exactly 2 SSE chunks (content, END).""" mock_sse = MagicMock() + mock_sse.store_streaming_inference = AsyncMock() mock_sse.format_sse = MagicMock(side_effect=["sse_content", "sse_end"]) executor = _make_executor(orchestration_service=mock_sse) @@ -188,6 +189,7 @@ async def test_yields_content_and_end(self) -> None: async def test_format_sse_called_with_buttons(self) -> None: """format_sse receives content and buttons on first call, 'END' on second.""" mock_sse = MagicMock() + mock_sse.store_streaming_inference = AsyncMock() mock_sse.format_sse = MagicMock(return_value="data: ...\n\n") executor = _make_executor(orchestration_service=mock_sse) diff --git a/tests/test_follow_up_detector.py b/tests/test_follow_up_detector.py new file mode 100644 index 00000000..135f8305 --- /dev/null +++ b/tests/test_follow_up_detector.py @@ -0,0 +1,335 @@ +"""Unit tests for FollowUpDetectorModule — DSPy follow-up query classification.""" + +import json +from collections.abc import Generator +from unittest.mock import MagicMock, patch + +import dspy +import pytest + +from tool_classifier.follow_up_detector import ( + FollowUpDetectorModule, + _validate_updated_params, +) + + +@pytest.fixture(autouse=True) +def mock_dspy_lm() -> Generator[MagicMock, None, None]: + """Mock DSPy LM to prevent 'No LM is loaded' errors during tests.""" + mock_lm = MagicMock() + mock_lm.history = [] + with patch("dspy.settings") as mock_settings: + mock_settings.lm = mock_lm + dspy.configure(lm=mock_lm) + yield mock_lm + + +def _make_mock_result( + follow_up_type: str, + updated_params: dict, +) -> MagicMock: + """Build a mock DSPy Predict result with the expected output attributes.""" + mock_result = MagicMock() + mock_result.follow_up_type = follow_up_type + mock_result.updated_params = json.dumps(updated_params, ensure_ascii=False) + return mock_result + + +_SAMPLE_SCHEMA = [ + { + "name": "year", + "type": "integer", + "required": True, + "description": "Year for the query", + }, + { + "name": "country", + "type": "string", + "required": True, + "description": "Country ISO code", + }, +] + +_SAMPLE_PARAMS = {"year": 2026, "country": "EE"} + + +class TestFollowUpDetectorModuleInit: + """FollowUpDetectorModule should initialise correctly.""" + + def test_module_has_detector_attribute(self) -> None: + module = FollowUpDetectorModule() + assert hasattr(module, "detector") + + def test_detector_is_dspy_predict(self) -> None: + module = FollowUpDetectorModule() + assert isinstance(module.detector, dspy.Predict) + + +class TestFollowUpDetectorForward: + """forward() should classify follow-up queries correctly.""" + + def test_response_question_classification(self) -> None: + """Query asking about returned data should be classified as response_question.""" + module = FollowUpDetectorModule() + mock_result = _make_mock_result("response_question", {}) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="Which of those holidays falls on a Monday?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "response_question" + assert result["updated_params"] == {} + + def test_param_update_classification(self) -> None: + """Query refining parameters should be classified as param_update with updated_params.""" + module = FollowUpDetectorModule() + mock_result = _make_mock_result("param_update", {"year": 2025}) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="What about 2025 instead?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "param_update" + assert result["updated_params"] == {"year": 2025} + + def test_new_intent_classification(self) -> None: + """Unrelated query should be classified as new_intent with empty updated_params.""" + module = FollowUpDetectorModule() + mock_result = _make_mock_result("new_intent", {}) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="How do I apply for a driving licence?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "new_intent" + assert result["updated_params"] == {} + + def test_invalid_follow_up_type_defaults_to_new_intent(self) -> None: + """An unrecognised follow_up_type value should fall back to new_intent.""" + module = FollowUpDetectorModule() + mock_result = _make_mock_result("unknown_value", {}) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="Some query", + previous_query="Previous query", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "new_intent" + assert result["updated_params"] == {} + + def test_json_parse_error_defaults_to_empty_params(self) -> None: + """Malformed JSON in updated_params for param_update should set params to {} but keep type.""" + module = FollowUpDetectorModule() + mock_result = MagicMock() + mock_result.follow_up_type = "param_update" + mock_result.updated_params = "not valid json {" + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="What about 2025?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + # Follow-up type is still recognized as param_update, but params default to {} + assert result["follow_up_type"] == "param_update" + assert result["updated_params"] == {} + + def test_exception_defaults_to_new_intent(self) -> None: + """A predictor exception should return the safe fallback without raising.""" + module = FollowUpDetectorModule() + + with patch.object(module, "detector", side_effect=RuntimeError("LLM failure")): + result = module.forward( + user_query="Some query", + previous_query="Previous query", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "new_intent" + assert result["updated_params"] == {} + + def test_invalid_updated_params_defaults_to_empty_dict(self) -> None: + """Non-dict updated_params (e.g. a JSON list) should be replaced with {}.""" + module = FollowUpDetectorModule() + mock_result = MagicMock() + mock_result.follow_up_type = "param_update" + mock_result.updated_params = json.dumps(["year", 2025]) # list, not dict + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="What about 2025?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "param_update" + assert result["updated_params"] == {} + + +class TestFollowUpDetectorOutputRulesEnforcement: + """OUTPUT RULES normalization: updated_params must be {} unless follow_up_type is param_update.""" + + def test_normalize_updated_params_to_empty_for_response_question(self) -> None: + """Verify that updated_params is forced to {} for response_question even if LLM returns non-empty. + + Per OUTPUT RULES: "For response_question or new_intent: return updated_params as empty {}" + This prevents unintended params from leaking downstream. + """ + module = FollowUpDetectorModule() + # LLM returns non-empty updated_params, but we should normalize to {} + mock_result = _make_mock_result("response_question", {"year": 2025}) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="Which of those falls on a Monday?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "response_question" + assert result["updated_params"] == {}, ( + "updated_params must be {} for response_question per OUTPUT RULES" + ) + + def test_normalize_updated_params_to_empty_for_new_intent(self) -> None: + """Verify that updated_params is forced to {} for new_intent even if LLM returns non-empty. + + Per OUTPUT RULES: "For response_question or new_intent: return updated_params as empty {}" + This prevents unintended params from leaking downstream. + """ + module = FollowUpDetectorModule() + # LLM returns non-empty updated_params, but we should normalize to {} + mock_result = _make_mock_result("new_intent", {"country": "US"}) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="How do I apply for a driving licence?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "new_intent" + assert result["updated_params"] == {}, ( + "updated_params must be {} for new_intent per OUTPUT RULES" + ) + + def test_preserve_updated_params_only_for_param_update(self) -> None: + """Verify that updated_params is only preserved when follow_up_type is param_update.""" + module = FollowUpDetectorModule() + mock_result = _make_mock_result("param_update", {"year": 2025, "country": "US"}) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="What about 2025 and US?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "param_update" + assert result["updated_params"] == {"year": 2025, "country": "US"}, ( + "updated_params should be preserved for param_update" + ) + + +class TestValidateUpdatedParams: + """_validate_updated_params should filter and validate against the schema.""" + + def test_keeps_valid_params_in_schema(self) -> None: + """Valid parameters that are in the schema should be kept.""" + updated_params = {"year": 2025, "country": "US"} + result = _validate_updated_params(updated_params, _SAMPLE_SCHEMA) + assert result == {"year": 2025, "country": "US"} + + def test_drops_unknown_params_not_in_schema(self) -> None: + """Parameters not in the schema should be dropped to prevent injection.""" + updated_params = {"year": 2025, "country": "US", "malicious_param": "value"} + result = _validate_updated_params(updated_params, _SAMPLE_SCHEMA) + assert result == {"year": 2025, "country": "US"} + assert "malicious_param" not in result + + def test_drops_multiple_unknown_params(self) -> None: + """Multiple injected parameters should all be dropped.""" + updated_params = { + "year": 2025, + "injected1": "danger1", + "injected2": "danger2", + } + result = _validate_updated_params(updated_params, _SAMPLE_SCHEMA) + assert result == {"year": 2025} + assert len(result) == 1 + + def test_empty_params_returns_empty_dict(self) -> None: + """Empty updated_params should return empty dict.""" + result = _validate_updated_params({}, _SAMPLE_SCHEMA) + assert result == {} + + def test_all_params_unknown_returns_empty_dict(self) -> None: + """If all parameters are unknown, result should be empty dict.""" + updated_params = {"unknown1": "val1", "unknown2": "val2"} + result = _validate_updated_params(updated_params, _SAMPLE_SCHEMA) + assert result == {} + + def test_partial_unknown_mixed_with_valid(self) -> None: + """Mix of valid and unknown params should keep only valid ones.""" + updated_params = { + "year": 2025, + "country": "EE", + "extra_field": "injected", + "another_injection": 123, + } + result = _validate_updated_params(updated_params, _SAMPLE_SCHEMA) + assert result == {"year": 2025, "country": "EE"} + + def test_schema_with_null_entries_handles_gracefully(self) -> None: + """Schema entries without 'name' key should be skipped without errors.""" + schema_with_invalid = [ + {"name": "year", "type": "integer", "required": True}, + {"type": "string"}, # Missing 'name' key + {"name": "country", "type": "string", "required": True}, + ] + updated_params = {"year": 2025, "country": "US"} + result = _validate_updated_params(updated_params, schema_with_invalid) + assert result == {"year": 2025, "country": "US"} + + def test_validation_inside_param_update_flow(self) -> None: + """Integration: param_update with injected params should be filtered.""" + module = FollowUpDetectorModule() + mock_result = _make_mock_result( + "param_update", + {"year": 2025, "country": "US", "injected_param": "danger"}, + ) + + with patch.object(module, "detector", return_value=mock_result): + result = module.forward( + user_query="What about 2025 and US?", + previous_query="Show public holidays in Estonia for 2026", + previous_params=_SAMPLE_PARAMS, + params_schema=_SAMPLE_SCHEMA, + ) + + assert result["follow_up_type"] == "param_update" + # Injected param should be dropped + assert result["updated_params"] == {"year": 2025, "country": "US"} + assert "injected_param" not in result["updated_params"] diff --git a/tests/test_history_integration.py b/tests/test_history_integration.py new file mode 100644 index 00000000..09d8c321 --- /dev/null +++ b/tests/test_history_integration.py @@ -0,0 +1,1466 @@ +"""Unit tests for conversation history integration in LLMOrchestrationService. + +Tests cover: +- should_save_history(): filtering logic for when to persist rounds (standalone function) +- save_history_round(): Redis persistence with error isolation (standalone function) +- _extract_content_from_sse(): SSE chunk parsing +- Non-streaming hook in process_orchestration_request() +- RAG streaming hook in _stream_rag_pipeline() (accumulated_response saved after END) +- Classifier streaming hook in stream_orchestration_response() (non-RAG workflows) +- ContextWorkflowExecutor._build_history(): Redis-first history retrieval +- ContextAnalyzer.detect_context_with_summary_fallback(): pre_computed_summary fast path +""" + +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from src.llm_orchestration_service import ( + LLMOrchestrationService, + _HISTORY_EXCLUDED_MESSAGES, +) +from src.llm_orchestrator_config.llm_ochestrator_constants import ( + INPUT_GUARDRAIL_VIOLATION_MESSAGE, + OUT_OF_SCOPE_MESSAGE, + OUTPUT_GUARDRAIL_VIOLATION_MESSAGE, + TECHNICAL_ISSUE_MESSAGE, +) +from src.models.conversation_history_models import ( + ConversationHistoryState, + ConversationRound, +) +from src.utils.conversation_history_store import should_save_history, save_history_round +from src.utils.sse_utils import extract_content_from_sse + +# Use the same import path as llm_orchestration_service.py uses internally +# (``from models.request_models import ...``) to avoid the Python dual-import +# problem where isinstance() fails across two module-path aliases. +from models.request_models import OrchestrationResponse, TestOrchestrationResponse + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_service() -> LLMOrchestrationService: + """Return a bare LLMOrchestrationService instance with __init__ bypassed.""" + svc: LLMOrchestrationService = object.__new__(LLMOrchestrationService) + svc.conversation_history_store = None # default off + return svc + + +def _make_ok_response(content: str = "Here is the answer.") -> OrchestrationResponse: + return OrchestrationResponse( + chatId="chat-1", + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content=content, + ) + + +def _make_sse(chat_id: str, content: str) -> str: + payload = { + "chatId": chat_id, + "payload": {"content": content}, + "timestamp": "1234567890", + "sentTo": [], + } + return f"data: {json.dumps(payload)}\n\n" + + +# --------------------------------------------------------------------------- +# _HISTORY_EXCLUDED_MESSAGES content +# --------------------------------------------------------------------------- + + +class TestHistoryExcludedMessages: + def test_english_out_of_scope_excluded(self): + assert OUT_OF_SCOPE_MESSAGE in _HISTORY_EXCLUDED_MESSAGES + + def test_english_technical_issue_excluded(self): + assert TECHNICAL_ISSUE_MESSAGE in _HISTORY_EXCLUDED_MESSAGES + + def test_english_input_guardrail_excluded(self): + assert INPUT_GUARDRAIL_VIOLATION_MESSAGE in _HISTORY_EXCLUDED_MESSAGES + + def test_english_output_guardrail_excluded(self): + assert OUTPUT_GUARDRAIL_VIOLATION_MESSAGE in _HISTORY_EXCLUDED_MESSAGES + + def test_normal_answer_not_excluded(self): + assert "Here is the answer." not in _HISTORY_EXCLUDED_MESSAGES + + def test_set_is_frozenset(self): + assert isinstance(_HISTORY_EXCLUDED_MESSAGES, frozenset) + + +# --------------------------------------------------------------------------- +# _should_save_history +# --------------------------------------------------------------------------- + + +class TestShouldSaveHistory: + def test_returns_false_when_store_is_none(self): + assert ( + should_save_history(None, _make_ok_response(), _HISTORY_EXCLUDED_MESSAGES) + is False + ) + + def test_returns_false_for_test_orchestration_response(self): + test_resp = TestOrchestrationResponse( + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=False, + content="answer", + ) + assert ( + should_save_history(MagicMock(), test_resp, _HISTORY_EXCLUDED_MESSAGES) + is False + ) + + def test_returns_false_when_input_guard_failed(self): + resp = OrchestrationResponse( + chatId="c1", + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=True, + content=INPUT_GUARDRAIL_VIOLATION_MESSAGE, + ) + assert ( + should_save_history(MagicMock(), resp, _HISTORY_EXCLUDED_MESSAGES) is False + ) + + def test_returns_false_when_out_of_scope(self): + resp = OrchestrationResponse( + chatId="c1", + llmServiceActive=True, + questionOutOfLLMScope=True, + inputGuardFailed=False, + content=OUT_OF_SCOPE_MESSAGE, + ) + assert ( + should_save_history(MagicMock(), resp, _HISTORY_EXCLUDED_MESSAGES) is False + ) + + def test_returns_false_when_content_is_excluded_message(self): + resp = _make_ok_response(content=TECHNICAL_ISSUE_MESSAGE) + assert ( + should_save_history(MagicMock(), resp, _HISTORY_EXCLUDED_MESSAGES) is False + ) + + def test_returns_true_for_valid_successful_response(self): + assert ( + should_save_history( + MagicMock(), _make_ok_response(), _HISTORY_EXCLUDED_MESSAGES + ) + is True + ) + + def test_returns_true_for_output_guardrail_violation_content_with_flags_false(self): + """Content match on OUTPUT_GUARDRAIL_VIOLATION_MESSAGE must also exclude.""" + resp = _make_ok_response(content=OUTPUT_GUARDRAIL_VIOLATION_MESSAGE) + assert ( + should_save_history(MagicMock(), resp, _HISTORY_EXCLUDED_MESSAGES) is False + ) + + +# --------------------------------------------------------------------------- +# save_history_round +# --------------------------------------------------------------------------- + + +class TestSaveHistoryRound: + @pytest.mark.asyncio + async def test_calls_store_save_round_with_correct_round(self): + store = AsyncMock() + + await save_history_round(store, "chat-99", "user question", "bot answer") + + store.save_round.assert_awaited_once() + call_args = store.save_round.call_args + chat_id_arg, round_arg = call_args.args + assert chat_id_arg == "chat-99" + assert isinstance(round_arg, ConversationRound) + assert round_arg.user_message == "user question" + assert round_arg.bot_message == "bot answer" + + @pytest.mark.asyncio + async def test_does_not_raise_when_store_raises(self): + store = AsyncMock() + store.save_round.side_effect = RuntimeError("Redis down") + + # Must not propagate + await save_history_round(store, "chat-1", "q", "a") + + +# --------------------------------------------------------------------------- +# extract_content_from_sse +# --------------------------------------------------------------------------- + + +class TestExtractContentFromSse: + def test_extracts_content_from_valid_chunk(self): + chunk = _make_sse("c1", "Hello world") + result = extract_content_from_sse(chunk) + assert result == "Hello world" + + def test_extracts_end_marker(self): + chunk = _make_sse("c1", "END") + result = extract_content_from_sse(chunk) + assert result == "END" + + def test_returns_none_for_non_sse_string(self): + result = extract_content_from_sse("not sse data") + assert result is None + + def test_returns_none_for_malformed_json(self): + result = extract_content_from_sse("data: {bad json}\n\n") + assert result is None + + def test_returns_none_when_payload_missing(self): + chunk = "data: " + json.dumps({"chatId": "c1", "sentTo": []}) + "\n\n" + result = extract_content_from_sse(chunk) + assert result is None + + def test_returns_none_when_content_key_missing(self): + chunk = "data: " + json.dumps({"chatId": "c1", "payload": {}}) + "\n\n" + result = extract_content_from_sse(chunk) + assert result is None + + def test_extracts_excluded_message_content(self): + chunk = _make_sse("c1", TECHNICAL_ISSUE_MESSAGE) + result = extract_content_from_sse(chunk) + assert result == TECHNICAL_ISSUE_MESSAGE + + +# --------------------------------------------------------------------------- +# Non-streaming hook: process_orchestration_request +# --------------------------------------------------------------------------- + + +class TestNonStreamingHistoryHook: + @pytest.mark.asyncio + async def test_save_called_on_successful_response(self): + store = AsyncMock() + response = _make_ok_response("The answer.") + + # Exercise the logic directly (mirrors what process_orchestration_request does) + if should_save_history(store, response, _HISTORY_EXCLUDED_MESSAGES): + await save_history_round(store, "chat-1", "my question", response.content) + + store.save_round.assert_awaited_once() + + @pytest.mark.asyncio + async def test_save_not_called_on_guardrail_blocked_response(self): + store = AsyncMock() + blocked = OrchestrationResponse( + chatId="c1", + llmServiceActive=True, + questionOutOfLLMScope=False, + inputGuardFailed=True, + content=INPUT_GUARDRAIL_VIOLATION_MESSAGE, + ) + + if should_save_history(store, blocked, _HISTORY_EXCLUDED_MESSAGES): + await save_history_round(store, "c1", "bad query", blocked.content) + + store.save_round.assert_not_awaited() + + @pytest.mark.asyncio + async def test_save_not_called_on_out_of_scope_response(self): + store = AsyncMock() + oos = OrchestrationResponse( + chatId="c1", + llmServiceActive=True, + questionOutOfLLMScope=True, + inputGuardFailed=False, + content=OUT_OF_SCOPE_MESSAGE, + ) + + if should_save_history(store, oos, _HISTORY_EXCLUDED_MESSAGES): + await save_history_round(store, "c1", "obscure query", oos.content) + + store.save_round.assert_not_awaited() + + @pytest.mark.asyncio + async def test_save_not_called_when_store_is_none(self): + response = _make_ok_response() + + if should_save_history(None, response, _HISTORY_EXCLUDED_MESSAGES): + await save_history_round(None, "c1", "q", response.content) + + # When store is None, save_history_round returns early without calling save_round + + +# --------------------------------------------------------------------------- +# RAG streaming hook: _stream_rag_pipeline logic +# --------------------------------------------------------------------------- + + +class TestRagStreamingHistoryHook: + @pytest.mark.asyncio + async def test_save_called_with_accumulated_response_after_successful_stream(self): + """Simulate the relevant section of _stream_rag_pipeline after streaming.""" + store = AsyncMock() + + accumulated_response = ["Hello", " world", " from", " RAG."] + _rag_bot_message = "".join(accumulated_response) + + # Mirrors the hook added to _stream_rag_pipeline + if store is not None: + if _rag_bot_message not in _HISTORY_EXCLUDED_MESSAGES: + await save_history_round( + store, "chat-5", "what is RAG?", _rag_bot_message + ) + + store.save_round.assert_awaited_once() + _, round_arg = store.save_round.call_args.args + assert round_arg.bot_message == "Hello world from RAG." + assert round_arg.user_message == "what is RAG?" + + @pytest.mark.asyncio + async def test_save_skipped_when_accumulated_is_excluded_message(self): + """OOS/violations yielded as single chunks must not be saved.""" + store = AsyncMock() + + # Simulate what happens when OOS message ends up in accumulated_response + accumulated_response = [OUT_OF_SCOPE_MESSAGE] + _rag_bot_message = "".join(accumulated_response) + + if store is not None: + if _rag_bot_message not in _HISTORY_EXCLUDED_MESSAGES: + await save_history_round(store, "chat-5", "q", _rag_bot_message) + + store.save_round.assert_not_awaited() + + @pytest.mark.asyncio + async def test_save_skipped_when_store_is_none(self): + store = None + + accumulated_response = ["some", " answer"] + _rag_bot_message = "".join(accumulated_response) + + if store is not None: + if _rag_bot_message not in _HISTORY_EXCLUDED_MESSAGES: + await save_history_round(store, "c1", "q", _rag_bot_message) + + # When store is None, save_history_round returns early without doing anything + + +# --------------------------------------------------------------------------- +# Classifier streaming hook: stream_orchestration_response accumulation logic +# --------------------------------------------------------------------------- + + +class TestClassifierStreamingHistoryHook: + @pytest.mark.asyncio + async def test_accumulates_non_end_non_excluded_content(self): + store = AsyncMock() + + tokens = ["The ", "answer ", "is 42."] + sse_chunks = [_make_sse("c1", t) for t in tokens] + sse_chunks.append(_make_sse("c1", "END")) + + # Mirrors the classifier streaming accumulation logic + _save_classifier_history = store is not None + _classifier_accumulated: list[str] = [] + + for sse_chunk in sse_chunks: + if _save_classifier_history: + extracted = extract_content_from_sse(sse_chunk) + if ( + extracted is not None + and extracted != "END" + and extracted not in _HISTORY_EXCLUDED_MESSAGES + ): + _classifier_accumulated.append(extracted) + + if _save_classifier_history and _classifier_accumulated: + await save_history_round( + store, "c1", "what is 42?", "".join(_classifier_accumulated) + ) + + store.save_round.assert_awaited_once() + + @pytest.mark.asyncio + async def test_does_not_save_when_only_violation_message_streamed(self): + store = AsyncMock() + + sse_chunks = [ + _make_sse("c1", INPUT_GUARDRAIL_VIOLATION_MESSAGE), + _make_sse("c1", "END"), + ] + + _save_classifier_history = store is not None + _classifier_accumulated: list[str] = [] + + for sse_chunk in sse_chunks: + if _save_classifier_history: + extracted = extract_content_from_sse(sse_chunk) + if ( + extracted is not None + and extracted != "END" + and extracted not in _HISTORY_EXCLUDED_MESSAGES + ): + _classifier_accumulated.append(extracted) + + if _save_classifier_history and _classifier_accumulated: + await save_history_round( + store, "c1", "bad q", "".join(_classifier_accumulated) + ) + + store.save_round.assert_not_awaited() + + @pytest.mark.asyncio + async def test_does_not_save_for_rag_workflow(self): + """RAG workflow has its own hook in _stream_rag_pipeline; skip classifier hook.""" + from src.tool_classifier import WorkflowType + + store = AsyncMock() + sse_chunks = [_make_sse("c1", "token"), _make_sse("c1", "END")] + + # Mirrors: _save_classifier_history = store is not None AND workflow != RAG + workflow_type = WorkflowType.RAG + _save_classifier_history = ( + store is not None and workflow_type != WorkflowType.RAG # False for RAG + ) + _classifier_accumulated: list[str] = [] + + for sse_chunk in sse_chunks: + if _save_classifier_history: + extracted = extract_content_from_sse(sse_chunk) + if extracted is not None and extracted != "END": + _classifier_accumulated.append(extracted) + + if _save_classifier_history and _classifier_accumulated: + await save_history_round(store, "c1", "q", "".join(_classifier_accumulated)) + + store.save_round.assert_not_awaited() + + @pytest.mark.asyncio + async def test_saves_for_non_rag_workflow(self): + """Non-RAG workflows (SERVICE, API_TOOL, etc.) should save via classifier hook.""" + from src.tool_classifier import WorkflowType + + store = AsyncMock() + tokens = ["answer"] + sse_chunks = [_make_sse("c1", t) for t in tokens] + sse_chunks.append(_make_sse("c1", "END")) + + # Use SERVICE workflow (not RAG) so the gate passes + workflow_type = WorkflowType.SERVICE + _save_classifier_history = ( + store is not None and workflow_type != WorkflowType.RAG # True for SERVICE + ) + _classifier_accumulated: list[str] = [] + + for sse_chunk in sse_chunks: + if _save_classifier_history: + extracted = extract_content_from_sse(sse_chunk) + if ( + extracted is not None + and extracted != "END" + and extracted not in _HISTORY_EXCLUDED_MESSAGES + ): + _classifier_accumulated.append(extracted) + + if _save_classifier_history and _classifier_accumulated: + await save_history_round(store, "c1", "q", "".join(_classifier_accumulated)) + + store.save_round.assert_awaited_once() + + @pytest.mark.asyncio + async def test_does_not_save_when_store_is_none(self): + store = None + sse_chunks = [_make_sse("c1", "answer"), _make_sse("c1", "END")] + _save_classifier_history = store is not None + _classifier_accumulated: list[str] = [] + + for sse_chunk in sse_chunks: + if _save_classifier_history: + extracted = extract_content_from_sse(sse_chunk) + if extracted is not None and extracted != "END": + _classifier_accumulated.append(extracted) + + if _save_classifier_history and _classifier_accumulated: + await save_history_round(store, "c1", "q", "".join(_classifier_accumulated)) + + # When store is None, save_history_round returns early without doing anything + + +# --------------------------------------------------------------------------- +# Helpers shared by the new test suites below +# --------------------------------------------------------------------------- + + +def _make_llm_manager() -> MagicMock: + mgr = MagicMock() + mgr.ensure_global_config = MagicMock() + mgr.use_task_local = MagicMock() + return mgr + + +def _make_orchestration_request( + chat_id: str = "chat-1", + message: str = "What did you say earlier?", + history: list | None = None, +) -> MagicMock: + """Return a lightweight mock that mimics the fields accessed by _build_history.""" + req = MagicMock() + req.chatId = chat_id + req.message = message + req.conversationHistory = history or [] + return req + + +def _make_round( + user: str = "What is the tax rate?", + bot: str = "The tax rate is 20%.", + ts: float = 1_700_000_000.0, +) -> ConversationRound: + return ConversationRound(user_message=user, bot_message=bot, timestamp=ts) + + +# --------------------------------------------------------------------------- +# ContextWorkflowExecutor._build_history — Redis-first retrieval +# --------------------------------------------------------------------------- + + +class TestBuildHistoryRedisFirst: + """Tests for the async _build_history method in ContextWorkflowExecutor.""" + + @pytest.mark.asyncio + async def test_uses_redis_rounds_when_available(self) -> None: + """When Redis has rounds, _build_history returns them and ignores request history.""" + from src.tool_classifier.workflows.context_workflow import ( + ContextWorkflowExecutor, + ) + + round_ = _make_round() + state = ConversationHistoryState( + chat_id="chat-1", rounds=[round_], summary=None + ) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + workflow = ContextWorkflowExecutor( + llm_manager=_make_llm_manager(), + conversation_history_store=store, + ) + req = _make_orchestration_request(chat_id="chat-1") + + history, summary = await workflow._build_history(req) + + assert len(history) == 2 # one round → two messages + assert history[0]["authorRole"] == "user" + assert history[0]["message"] == round_.user_message + assert history[1]["authorRole"] == "bot" + assert history[1]["message"] == round_.bot_message + assert summary is None + + @pytest.mark.asyncio + async def test_returns_redis_summary_with_rounds(self) -> None: + """Summary stored in Redis is returned as pre_computed_summary.""" + from src.tool_classifier.workflows.context_workflow import ( + ContextWorkflowExecutor, + ) + + round_ = _make_round() + state = ConversationHistoryState( + chat_id="chat-1", + rounds=[round_], + summary="Earlier we discussed tax rates.", + ) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + workflow = ContextWorkflowExecutor( + llm_manager=_make_llm_manager(), + conversation_history_store=store, + ) + req = _make_orchestration_request(chat_id="chat-1") + + history, summary = await workflow._build_history(req) + + assert len(history) == 2 + assert summary == "Earlier we discussed tax rates." + + @pytest.mark.asyncio + async def test_falls_back_to_request_when_redis_raises(self) -> None: + """When get_context() raises, _build_history falls back to request.conversationHistory.""" + from src.tool_classifier.workflows.context_workflow import ( + ContextWorkflowExecutor, + ) + + store = AsyncMock() + store.get_context = AsyncMock(side_effect=RuntimeError("Redis down")) + + workflow = ContextWorkflowExecutor( + llm_manager=_make_llm_manager(), + conversation_history_store=store, + ) + + item = MagicMock() + item.authorRole = "user" + item.message = "fallback message" + item.timestamp = "2024-01-01T00:00:00" + + req = _make_orchestration_request(chat_id="chat-1", history=[item]) + + history, summary = await workflow._build_history(req) + + assert len(history) == 1 + assert history[0]["message"] == "fallback message" + assert summary is None + + @pytest.mark.asyncio + async def test_falls_back_to_request_when_store_is_none(self) -> None: + """When conversation_history_store is None, request history is used.""" + from src.tool_classifier.workflows.context_workflow import ( + ContextWorkflowExecutor, + ) + + workflow = ContextWorkflowExecutor( + llm_manager=_make_llm_manager(), + conversation_history_store=None, + ) + + item = MagicMock() + item.authorRole = "bot" + item.message = "bot reply" + item.timestamp = "2024-01-01T00:00:01" + + req = _make_orchestration_request(chat_id="chat-1", history=[item]) + + history, summary = await workflow._build_history(req) + + assert len(history) == 1 + assert history[0]["authorRole"] == "bot" + assert summary is None + + @pytest.mark.asyncio + async def test_falls_back_to_request_when_redis_returns_empty_rounds(self) -> None: + """Redis state with no rounds → fall back to request.conversationHistory.""" + from src.tool_classifier.workflows.context_workflow import ( + ContextWorkflowExecutor, + ) + + state = ConversationHistoryState( + chat_id="chat-1", rounds=[], summary="old summary" + ) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + workflow = ContextWorkflowExecutor( + llm_manager=_make_llm_manager(), + conversation_history_store=store, + ) + + item = MagicMock() + item.authorRole = "user" + item.message = "from request" + item.timestamp = "2024-01-01T00:00:00" + + req = _make_orchestration_request(chat_id="chat-1", history=[item]) + + history, summary = await workflow._build_history(req) + + # Empty Redis rounds → fall back to request; summary not returned + assert len(history) == 1 + assert history[0]["message"] == "from request" + assert summary is None + + +# --------------------------------------------------------------------------- +# ContextAnalyzer.detect_context_with_summary_fallback — pre_computed_summary +# --------------------------------------------------------------------------- + + +class TestDetectContextWithPrecomputedSummary: + """Tests for the pre_computed_summary fast-path in detect_context_with_summary_fallback.""" + + def _make_analyzer(self) -> object: + from src.tool_classifier.context_analyzer import ContextAnalyzer + + return ContextAnalyzer(_make_llm_manager()) + + def _no_answer_detection(self) -> tuple: + from src.tool_classifier.context_analyzer import ContextDetectionResult + + result = ContextDetectionResult( + is_greeting=False, + can_answer_from_context=False, + reasoning="cannot answer", + ) + cost: dict = {"total_cost": 0.001, "total_tokens": 10, "num_calls": 1} + return result, cost + + def _answer_from_summary(self) -> tuple: + from src.tool_classifier.context_analyzer import ContextAnalysisResult + + result = ContextAnalysisResult( + is_greeting=False, + can_answer_from_context=True, + answer="The tax rate is 20%.", + reasoning="found in summary", + ) + cost: dict = {"total_cost": 0.002, "total_tokens": 20, "num_calls": 1} + return result, cost + + @pytest.mark.asyncio + async def test_skips_generate_summary_when_pre_computed_provided(self) -> None: + """_generate_conversation_summary must NOT be called when pre_computed_summary is set.""" + analyzer = self._make_analyzer() + + with ( + patch.object( + analyzer, + "detect_context", + new_callable=AsyncMock, + return_value=self._no_answer_detection(), + ), + patch.object( + analyzer, + "_generate_conversation_summary", + new_callable=AsyncMock, + ) as mock_generate, + patch.object( + analyzer, + "_analyze_from_summary", + new_callable=AsyncMock, + return_value=self._answer_from_summary(), + ), + ): + await analyzer.detect_context_with_summary_fallback( + query="What was the tax rate?", + conversation_history=[], + pre_computed_summary="Tax rate is 20%.", + ) + + mock_generate.assert_not_awaited() + + @pytest.mark.asyncio + async def test_still_runs_analyze_from_summary_when_pre_computed_provided( + self, + ) -> None: + """_analyze_from_summary IS called even when the summary comes from Redis.""" + analyzer = self._make_analyzer() + + with ( + patch.object( + analyzer, + "detect_context", + new_callable=AsyncMock, + return_value=self._no_answer_detection(), + ), + patch.object( + analyzer, + "_generate_conversation_summary", + new_callable=AsyncMock, + ), + patch.object( + analyzer, + "_analyze_from_summary", + new_callable=AsyncMock, + return_value=self._answer_from_summary(), + ) as mock_analyze, + ): + await analyzer.detect_context_with_summary_fallback( + query="What was the tax rate?", + conversation_history=[], + pre_computed_summary="Tax rate is 20%.", + ) + + mock_analyze.assert_awaited_once() + call_kwargs = mock_analyze.call_args.kwargs + assert call_kwargs["summary"] == "Tax rate is 20%." + + @pytest.mark.asyncio + async def test_returns_answer_from_pre_computed_summary(self) -> None: + """When the summary analysis succeeds, result has answered_from_summary=True.""" + from src.tool_classifier.context_analyzer import ContextDetectionResult + + analyzer = self._make_analyzer() + + with ( + patch.object( + analyzer, + "detect_context", + new_callable=AsyncMock, + return_value=self._no_answer_detection(), + ), + patch.object( + analyzer, "_generate_conversation_summary", new_callable=AsyncMock + ), + patch.object( + analyzer, + "_analyze_from_summary", + new_callable=AsyncMock, + return_value=self._answer_from_summary(), + ), + ): + result, _ = await analyzer.detect_context_with_summary_fallback( + query="What was the tax rate?", + conversation_history=[], + pre_computed_summary="Tax rate is 20%.", + ) + + assert isinstance(result, ContextDetectionResult) + assert result.can_answer_from_context is True + assert result.answered_from_summary is True + assert result.context_snippet == "The tax rate is 20%." + + @pytest.mark.asyncio + async def test_summary_path_attempted_for_short_history_with_pre_computed( + self, + ) -> None: + """Summary analysis runs even with <=10 turns when pre_computed_summary is set.""" + analyzer = self._make_analyzer() + + # Only 2 items in history (well below 10) + short_history = [ + {"authorRole": "user", "message": "hi", "timestamp": "0"}, + {"authorRole": "bot", "message": "hello", "timestamp": "1"}, + ] + + with ( + patch.object( + analyzer, + "detect_context", + new_callable=AsyncMock, + return_value=self._no_answer_detection(), + ), + patch.object( + analyzer, "_generate_conversation_summary", new_callable=AsyncMock + ) as mock_gen, + patch.object( + analyzer, + "_analyze_from_summary", + new_callable=AsyncMock, + return_value=self._answer_from_summary(), + ) as mock_analyze, + ): + await analyzer.detect_context_with_summary_fallback( + query="What was the tax rate?", + conversation_history=short_history, + pre_computed_summary="Tax rate was discussed previously.", + ) + + mock_gen.assert_not_awaited() # skipped because pre_computed_summary is set + mock_analyze.assert_awaited_once() # still validated against query + + +# --------------------------------------------------------------------------- +# End-to-end: no LLM summarisation when Redis supplies a summary +# --------------------------------------------------------------------------- + + +class TestEndToEndNoLLMSummarisationWhenRedisHasSummary: + """Verify that _generate_conversation_summary is never called when the + workflow retrieves a summary from Redis via _build_history.""" + + @pytest.mark.asyncio + async def test_no_summarisation_call_when_redis_has_summary(self) -> None: + """Full _detect() path: Redis summary present → no LLM summarisation.""" + from src.tool_classifier.context_analyzer import ( + ContextAnalysisResult, + ContextDetectionResult, + ) + from src.tool_classifier.workflows.context_workflow import ( + ContextWorkflowExecutor, + ) + + # Redis returns one round + a running summary + round_ = _make_round() + state = ConversationHistoryState( + chat_id="chat-e2e", + rounds=[round_], + summary="We discussed tax rates earlier.", + ) + store = AsyncMock() + store.get_context = AsyncMock(return_value=state) + + # detect_context returns "cannot answer" so the summary path is tried + cannot_answer = ContextDetectionResult( + is_greeting=False, + can_answer_from_context=False, + reasoning="not in recent history", + ) + + llm_manager = _make_llm_manager() + workflow = ContextWorkflowExecutor( + llm_manager=llm_manager, + conversation_history_store=store, + ) + + with ( + patch.object( + workflow.context_analyzer, + "detect_context", + new_callable=AsyncMock, + return_value=( + cannot_answer, + {"total_cost": 0.001, "total_tokens": 10, "num_calls": 1}, + ), + ), + patch.object( + workflow.context_analyzer, + "_generate_conversation_summary", + new_callable=AsyncMock, + ) as mock_gen, + patch.object( + workflow.context_analyzer, + "_analyze_from_summary", + new_callable=AsyncMock, + return_value=( + ContextAnalysisResult( + is_greeting=False, + can_answer_from_context=True, + answer="The tax rate is 20%.", + reasoning="from summary", + ), + {"total_cost": 0.002, "total_tokens": 20, "num_calls": 1}, + ), + ), + ): + time_metric: dict = {} + costs_metric: dict = {} + history, pre_computed_summary = await workflow._build_history( + _make_orchestration_request(chat_id="chat-e2e") + ) + result = await workflow._detect( + message="What was the tax rate?", + history=history, + time_metric=time_metric, + costs_metric=costs_metric, + pre_computed_summary=pre_computed_summary, + ) + + mock_gen.assert_not_awaited() + assert result is not None + assert result.can_answer_from_context is True + assert result.answered_from_summary is True + + +# --------------------------------------------------------------------------- +# RAG Workflow: Redis history used in _refine_user_prompt +# --------------------------------------------------------------------------- + + +class TestRagWorkflowUsesRedisHistory: + """Verify that the RAG pipeline fetches history from Redis before refinement.""" + + @pytest.mark.asyncio + async def test_redis_history_passed_to_refine_when_available(self) -> None: + """get_conversation_history is called and its result replaces request history.""" + + svc = _make_service() + svc.conversation_history_store = AsyncMock() + + with patch( + "src.llm_orchestration_service.get_conversation_history", + new_callable=AsyncMock, + return_value=([], None), + ) as mock_get_history: + # Stub out the rest of the pipeline so we only test the history fetch + svc._refine_user_prompt = MagicMock( + return_value=( + MagicMock( + original_question="q", + refined_questions=["q1"], + ), + {}, + ) + ) + svc._safe_retrieve_contextual_chunks = AsyncMock(return_value=[]) + svc.format_sse = MagicMock(return_value="data: {}\n\n") + + components = { + "llm_manager": _make_llm_manager(), + "contextual_retriever": AsyncMock(), + "response_generator": MagicMock(), + "guardrails_adapter": None, + } + stream_ctx = MagicMock() + stream_ctx.stream_id = "sid" + stream_ctx.mark_completed = MagicMock() + + request = _make_orchestration_request(chat_id="chat-rag") + + # Drain the generator to trigger the history fetch + async for _ in svc._stream_rag_pipeline( + request=request, + components=components, + stream_ctx=stream_ctx, + costs_metric={}, + time_metric={}, + ): + pass + + mock_get_history.assert_awaited_once_with( + chat_id="chat-rag", + store=svc.conversation_history_store, + fallback=request.conversationHistory, + ) + + @pytest.mark.asyncio + async def test_redis_fallback_used_when_store_is_none(self) -> None: + """When conversation_history_store is None, request history is used.""" + svc = _make_service() + svc.conversation_history_store = None + + with patch( + "src.llm_orchestration_service.get_conversation_history", + new_callable=AsyncMock, + return_value=([], None), + ) as mock_get_history: + svc._refine_user_prompt = MagicMock( + return_value=( + MagicMock(original_question="q", refined_questions=["q1"]), + {}, + ) + ) + svc._safe_retrieve_contextual_chunks = AsyncMock(return_value=[]) + svc.format_sse = MagicMock(return_value="data: {}\n\n") + + request = _make_orchestration_request(chat_id="chat-rag-fallback") + components = { + "llm_manager": _make_llm_manager(), + "contextual_retriever": AsyncMock(), + "response_generator": MagicMock(), + "guardrails_adapter": None, + } + stream_ctx = MagicMock() + stream_ctx.stream_id = "sid" + stream_ctx.mark_completed = MagicMock() + + # Drain the generator to trigger the history fetch + async for _ in svc._stream_rag_pipeline( + request=request, + components=components, + stream_ctx=stream_ctx, + costs_metric={}, + time_metric={}, + ): + pass + + # Fallback is request.conversationHistory + mock_get_history.assert_awaited_once_with( + chat_id="chat-rag-fallback", + store=None, + fallback=request.conversationHistory, + ) + + +class TestRefineUserPromptSummary: + """Verify that _refine_user_prompt prepends summary as a system turn.""" + + def test_summary_prepended_to_dspy_history(self) -> None: + """When conversation_summary is provided, a system message is first in history.""" + svc = _make_service() + svc.langfuse_config = MagicMock() + svc.langfuse_config.langfuse_client = None + + captured_history: list = [] + + class _FakeRefiner: + def forward_structured( + self, history: list, question: str, **_: object + ) -> dict: + captured_history.extend(history) + return { + "original_question": question, + "refined_questions": [question], + "usage": {}, + "module_info": {}, + } + + llm_manager = _make_llm_manager() + llm_manager.use_task_local = MagicMock() + llm_manager.use_task_local.return_value.__enter__ = MagicMock(return_value=None) + llm_manager.use_task_local.return_value.__exit__ = MagicMock(return_value=False) + + with patch( + "src.llm_orchestration_service.PromptRefinerAgent", + return_value=_FakeRefiner(), + ): + svc._refine_user_prompt( + llm_manager=llm_manager, + original_message="What is the rate?", + conversation_history=[], + conversation_summary="We discussed tax earlier.", + ) + + assert len(captured_history) >= 1 + first = captured_history[0] + assert first["role"] == "system" + assert "We discussed tax earlier." in first["content"] + + def test_no_summary_entry_when_summary_is_none(self) -> None: + """When conversation_summary is None, no system turn is prepended.""" + svc = _make_service() + svc.langfuse_config = MagicMock() + svc.langfuse_config.langfuse_client = None + + captured_history: list = [] + + class _FakeRefiner: + def forward_structured( + self, history: list, question: str, **_: object + ) -> dict: + captured_history.extend(history) + return { + "original_question": question, + "refined_questions": [question], + "usage": {}, + "module_info": {}, + } + + llm_manager = _make_llm_manager() + + with patch( + "src.llm_orchestration_service.PromptRefinerAgent", + return_value=_FakeRefiner(), + ): + svc._refine_user_prompt( + llm_manager=llm_manager, + original_message="What is the rate?", + conversation_history=[], + conversation_summary=None, + ) + + assert all(item.get("role") != "system" for item in captured_history) + + +# --------------------------------------------------------------------------- +# Service Workflow: Redis history used in _process_intent_detection +# --------------------------------------------------------------------------- + + +class TestServiceWorkflowUsesRedisHistory: + """Verify ServiceWorkflowExecutor fetches history from Redis before intent detection.""" + + @pytest.mark.asyncio + async def test_get_conversation_history_called_in_process_intent_detection( + self, + ) -> None: + """get_conversation_history is called with the history store from the service.""" + from src.tool_classifier.workflows.service_workflow import ( + ServiceWorkflowExecutor, + ) + + orchestration_service = MagicMock() + history_store = AsyncMock() + orchestration_service.conversation_history_store = history_store + + executor = ServiceWorkflowExecutor( + llm_manager=_make_llm_manager(), + orchestration_service=orchestration_service, + ) + + with ( + patch( + "src.tool_classifier.workflows.service_workflow.get_conversation_history", + new_callable=AsyncMock, + return_value=([], None), + ) as mock_get_history, + patch.object( + executor, + "_detect_service_intent", + new_callable=AsyncMock, + return_value=(None, {}), + ), + ): + request = _make_orchestration_request(chat_id="chat-svc") + await executor._process_intent_detection( + services=[], + request=request, + chat_id="chat-svc", + context={}, + costs_metric={}, + ) + + mock_get_history.assert_awaited_once_with( + chat_id="chat-svc", + store=history_store, + fallback=request.conversationHistory, + ) + + @pytest.mark.asyncio + async def test_summary_passed_to_detect_service_intent(self) -> None: + """Summary from Redis is forwarded to _detect_service_intent.""" + from src.tool_classifier.workflows.service_workflow import ( + ServiceWorkflowExecutor, + ) + + executor = ServiceWorkflowExecutor(llm_manager=_make_llm_manager()) + + with ( + patch( + "src.tool_classifier.workflows.service_workflow.get_conversation_history", + new_callable=AsyncMock, + return_value=([], "Earlier we discussed registration."), + ), + patch.object( + executor, + "_detect_service_intent", + new_callable=AsyncMock, + return_value=(None, {}), + ) as mock_detect, + ): + request = _make_orchestration_request(chat_id="chat-svc2") + await executor._process_intent_detection( + services=[], + request=request, + chat_id="chat-svc2", + context={}, + costs_metric={}, + ) + + _, kwargs = mock_detect.call_args + assert ( + kwargs.get("conversation_summary") == "Earlier we discussed registration." + ) + + @pytest.mark.asyncio + async def test_get_conversation_history_store_returns_none_when_no_service( + self, + ) -> None: + """_get_conversation_history_store returns None when orchestration_service is None.""" + from src.tool_classifier.workflows.service_workflow import ( + ServiceWorkflowExecutor, + ) + + executor = ServiceWorkflowExecutor() + assert executor._get_conversation_history_store() is None + + @pytest.mark.asyncio + async def test_summary_prepended_to_history_dicts_in_detect_service_intent( + self, + ) -> None: + """When conversation_summary is set, a 'system' message is first in history_dicts.""" + from src.tool_classifier.workflows.service_workflow import ( + ServiceWorkflowExecutor, + ) + import dspy + + executor = ServiceWorkflowExecutor(llm_manager=_make_llm_manager()) + captured: list = [] + + class _FakeModule: + def forward( + self, + user_query: str, + services: list, + conversation_history: list | None = None, + ) -> dict: + if conversation_history: + captured.extend(conversation_history) + return { + "matched_service_id": None, + "confidence": 0.0, + "entities": {}, + "reasoning": "", + } + + with ( + patch( + "src.tool_classifier.workflows.service_workflow.IntentDetectionModule", + return_value=_FakeModule(), + ), + patch.object(executor.llm_manager, "ensure_global_config"), + patch.object(executor.llm_manager, "use_task_local"), + ): + # Patch dspy.settings.lm to avoid NoneType + mock_lm = MagicMock() + mock_lm.history = [] + with patch.object(dspy, "settings", MagicMock(lm=mock_lm)): + await executor._detect_service_intent( + user_query="register my car", + services=[], + conversation_history=[], + chat_id="c1", + conversation_summary="User was asking about vehicle registration.", + ) + + assert len(captured) >= 1 + assert captured[0]["authorRole"] == "system" + assert "vehicle registration" in captured[0]["message"] + + +# --------------------------------------------------------------------------- +# ATC Workflow: Redis history used in _compute_loop_step +# --------------------------------------------------------------------------- + + +class TestATCWorkflowUsesRedisHistory: + """Verify APIToolWorkflowExecutor fetches history from Redis on turns > 0.""" + + def _make_executor( + self, + history_store: object = None, + ) -> object: + from src.tool_classifier.workflows.api_tool_workflow import ( + APIToolWorkflowExecutor, + ) + + orchestration_service = MagicMock() + orchestration_service.conversation_history_store = history_store + orchestration_service.session_store = None + orchestration_service.prompt_config_loader = None + executor = APIToolWorkflowExecutor(orchestration_service=orchestration_service) + return executor + + def test_get_conversation_history_store_returns_store(self) -> None: + """_get_conversation_history_store returns the store from orchestration_service.""" + from src.tool_classifier.workflows.api_tool_workflow import ( + APIToolWorkflowExecutor, + ) + + store = AsyncMock() + orchestration_service = MagicMock() + orchestration_service.conversation_history_store = store + executor = APIToolWorkflowExecutor(orchestration_service=orchestration_service) + assert executor._get_conversation_history_store() is store + + def test_get_conversation_history_store_returns_none_when_no_service(self) -> None: + """_get_conversation_history_store returns None when orchestration_service is None.""" + from src.tool_classifier.workflows.api_tool_workflow import ( + APIToolWorkflowExecutor, + ) + + executor = APIToolWorkflowExecutor(orchestration_service=None) + assert executor._get_conversation_history_store() is None + + @pytest.mark.asyncio + async def test_no_redis_call_on_turn_zero(self) -> None: + """On turn 0 get_conversation_history must NOT be called.""" + from src.models.session_models import APIToolSession + from src.tool_classifier.enums import ExecutionMode + + executor = self._make_executor() + + session = MagicMock(spec=APIToolSession) + session.turn_count = 0 + session.selected_endpoint = {"name": "test_endpoint", "params": []} + session.collected_params = {} + session.max_turns = 5 + session.awaiting_continuation = False + session.detected_language = "en" + session.original_query = "book a slot" + session.execution_mode = ExecutionMode.SINGLE.value + session.parallel_endpoints = [] + + with ( + patch( + "src.tool_classifier.workflows.api_tool_workflow.get_conversation_history", + new_callable=AsyncMock, + ) as mock_get_history, + patch.object(executor, "_get_session_store", return_value=None), + patch.object( + executor, + "_get_custom_instructions", + new_callable=AsyncMock, + return_value="", + ), + patch.object( + executor, + "_build_agentic_loop", + return_value=MagicMock( + stream_run_turn=AsyncMock( + return_value=( + MagicMock( + status=MagicMock(value="NEEDS_INPUT"), + clarifying_question="What time?", + turn_count=1, + ), + ["What", " time", "?"], + ) + ) + ), + ), + ): + # Force session_store.get to return our mocked session + with patch.object( + executor, + "_get_session_store", + return_value=MagicMock( + get=AsyncMock(return_value=session), + delete=AsyncMock(), + ), + ): + request = _make_orchestration_request(chat_id="chat-atc-t0") + await executor._compute_loop_step( + request=request, + context={"matched_endpoint": {"name": "ep", "params": []}}, + ) + + mock_get_history.assert_not_awaited() + + @pytest.mark.asyncio + async def test_redis_history_fetched_on_turn_greater_than_zero(self) -> None: + """On turn > 0 get_conversation_history IS called.""" + from src.models.session_models import APIToolSession + from src.tool_classifier.enums import ExecutionMode + + executor = self._make_executor() + + session = MagicMock(spec=APIToolSession) + session.turn_count = 1 + session.selected_endpoint = {"name": "test_endpoint", "params": []} + session.collected_params = {} + session.max_turns = 5 + session.awaiting_continuation = False + session.detected_language = "en" + session.original_query = "book a slot" + session.execution_mode = ExecutionMode.SINGLE.value + session.parallel_endpoints = [] + + with ( + patch( + "src.tool_classifier.workflows.api_tool_workflow.get_conversation_history", + new_callable=AsyncMock, + return_value=([], None), + ) as mock_get_history, + patch.object( + executor, + "_get_custom_instructions", + new_callable=AsyncMock, + return_value="", + ), + patch.object( + executor, + "_build_agentic_loop", + return_value=MagicMock( + stream_run_turn=AsyncMock( + return_value=( + MagicMock( + status=MagicMock(value="NEEDS_INPUT"), + clarifying_question="What time?", + turn_count=2, + ), + ["What", " time", "?"], + ) + ) + ), + ), + patch.object( + executor, + "_get_session_store", + return_value=MagicMock( + get=AsyncMock(return_value=session), + delete=AsyncMock(), + ), + ), + ): + request = _make_orchestration_request(chat_id="chat-atc-t1") + await executor._compute_loop_step( + request=request, + context={"matched_endpoint": {"name": "ep", "params": []}}, + ) + + mock_get_history.assert_awaited_once() + _, kwargs = mock_get_history.call_args + assert kwargs["chat_id"] == "chat-atc-t1" diff --git a/tests/test_multi_agentic_loop.py b/tests/test_multi_agentic_loop.py new file mode 100644 index 00000000..3e0d44f0 --- /dev/null +++ b/tests/test_multi_agentic_loop.py @@ -0,0 +1,2217 @@ +"""Unit tests for MultiEndpointAgenticLoop.""" + +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from models.session_models import EndpointSessionState +from tool_classifier.enums import AgenticLoopStatus +from tool_classifier.multi_agentic_loop import MultiEndpointAgenticLoop +from tool_classifier.param_extractor import ParamExtractionResult + + +# --------------------------------------------------------------------------- +# Helpers / Fixtures +# --------------------------------------------------------------------------- + +_CHAT_ID = "test-multi-chat-1" + +_HISTORY: List[Dict[str, Any]] = [ + { + "authorRole": "user", + "message": "Show me the weather and public holidays for Estonia", + }, + {"authorRole": "bot", "message": "I can help with both!"}, +] + +# Endpoint A — get_public_holidays: languageIsoCode (required, shared with B), +# countryIsoCode + validFrom + validTo (required, unique to A) +_ENDPOINT_A: Dict[str, Any] = { + "name": "get_public_holidays", + "url": "https://openholidaysapi.org/PublicHolidays", + "params": [ + { + "name": "languageIsoCode", + "type": "string", + "required": True, + "description": "Response language code (e.g. ET, EN)", + }, + { + "name": "countryIsoCode", + "type": "string", + "required": True, + "description": "Two-letter country ISO code (e.g. EE, LV)", + }, + { + "name": "validFrom", + "type": "date", + "required": True, + "description": "Start date (YYYY-MM-DD)", + }, + { + "name": "validTo", + "type": "date", + "required": True, + "description": "End date (YYYY-MM-DD)", + }, + ], +} + +# Endpoint B — get_current_weather: languageIsoCode (required, shared with A), +# station (required, unique to B) +_ENDPOINT_B: Dict[str, Any] = { + "name": "get_current_weather", + "url": "https://publicapi.envir.ee/v1/combinedWeatherData", + "params": [ + { + "name": "languageIsoCode", + "type": "string", + "required": True, + "description": "Response language code", + }, + { + "name": "station", + "type": "string", + "required": True, + "description": "Weather station identifier", + }, + ], +} + +# Endpoint C — get_current_weather with no required params +_ENDPOINT_C: Dict[str, Any] = { + "name": "get_current_weather", + "url": "https://publicapi.envir.ee/v1/combinedWeatherData", + "params": [ + { + "name": "station", + "type": "string", + "required": False, + "description": "Weather station identifier (optional)", + }, + ], +} + +# Endpoint D — get_current_weather with single required param (station) +_ENDPOINT_D: Dict[str, Any] = { + "name": "get_current_weather", + "url": "https://publicapi.envir.ee/v1/combinedWeatherData", + "params": [ + { + "name": "station", + "type": "string", + "required": True, + "description": "Weather station identifier", + }, + ], +} + + +def _make_state( + endpoint: Dict[str, Any], + collected: Dict[str, Any] | None = None, + completed: bool = False, +) -> EndpointSessionState: + return EndpointSessionState( + endpoint=endpoint, + collected_params=collected or {}, + completed=completed, + ) + + +def _make_session_store_mock() -> AsyncMock: + mock = AsyncMock() + mock.update = AsyncMock(return_value=None) + return mock + + +def _make_extractor_mock(result: ParamExtractionResult) -> MagicMock: + mock = MagicMock(return_value=result) + return mock + + +def _make_loop( + extractor_mock: MagicMock, + session_store_mock: AsyncMock | None = None, +) -> MultiEndpointAgenticLoop: + return MultiEndpointAgenticLoop( + session_store=session_store_mock or _make_session_store_mock(), + param_extractor=extractor_mock, + ) + + +def _extraction( + extracted: Dict[str, Any], + missing: List[str], + question: str, +) -> ParamExtractionResult: + return ParamExtractionResult( + extracted_params=extracted, + missing_required=missing, + clarifying_question=question, + ) + + +# --------------------------------------------------------------------------- +# Schema merging & deduplication +# --------------------------------------------------------------------------- + + +class TestSchemaMerging: + def test_shared_param_namespaced_when_both_endpoints_incomplete(self) -> None: + """'languageIsoCode' is shared by A and B (both incomplete) — it must be + namespaced to 'languageIsoCode__0' and 'languageIsoCode__1', not merged.""" + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + merged_schema, param_owners, namespace_map = loop._build_merged_schema(states) + + names = [p["name"] for p in merged_schema] + # Original name must not appear (it's been namespaced) + assert "languageIsoCode" not in names + # Both namespaced entries must appear + assert "languageIsoCode__0" in names + assert "languageIsoCode__1" in names + + def test_param_owners_for_namespaced_params(self) -> None: + """Each namespaced entry must be owned by exactly its corresponding endpoint.""" + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, param_owners, namespace_map = loop._build_merged_schema(states) + + assert param_owners["languageIsoCode__0"] == [0] + assert param_owners["languageIsoCode__1"] == [1] + + def test_namespace_map_populated_for_conflicting_params(self) -> None: + """namespace_map must map namespaced keys to (ep_idx, original_name).""" + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, _, namespace_map = loop._build_merged_schema(states) + + assert namespace_map["languageIsoCode__0"] == (0, "languageIsoCode") + assert namespace_map["languageIsoCode__1"] == (1, "languageIsoCode") + + def test_unique_params_included(self) -> None: + """'countryIsoCode' (unique to A) and 'station' (unique to B) must appear.""" + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + merged_schema, _, _ = loop._build_merged_schema(states) + + names = [p["name"] for p in merged_schema] + assert "countryIsoCode" in names + assert "station" in names + + def test_completed_endpoint_skipped_in_schema(self) -> None: + """Params from a completed endpoint should NOT appear in merged schema.""" + states = [ + _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + completed=True, + ), + _make_state(_ENDPOINT_D), + ] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + merged_schema, _, _ = loop._build_merged_schema(states) + + names = [p["name"] for p in merged_schema] + assert "countryIsoCode" not in names + assert "languageIsoCode" not in names + assert "station" in names + + def test_namespace_map_empty_when_only_one_endpoint_incomplete(self) -> None: + """When only one incomplete endpoint exists, no conflicts arise and + namespace_map must be empty.""" + states = [ + _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + completed=True, + ), + _make_state(_ENDPOINT_D), + ] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, _, namespace_map = loop._build_merged_schema(states) + + assert namespace_map == {} + + def test_required_promoted_if_any_owner_requires_it(self) -> None: + """For a shared *non-conflicting* param (one completed + one incomplete + endpoint), Pass 2 must promote it to required=True when the completed + endpoint marks it required while the incomplete endpoint marks it optional.""" + incomplete_ep = { + "name": "get_current_weather", + "url": "https://publicapi.envir.ee/v1/combinedWeatherData", + "params": [ + { + "name": "languageIsoCode", + "type": "string", + "required": False, + "description": "Language code (optional)", + }, + { + "name": "station", + "type": "string", + "required": True, + "description": "Weather station identifier", + }, + ], + } + states = [ + _make_state(incomplete_ep), + _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + completed=True, + ), + ] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + merged_schema, _, namespace_map = loop._build_merged_schema(states) + + # Only incomplete_ep is counted in Pass 0 — no conflict, namespace_map empty. + assert namespace_map == {} + lang_param = next(p for p in merged_schema if p["name"] == "languageIsoCode") + # Pass 2 promotes to required because completed _ENDPOINT_A marks it required. + assert lang_param["required"] is True + + def test_required_promoted_by_completed_endpoint_owner(self) -> None: + """Pass 2 must promote a shared param to required=True when the sole + incomplete endpoint marks it optional but a completed endpoint marks it + required — "any owner" semantics apply across both passes.""" + incomplete_ep = { + "name": "get_current_weather", + "url": "https://publicapi.envir.ee/v1/combinedWeatherData", + "params": [ + { + "name": "languageIsoCode", + "type": "string", + "required": False, + "description": "Language code (optional for weather)", + }, + { + "name": "station", + "type": "string", + "required": True, + "description": "Weather station identifier", + }, + ], + } + completed_ep = { + "name": "get_public_holidays", + "url": "https://openholidaysapi.org/PublicHolidays", + "params": [ + { + "name": "languageIsoCode", + "type": "string", + "required": True, + "description": "Response language code (required for holidays)", + }, + ], + } + states = [ + _make_state(incomplete_ep), + _make_state( + completed_ep, collected={"languageIsoCode": "ET"}, completed=True + ), + ] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + merged_schema, _, _ = loop._build_merged_schema(states) + + lang_param = next(p for p in merged_schema if p["name"] == "languageIsoCode") + assert lang_param["required"] is True + + def test_type_conflict_for_conflicting_params_each_preserves_own_type(self) -> None: + """When two incomplete endpoints share a param name with different types, + both get their own namespaced entry with their own type (no merging).""" + endpoint_date_type = { + "name": "get_something", + "url": "https://example.com/api", + "params": [ + { + "name": "languageIsoCode", + "type": "date", # different type from _ENDPOINT_A + "required": True, + "description": "Some date field", + }, + ], + } + states = [_make_state(_ENDPOINT_A), _make_state(endpoint_date_type)] + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + merged_schema, _, namespace_map = loop._build_merged_schema(states) + + # Both endpoints conflict — both get namespaced entries. + assert "languageIsoCode__0" in namespace_map + assert "languageIsoCode__1" in namespace_map + schema_by_name = {p["name"]: p for p in merged_schema} + assert ( + schema_by_name["languageIsoCode__0"]["type"] == "string" + ) # from _ENDPOINT_A + assert ( + schema_by_name["languageIsoCode__1"]["type"] == "date" + ) # from endpoint_date_type + + +# --------------------------------------------------------------------------- +# Param distribution +# --------------------------------------------------------------------------- + + +class TestParamDistribution: + def test_namespaced_param_written_to_correct_endpoint(self) -> None: + """Namespaced key 'languageIsoCode__0' must be reverse-translated and + written to endpoint 0's collected_params under the original name.""" + state_a = _make_state(_ENDPOINT_A) + state_b = _make_state(_ENDPOINT_B) + states = [state_a, state_b] + + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, param_owners, namespace_map = loop._build_merged_schema(states) + loop._distribute_params( + {"languageIsoCode__0": "ET", "languageIsoCode__1": "EN"}, + states, + param_owners, + namespace_map, + ) + + assert state_a.collected_params["languageIsoCode"] == "ET" + assert state_b.collected_params["languageIsoCode"] == "EN" + + def test_unique_param_written_only_to_owner(self) -> None: + """Extracted 'countryIsoCode' must be written only to endpoint A, not B.""" + state_a = _make_state(_ENDPOINT_A) + state_b = _make_state(_ENDPOINT_B) + states = [state_a, state_b] + + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, param_owners, namespace_map = loop._build_merged_schema(states) + loop._distribute_params( + {"countryIsoCode": "EE"}, states, param_owners, namespace_map + ) + + assert state_a.collected_params.get("countryIsoCode") == "EE" + assert "countryIsoCode" not in state_b.collected_params + + def test_endpoint_marked_completed_when_all_required_present(self) -> None: + """Endpoint A must be marked completed once all required params are set.""" + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + # A and B share languageIsoCode (conflicting); use A+D instead so there + # are no shared params and no namespace conflict. + state_a2 = _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + ) + state_d = _make_state(_ENDPOINT_D) + states2 = [state_a2, state_d] + _, param_owners2, namespace_map2 = loop._build_merged_schema(states2) + # namespace_map2 is empty (no conflicts between A and D) + assert namespace_map2 == {} + loop._distribute_params( + {"countryIsoCode": "EE"}, states2, param_owners2, namespace_map2 + ) + + assert state_a2.completed is True + assert state_d.completed is False + + def test_endpoint_with_no_required_params_immediately_completed(self) -> None: + """Endpoint C has no required params — distribute with empty dict should + mark it completed.""" + state_c = _make_state(_ENDPOINT_C) + states = [state_c] + + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, param_owners, namespace_map = loop._build_merged_schema(states) + loop._distribute_params({}, states, param_owners, namespace_map) + + assert state_c.completed is True + + def test_completed_endpoint_not_overwritten_when_value_unchanged(self) -> None: + """A completed endpoint's param must NOT be touched when the extracted + value is identical to the already-stored value.""" + state_a = _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + completed=True, + ) + state_b = _make_state(_ENDPOINT_B) + states = [state_a, state_b] + + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + # A is completed — Pass 0 only sees B's params, so no conflict. + _, param_owners, namespace_map = loop._build_merged_schema(states) + assert namespace_map == {} + original_collected = dict(state_a.collected_params) + # languageIsoCode is in param_owners (non-namespaced) — same value → no-op. + loop._distribute_params( + {"languageIsoCode": "ET"}, states, param_owners, namespace_map + ) + + assert state_a.collected_params == original_collected + assert state_a.completed is True + + def test_completed_endpoint_overwritten_when_value_differs(self) -> None: + """A completed endpoint's param IS overwritten when the extracted value + differs — the intentional shared-param correction path.""" + state_a = _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + completed=True, + ) + state_b = _make_state(_ENDPOINT_B) + states = [state_a, state_b] + + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + # A is completed — no conflict. + _, param_owners, namespace_map = loop._build_merged_schema(states) + assert namespace_map == {} + loop._distribute_params( + {"languageIsoCode": "EN"}, states, param_owners, namespace_map + ) + + assert state_a.collected_params["languageIsoCode"] == "EN" + # Still completed — the flag is not rolled back + assert state_a.completed is True + + +# --------------------------------------------------------------------------- +# Single-turn completion +# --------------------------------------------------------------------------- + + +class TestSingleTurnCompletion: + @pytest.mark.asyncio + async def test_completed_when_all_params_provided_in_first_message(self) -> None: + """If the user provides all params in one message, status should be COMPLETED. + With A+B both incomplete, languageIsoCode is conflicting — the extractor + must return namespaced keys 'languageIsoCode__0' and 'languageIsoCode__1'.""" + extractor_mock = _make_extractor_mock( + _extraction( + { + "languageIsoCode__0": "ET", + "languageIsoCode__1": "EN", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + "station": "Tallinn", + }, + [], + "none", + ) + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="ET, EE, 2026-01-01, 2026-12-31, Tallinn", + conversation_history=_HISTORY, + endpoint_states=states, + turn_count=0, + ) + + assert result.status == AgenticLoopStatus.COMPLETED + assert result.turn_count == 1 + + @pytest.mark.asyncio + async def test_endpoint_with_no_required_params_completes_on_any_turn(self) -> None: + """Endpoint C has no required params — should be marked completed immediately.""" + extractor_mock = _make_extractor_mock(_extraction({}, [], "none")) + states = [_make_state(_ENDPOINT_C)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="go", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + assert result.status == AgenticLoopStatus.COMPLETED + + +# --------------------------------------------------------------------------- +# Multi-turn accumulation +# --------------------------------------------------------------------------- + + +class TestMultiTurnAccumulation: + @pytest.mark.asyncio + async def test_params_accumulate_across_turns(self) -> None: + """Turn 1 collects namespaced 'languageIsoCode__0/1', turn 2 completes both.""" + state_a = _make_state(_ENDPOINT_A) + state_b = _make_state(_ENDPOINT_B) + states = [state_a, state_b] + store_mock = _make_session_store_mock() + + # Turn 1: collect shared languageIsoCode (namespaced because both incomplete) + extractor1 = _make_extractor_mock( + _extraction( + {"languageIsoCode__0": "ET", "languageIsoCode__1": "EN"}, + ["countryIsoCode", "validFrom", "validTo", "station"], + "Which country, date range, and weather station?", + ) + ) + loop = _make_loop(extractor1, store_mock) + result1 = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + assert result1.status == AgenticLoopStatus.NEEDS_INPUT + assert state_a.collected_params["languageIsoCode"] == "ET" + assert state_b.collected_params["languageIsoCode"] == "EN" + + # Turn 2: collect remaining params for both endpoints + extractor2 = _make_extractor_mock( + _extraction( + { + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + "station": "Tallinn", + }, + [], + "none", + ) + ) + loop2 = _make_loop(extractor2, store_mock) + result2 = await loop2.run_turn( + chat_id=_CHAT_ID, + user_message="EE, 2026-01-01, 2026-12-31, Tallinn", + conversation_history=[], + endpoint_states=states, + turn_count=1, + ) + assert result2.status == AgenticLoopStatus.COMPLETED + assert state_a.completed is True + assert state_b.completed is True + + @pytest.mark.asyncio + async def test_per_endpoint_partial_completion(self) -> None: + """One endpoint can complete before the others.""" + state_a = _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + ) + state_d = _make_state(_ENDPOINT_D) # needs 'station' + states = [state_a, state_d] + + extractor_mock = _make_extractor_mock( + _extraction({"countryIsoCode": "EE"}, ["station"], "Which weather station?") + ) + loop = _make_loop(extractor_mock) + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="EE", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + # A is now complete (has all 4 required params), D still missing station + assert state_a.completed is True + assert state_d.completed is False + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + +# --------------------------------------------------------------------------- +# Turn limit enforcement +# --------------------------------------------------------------------------- + + +class TestTurnLimitEnforcement: + @pytest.mark.asyncio + async def test_max_turns_reached_when_turn_count_equals_limit(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction({}, ["languageIsoCode"], "Which language code?") + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + # max_turns = min(3 * 2, 9) = 6 + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=6, + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + extractor_mock.assert_not_called() + + @pytest.mark.asyncio + async def test_max_turns_capped_at_multi_api_max_turns(self) -> None: + """4 endpoints → 3*4=12 → capped at MULTI_API_MAX_TURNS=9.""" + endpoint_e = { + "name": "get_public_holidays_lv", + "url": "https://openholidaysapi.org/PublicHolidays", + "params": [ + { + "name": "countryIsoCode", + "type": "string", + "required": True, + "description": "Country ISO code for Latvia", + } + ], + } + endpoint_f = { + "name": "get_current_weather_lv", + "url": "https://publicapi.envir.ee/v1/combinedWeatherData", + "params": [ + { + "name": "station", + "type": "string", + "required": True, + "description": "Latvian weather station identifier", + } + ], + } + states = [ + _make_state(_ENDPOINT_A), + _make_state(_ENDPOINT_B), + _make_state(endpoint_e), + _make_state(endpoint_f), + ] + extractor_mock = _make_extractor_mock(_extraction({}, [], "none")) + loop = _make_loop(extractor_mock) + + # turn_count=9 → max_turns=9 → should hit guard + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=9, + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + extractor_mock.assert_not_called() + + @pytest.mark.asyncio + async def test_turn_count_incremented_on_max_turns(self) -> None: + extractor_mock = _make_extractor_mock(_extraction({}, [], "none")) + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=3, # max_turns = min(3*1, 9) = 3 (holidays has 4 required params) + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + assert result.turn_count == 4 + + @pytest.mark.asyncio + async def test_session_not_saved_on_max_turns(self) -> None: + store_mock = _make_session_store_mock() + extractor_mock = _make_extractor_mock(_extraction({}, [], "none")) + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=3, + ) + + store_mock.update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_multi_intent_last_valid_turn_executes_normally(self) -> None: + """With 2 endpoints (multi-intent), max_turns = MULTI_INTENT_MAX_TURNS = 6. + + turn_count=5 → updated_turn_count=6 is the last turn that passes the + guard (5 < 6). Extraction must be attempted and the result must NOT be + MAX_TURNS_REACHED, proving that 6 full turns execute before the fallback. + """ + extractor_mock = _make_extractor_mock( + _extraction({}, ["languageIsoCode"], "Which language code?") + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="not sure", + conversation_history=[], + endpoint_states=states, + turn_count=5, # updated_turn_count=6 — last turn before guard fires + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + assert result.turn_count == 6 + extractor_mock.assert_called_once() + + @pytest.mark.asyncio + async def test_multi_intent_fixed_cap_overrides_per_endpoint_formula(self) -> None: + """The multi-intent cap is fixed at MULTI_INTENT_MAX_TURNS=6, regardless of + how many endpoints are active. + + With 3 endpoints the single-intent formula would give + min(3 * 3, MULTI_API_MAX_TURNS) = min(9, 9) = 9 turns, but the + multi-intent path uses MULTI_INTENT_MAX_TURNS=6 instead. Asserting that + turn_count=6 already triggers the guard for 3 endpoints confirms the + fixed cap is applied and the per-endpoint formula is not. + """ + endpoint_third = { + "name": "get_public_holidays_lv", + "url": "https://openholidaysapi.org/PublicHolidays", + "params": [ + { + "name": "countryIsoCode", + "type": "string", + "required": True, + "description": "Country ISO code for Latvia", + } + ], + } + states = [ + _make_state(_ENDPOINT_A), + _make_state(_ENDPOINT_B), + _make_state(endpoint_third), + ] + extractor_mock = _make_extractor_mock(_extraction({}, [], "none")) + loop = _make_loop(extractor_mock) + + # With the single-endpoint formula: min(3*3, 9)=9 → turn_count=6 would NOT + # trigger the guard. With MULTI_INTENT_MAX_TURNS=6 it must. + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=6, # == MULTI_INTENT_MAX_TURNS — must fire for 3 endpoints too + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + assert result.turn_count == 7 + extractor_mock.assert_not_called() + + +# --------------------------------------------------------------------------- +# Continuation threshold +# --------------------------------------------------------------------------- + + +class TestContinuationThreshold: + @pytest.mark.asyncio + async def test_continuation_asked_at_threshold(self) -> None: + """With 2 endpoints (multi-intent), continuation_turn = MULTI_INTENT_CONTINUATION_TURN = 4. + At turn_count=3, updated=4 → trigger.""" + extractor_mock = _make_extractor_mock( + _extraction( + {}, + ["languageIsoCode", "countryIsoCode", "station"], + "Still missing params", + ) + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="not sure", + conversation_history=[], + endpoint_states=states, + turn_count=3, # updated_turn_count = 4 == MULTI_INTENT_CONTINUATION_TURN + ) + + assert result.status == AgenticLoopStatus.AWAITING_CONTINUATION_DECISION + assert result.clarifying_question != "" + assert result.turn_count == 4 + + @pytest.mark.asyncio + async def test_continuation_not_asked_before_threshold(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction({}, ["languageIsoCode"], "Which language code?") + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hmm", + conversation_history=[], + endpoint_states=states, + turn_count=0, # updated = 1, threshold = 3 + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + @pytest.mark.asyncio + async def test_continuation_not_asked_after_threshold(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction({}, ["languageIsoCode"], "Which language code?") + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="still no", + conversation_history=[], + endpoint_states=states, + turn_count=4, # updated = 5, past MULTI_INTENT_CONTINUATION_TURN=4 + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + @pytest.mark.asyncio + async def test_single_endpoint_continuation_threshold_is_two(self) -> None: + """Single endpoint → continuation_turn = 2. turn_count=1 → updated=2 → trigger.""" + extractor_mock = _make_extractor_mock( + _extraction({}, ["station"], "Which weather station?") + ) + states = [_make_state(_ENDPOINT_D)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hmm", + conversation_history=[], + endpoint_states=states, + turn_count=1, # updated = 2 == num_endpoints+1 = 2 + ) + + assert result.status == AgenticLoopStatus.AWAITING_CONTINUATION_DECISION + + +# --------------------------------------------------------------------------- +# Continuation yes/no detection +# --------------------------------------------------------------------------- + + +class TestContinuationYesNo: + @pytest.mark.asyncio + async def test_user_yes_continues_normally(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction({}, ["languageIsoCode"], "Which language code?") + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="yes", + conversation_history=[], + endpoint_states=states, + turn_count=4, # updated=5, past MULTI_INTENT_CONTINUATION_TURN=4 + awaiting_continuation=True, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + @pytest.mark.asyncio + async def test_estonian_yes_continues(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction({}, ["languageIsoCode"], "Mis keelekood?") + ) + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="jah", + conversation_history=[], + endpoint_states=states, + turn_count=2, + awaiting_continuation=True, + ) + + assert result.status != AgenticLoopStatus.MAX_TURNS_REACHED + + @pytest.mark.asyncio + async def test_user_no_returns_max_turns_reached(self) -> None: + extractor_mock = _make_extractor_mock(_extraction({}, [], "none")) + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="no", + conversation_history=[], + endpoint_states=states, + turn_count=2, + awaiting_continuation=True, + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + extractor_mock.assert_not_called() + + @pytest.mark.asyncio + async def test_ambiguous_response_treated_as_no(self) -> None: + extractor_mock = _make_extractor_mock(_extraction({}, [], "none")) + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="maybe later", + conversation_history=[], + endpoint_states=states, + turn_count=2, + awaiting_continuation=True, + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + + +# --------------------------------------------------------------------------- +# Session persistence +# --------------------------------------------------------------------------- + + +class TestSessionPersistence: + @pytest.mark.asyncio + async def test_session_saved_on_needs_input(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction( + {"languageIsoCode": "ET"}, + ["countryIsoCode", "validFrom", "validTo"], + "Which country and date range?", + ) + ) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + store_mock.update.assert_awaited_once() + call_kwargs = store_mock.update.await_args + assert call_kwargs.args[0] == _CHAT_ID + assert call_kwargs.kwargs["turn_count"] == 1 + assert call_kwargs.kwargs["awaiting_continuation"] is False + + @pytest.mark.asyncio + async def test_session_saved_on_completed(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction( + { + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + "station": "Tallinn", + }, + [], + "none", + ) + ) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.run_turn( + chat_id=_CHAT_ID, + user_message="all params", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + store_mock.update.assert_awaited_once() + call_kwargs = store_mock.update.await_args + assert call_kwargs.kwargs["awaiting_continuation"] is False + + @pytest.mark.asyncio + async def test_session_save_failure_does_not_raise(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction({}, ["languageIsoCode"], "Which language code?") + ) + store_mock = _make_session_store_mock() + store_mock.update.side_effect = RuntimeError("Redis unavailable") + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + @pytest.mark.asyncio + async def test_awaiting_continuation_saved_as_true_at_threshold(self) -> None: + extractor_mock = _make_extractor_mock( + _extraction( + {}, ["languageIsoCode", "countryIsoCode", "station"], "Still missing" + ) + ) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.run_turn( + chat_id=_CHAT_ID, + user_message="not sure", + conversation_history=[], + endpoint_states=states, + turn_count=3, # triggers continuation threshold (updated=4 == MULTI_INTENT_CONTINUATION_TURN) + ) + + call_kwargs = store_mock.update.await_args + assert call_kwargs.kwargs["awaiting_continuation"] is True + + @pytest.mark.asyncio + async def test_session_saved_on_extractor_error(self) -> None: + """turn_count must be persisted to Redis even when param extraction raises.""" + extractor_mock = MagicMock(side_effect=RuntimeError("LLM timeout")) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + result = await loop.run_turn( + chat_id=_CHAT_ID, + user_message="Estonia", + conversation_history=[], + endpoint_states=states, + turn_count=1, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + store_mock.update.assert_awaited_once() + assert store_mock.update.await_args.kwargs["turn_count"] == 2 + + +# --------------------------------------------------------------------------- +# Merged collected params helper +# --------------------------------------------------------------------------- + + +class TestMergedCollected: + def test_merged_collected_returns_union(self) -> None: + state_a = _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + ) + state_b = _make_state(_ENDPOINT_B, collected={"station": "Tallinn"}) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + merged = loop._merged_collected([state_a, state_b]) + assert merged == { + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + "station": "Tallinn", + } + + def test_merged_collected_empty_states(self) -> None: + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + assert loop._merged_collected([]) == {} + + +# --------------------------------------------------------------------------- +# stream_run_turn — streaming path +# --------------------------------------------------------------------------- + + +def _make_stream_extractor_mock( + tokens: List[str], + result: ParamExtractionResult, +) -> MagicMock: + """Return a MagicMock whose stream_forward() coroutine returns (tokens, result).""" + mock = MagicMock() + mock.stream_forward = AsyncMock(return_value=(tokens, result)) + return mock + + +class TestStreamRunTurn: + """stream_run_turn() must mirror run_turn() semantics while returning (result, tokens).""" + + @pytest.mark.asyncio + async def test_completed_returns_empty_tokens(self) -> None: + """All params collected → COMPLETED result with no question tokens. + With A+B both incomplete, languageIsoCode is namespaced.""" + extractor_mock = _make_stream_extractor_mock( + [], + _extraction( + { + "languageIsoCode__0": "ET", + "languageIsoCode__1": "EN", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + "station": "Tallinn", + }, + [], + "none", + ), + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ET, EE, 2026-01-01, 2026-12-31, Tallinn", + conversation_history=_HISTORY, + endpoint_states=states, + turn_count=0, + ) + + assert result.status == AgenticLoopStatus.COMPLETED + assert tokens == [] + assert result.turn_count == 1 + + @pytest.mark.asyncio + async def test_needs_input_returns_question_tokens(self) -> None: + """Missing params → NEEDS_INPUT result with streamed question tokens.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " country", "?"], + _extraction( + {"languageIsoCode": "ET"}, + ["countryIsoCode", "station"], + "Which country?", + ), + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + assert tokens == ["Which", " country", "?"] + assert result.clarifying_question == "Which country?" + + @pytest.mark.asyncio + async def test_re_extracted_param_overrides_prior_value(self) -> None: + """Re-extracted value overwrites the previously collected one (correction allowed). + A+B both incomplete — languageIsoCode is namespaced.""" + extractor_mock = _make_stream_extractor_mock( + [], + _extraction( + {"languageIsoCode__0": "EN", "languageIsoCode__1": "FR"}, [], "none" + ), + ) + # Both endpoints have unique params but not the shared languageIsoCode yet + state_a = _make_state( + _ENDPOINT_A, + collected={ + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + ) + state_b = _make_state(_ENDPOINT_B, collected={"station": "Tallinn"}) + states = [state_a, state_b] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="English please", + conversation_history=[], + endpoint_states=states, + turn_count=1, + ) + + assert result.status == AgenticLoopStatus.COMPLETED + assert state_a.collected_params["languageIsoCode"] == "EN" + assert state_b.collected_params["languageIsoCode"] == "FR" + assert tokens == [] + + @pytest.mark.asyncio + async def test_max_turns_reached_returns_empty_tokens(self) -> None: + """Turn limit guard returns MAX_TURNS_REACHED with empty token list.""" + extractor_mock = _make_stream_extractor_mock( + ["Some", " question?"], + _extraction({}, ["languageIsoCode"], "Which language code?"), + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + # max_turns = min(3 * 2, 9) = 6 + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=6, + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + assert tokens == [] + extractor_mock.stream_forward.assert_not_awaited() + + @pytest.mark.asyncio + async def test_stream_extraction_exception_returns_safe_defaults(self) -> None: + """An exception from stream_forward must return NEEDS_INPUT with empty tokens.""" + extractor_mock = MagicMock() + extractor_mock.stream_forward = AsyncMock( + side_effect=RuntimeError("stream failure") + ) + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="hello", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + assert tokens == [] + assert result.collected_params == {} + + @pytest.mark.asyncio + async def test_session_language_forwarded_to_stream_forward(self) -> None: + """session_language must be passed through to stream_forward.""" + extractor_mock = _make_stream_extractor_mock( + ["Millist", " keelt?"], + _extraction({}, ["language"], "Millist keelt?"), + ) + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="Mis keel?", + conversation_history=_HISTORY, + endpoint_states=states, + turn_count=0, + session_language="et", + ) + + extractor_mock.stream_forward.assert_awaited_once() + call_kwargs = extractor_mock.stream_forward.call_args.kwargs + assert call_kwargs["session_language"] == "et" + + @pytest.mark.asyncio + async def test_user_exit_during_stream_returns_empty_tokens(self) -> None: + """awaiting_continuation=True + 'no' → MAX_TURNS_REACHED with empty tokens.""" + extractor_mock = _make_stream_extractor_mock( + ["Some", " tokens"], + _extraction({}, ["language"], "Which language?"), + ) + states = [_make_state(_ENDPOINT_A, collected={"countryIsoCode": "EE"})] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="no", + conversation_history=[], + endpoint_states=states, + turn_count=2, + awaiting_continuation=True, + ) + + assert result.status == AgenticLoopStatus.MAX_TURNS_REACHED + assert tokens == [] + assert result.collected_params == {"countryIsoCode": "EE"} + extractor_mock.stream_forward.assert_not_awaited() + + @pytest.mark.asyncio + async def test_user_yes_continues_normally(self) -> None: + """awaiting_continuation=True + 'yes' → loop continues with NEEDS_INPUT.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " language?"], + _extraction({}, ["languageIsoCode"], "Which language?"), + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="yes", + conversation_history=[], + endpoint_states=states, + turn_count=4, # updated=5, past MULTI_INTENT_CONTINUATION_TURN=4 + awaiting_continuation=True, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + assert tokens == ["Which", " language?"] + + @pytest.mark.asyncio + async def test_continuation_threshold_returns_word_tokens(self) -> None: + """At the continuation threshold, tokens are the continuation question split word-by-word.""" + from tool_classifier.constants import CONTINUATION_QUESTION + + extractor_mock = _make_stream_extractor_mock( + [], + _extraction( + {}, + ["languageIsoCode", "countryIsoCode", "station"], + "Still missing", + ), + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="not sure", + conversation_history=[], + endpoint_states=states, + turn_count=3, # updated_turn_count = 4 == MULTI_INTENT_CONTINUATION_TURN + ) + + assert result.status == AgenticLoopStatus.AWAITING_CONTINUATION_DECISION + assert "".join(tokens) == CONTINUATION_QUESTION + assert all(t.endswith(" ") for t in tokens[:-1]) + assert not tokens[-1].endswith(" ") + + @pytest.mark.asyncio + async def test_continuation_threshold_uses_continuation_language(self) -> None: + """continuation_language overrides session_language for the threshold tokens.""" + from tool_classifier.constants import CONTINUATION_QUESTION_ET + + extractor_mock = _make_stream_extractor_mock( + [], + _extraction( + {}, + ["languageIsoCode", "countryIsoCode", "station"], + "Veel puudub", + ), + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ei tea", + conversation_history=[], + endpoint_states=states, + turn_count=3, # updated_turn_count = 4 == MULTI_INTENT_CONTINUATION_TURN + session_language="et", + continuation_language="et", + ) + + assert result.status == AgenticLoopStatus.AWAITING_CONTINUATION_DECISION + assert "".join(tokens) == CONTINUATION_QUESTION_ET + + @pytest.mark.asyncio + async def test_shared_params_distributed_to_both_endpoints(self) -> None: + """Namespaced params are written to their corresponding endpoint states.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " country?"], + _extraction( + {"languageIsoCode__0": "ET", "languageIsoCode__1": "EN"}, + ["countryIsoCode", "station"], + "Which country?", + ), + ) + state_a = _make_state(_ENDPOINT_A) + state_b = _make_state(_ENDPOINT_B) + states = [state_a, state_b] + loop = _make_loop(extractor_mock) + + result, _ = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + assert state_a.collected_params["languageIsoCode"] == "ET" + assert state_b.collected_params["languageIsoCode"] == "EN" + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + @pytest.mark.asyncio + async def test_session_saved_on_completed(self) -> None: + """Session is persisted with awaiting_continuation=False on COMPLETED.""" + extractor_mock = _make_stream_extractor_mock( + [], + _extraction( + { + "languageIsoCode__0": "ET", + "languageIsoCode__1": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + "station": "Tallinn", + }, + [], + "none", + ), + ) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="all params", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + store_mock.update.assert_awaited_once() + assert store_mock.update.await_args.kwargs["awaiting_continuation"] is False + + @pytest.mark.asyncio + async def test_session_saved_on_needs_input(self) -> None: + """Session is persisted with awaiting_continuation=False on NEEDS_INPUT.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " country?"], + _extraction( + {"languageIsoCode": "ET"}, + ["countryIsoCode", "validFrom", "validTo"], + "Which country?", + ), + ) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + store_mock.update.assert_awaited_once() + call_kwargs = store_mock.update.await_args.kwargs + assert call_kwargs["awaiting_continuation"] is False + assert call_kwargs["turn_count"] == 1 + + @pytest.mark.asyncio + async def test_session_saved_with_awaiting_continuation_true_at_threshold( + self, + ) -> None: + """Session is persisted with awaiting_continuation=True at the threshold.""" + extractor_mock = _make_stream_extractor_mock( + [], + _extraction( + {}, + ["languageIsoCode", "countryIsoCode", "station"], + "Still missing", + ), + ) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="not sure", + conversation_history=[], + endpoint_states=states, + turn_count=3, # triggers continuation threshold (updated=4 == MULTI_INTENT_CONTINUATION_TURN) + ) + + call_kwargs = store_mock.update.await_args.kwargs + assert call_kwargs["awaiting_continuation"] is True + assert call_kwargs["turn_count"] == 4 + + @pytest.mark.asyncio + async def test_session_not_saved_on_max_turns(self) -> None: + """No session save when the turn limit is hit.""" + extractor_mock = _make_stream_extractor_mock([], _extraction({}, [], "none")) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=3, # max_turns = min(3 * 1, 9) = 3 + ) + + store_mock.update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_session_not_saved_on_user_exit(self) -> None: + """No session save when the user chooses to exit (awaiting_continuation + 'no').""" + extractor_mock = _make_stream_extractor_mock([], _extraction({}, [], "none")) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="no", + conversation_history=[], + endpoint_states=states, + turn_count=2, + awaiting_continuation=True, + ) + + store_mock.update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_session_save_failure_does_not_raise(self) -> None: + """A Redis failure must not propagate — NEEDS_INPUT is still returned.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " language?"], + _extraction({}, ["languageIsoCode"], "Which language?"), + ) + store_mock = _make_session_store_mock() + store_mock.update.side_effect = RuntimeError("Redis unavailable") + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + result, _ = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + + @pytest.mark.asyncio + async def test_session_saved_on_stream_extractor_error(self) -> None: + """turn_count must be persisted to Redis even when stream_forward raises.""" + extractor_mock = MagicMock() + extractor_mock.stream_forward = AsyncMock( + side_effect=RuntimeError("stream failure") + ) + store_mock = _make_session_store_mock() + states = [_make_state(_ENDPOINT_A)] + loop = _make_loop(extractor_mock, store_mock) + + result, tokens = await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="hi", + conversation_history=[], + endpoint_states=states, + turn_count=1, + ) + + assert result.status == AgenticLoopStatus.NEEDS_INPUT + assert tokens == [] + store_mock.update.assert_awaited_once() + assert store_mock.update.await_args.kwargs["turn_count"] == 2 + + +# --------------------------------------------------------------------------- +# _build_intent_groups — unit tests +# --------------------------------------------------------------------------- + + +class TestBuildIntentGroups: + """Unit tests for MultiEndpointAgenticLoop._build_intent_groups().""" + + def test_two_intents_with_missing_required_params_returns_two_groups( + self, + ) -> None: + """Two incomplete endpoints each with missing required params → two groups.""" + state_a = _make_state(_ENDPOINT_A) + state_b = _make_state(_ENDPOINT_B) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups( + [state_a, state_b], already_collected={}, namespace_map={} + ) + + assert len(groups) == 2 + intents = {g["intent"] for g in groups} + assert "get_public_holidays" in intents + assert "get_current_weather" in intents + + def test_single_incomplete_endpoint_returns_single_element_list(self) -> None: + """Only one incomplete endpoint has missing required params → single-element list.""" + state_a = _make_state(_ENDPOINT_A) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups( + [state_a], already_collected={}, namespace_map={} + ) + + assert len(groups) == 1 + + def test_no_endpoints_with_missing_returns_empty_list(self) -> None: + """No incomplete endpoints → empty list.""" + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups([], already_collected={}, namespace_map={}) + + assert groups == [] + + def test_completed_endpoints_are_skipped(self) -> None: + """A completed endpoint must not contribute a group.""" + state_a = _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + completed=True, + ) + state_b = _make_state(_ENDPOINT_B) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + # Only one incomplete endpoint → single-element list for the incomplete endpoint + groups = loop._build_intent_groups( + [state_a, state_b], + already_collected={"languageIsoCode": "ET"}, + namespace_map={}, + ) + + assert len(groups) == 1 + assert groups[0]["intent"] == "get_current_weather" + + def test_already_collected_params_excluded_from_descriptions(self) -> None: + """Params already in already_collected must not appear in missing_param_descriptions.""" + state_a = _make_state(_ENDPOINT_A) + state_b = _make_state(_ENDPOINT_B) + already_collected = {"languageIsoCode": "ET"} + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups( + [state_a, state_b], + already_collected=already_collected, + namespace_map={}, + ) + + # Groups still returned because both endpoints still have missing params + assert len(groups) == 2 + for group in groups: + # languageIsoCode is collected — its description must not appear + assert not any( + "language" in desc.lower() + for desc in group["missing_param_descriptions"] + ) + + def test_optional_params_excluded_from_descriptions(self) -> None: + """Optional (required=False) params must not appear in missing_param_descriptions.""" + endpoint_with_optional = { + "name": "get_data_a", + "url": "https://api.example.com/a", + "params": [ + { + "name": "requiredParam", + "type": "string", + "required": True, + "description": "The required value", + }, + { + "name": "optionalParam", + "type": "string", + "required": False, + "description": "An optional filter", + }, + ], + } + endpoint_b_required = { + "name": "get_data_b", + "url": "https://api.example.com/b", + "params": [ + { + "name": "anotherRequired", + "type": "string", + "required": True, + "description": "Another required field", + }, + ], + } + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups( + [_make_state(endpoint_with_optional), _make_state(endpoint_b_required)], + already_collected={}, + namespace_map={}, + ) + + assert len(groups) == 2 + group_a = next(g for g in groups if g["intent"] == "get_data_a") + descriptions = group_a["missing_param_descriptions"] + assert "The required value" in descriptions + assert not any("optional" in d.lower() for d in descriptions) + + def test_format_hints_stripped_from_descriptions(self) -> None: + """(YYYY-MM-DD) format hints must be stripped from missing_param_descriptions.""" + state_a = _make_state(_ENDPOINT_A) # validFrom/validTo have (YYYY-MM-DD) hints + state_b = _make_state(_ENDPOINT_B) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups( + [state_a, state_b], + already_collected={"languageIsoCode": "ET"}, + namespace_map={}, + ) + + group_a = next(g for g in groups if g["intent"] == "get_public_holidays") + for desc in group_a["missing_param_descriptions"]: + assert "YYYY" not in desc + assert "MM-DD" not in desc + + def test_intent_name_from_endpoint_name_key(self) -> None: + """Group 'intent' field must be taken from the endpoint's 'name' key.""" + state_a = _make_state(_ENDPOINT_A) + state_b = _make_state(_ENDPOINT_B) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups( + [state_a, state_b], already_collected={}, namespace_map={} + ) + + intents = {g["intent"] for g in groups} + assert "get_public_holidays" in intents + assert "get_current_weather" in intents + + def test_intent_name_falls_back_to_description_when_name_absent(self) -> None: + """When endpoint has no 'name' key, 'intent' must fall back to 'description'.""" + endpoint_no_name = { + "description": "lookup_initiative", + "url": "https://api.example.com/initiative", + "params": [ + { + "name": "initiativeId", + "type": "string", + "required": True, + "description": "Unique initiative identifier", + }, + ], + } + endpoint_b = { + "name": "search_address", + "url": "https://api.example.com/address", + "params": [ + { + "name": "address", + "type": "string", + "required": True, + "description": "Street address or place name", + }, + ], + } + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + groups = loop._build_intent_groups( + [_make_state(endpoint_no_name), _make_state(endpoint_b)], + already_collected={}, + namespace_map={}, + ) + + assert len(groups) == 2 + intents = {g["intent"] for g in groups} + assert "lookup_initiative" in intents + assert "search_address" in intents + + def test_endpoint_with_only_optional_params_excluded_from_groups(self) -> None: + """An endpoint with no missing required params contributes no group; + the remaining endpoint's required params produce a single-element list.""" + state_c = _make_state(_ENDPOINT_C) # only optional params + state_b = _make_state(_ENDPOINT_B) # has required 'station' + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + # C has no required params → no group for C → only 1 group (for B) + groups = loop._build_intent_groups( + [state_c, state_b], already_collected={}, namespace_map={} + ) + + assert len(groups) == 1 + + +# --------------------------------------------------------------------------- +# intent_groups forwarding in stream_run_turn +# --------------------------------------------------------------------------- + + +class TestStreamRunTurnIntentGroupsForwarding: + """Verify that intent_groups is built correctly and forwarded to stream_forward.""" + + @pytest.mark.asyncio + async def test_intent_groups_forwarded_for_two_incomplete_endpoints(self) -> None: + """stream_forward must receive a non-empty intent_groups list when two + incomplete endpoints each have missing required params.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " country?"], + _extraction( + {}, + ["languageIsoCode", "countryIsoCode", "station"], + "Which country and station?", + ), + ) + states = [_make_state(_ENDPOINT_A), _make_state(_ENDPOINT_B)] + loop = _make_loop(extractor_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="not sure", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + extractor_mock.stream_forward.assert_awaited_once() + call_kwargs = extractor_mock.stream_forward.call_args.kwargs + intent_groups = call_kwargs["intent_groups"] + assert isinstance(intent_groups, list) + assert len(intent_groups) == 2 + intents = {g["intent"] for g in intent_groups} + assert "get_public_holidays" in intents + assert "get_current_weather" in intents + for group in intent_groups: + assert "missing_param_descriptions" in group + assert isinstance(group["missing_param_descriptions"], list) + assert len(group["missing_param_descriptions"]) > 0 + + @pytest.mark.asyncio + async def test_intent_groups_single_element_for_single_endpoint(self) -> None: + """stream_forward must receive a single-element intent_groups list when only one endpoint is active.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " station?"], + _extraction({}, ["station"], "Which station?"), + ) + states = [_make_state(_ENDPOINT_D)] + loop = _make_loop(extractor_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="not sure", + conversation_history=[], + endpoint_states=states, + turn_count=0, + ) + + extractor_mock.stream_forward.assert_awaited_once() + call_kwargs = extractor_mock.stream_forward.call_args.kwargs + assert len(call_kwargs["intent_groups"]) == 1 + assert call_kwargs["intent_groups"][0]["intent"] == "get_current_weather" + + @pytest.mark.asyncio + async def test_intent_groups_single_element_when_other_endpoint_complete( + self, + ) -> None: + """stream_forward must receive a single-element intent_groups list when only one of + two endpoints has missing required params (the other has all params collected).""" + state_a = _make_state( + _ENDPOINT_A, + collected={ + "languageIsoCode": "ET", + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, + completed=True, + ) + state_b = _make_state(_ENDPOINT_B) # still missing 'station' + extractor_mock = _make_stream_extractor_mock( + ["Which", " station?"], + _extraction({"languageIsoCode": "ET"}, ["station"], "Which station?"), + ) + loop = _make_loop(extractor_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=[state_a, state_b], + turn_count=0, + ) + + extractor_mock.stream_forward.assert_awaited_once() + call_kwargs = extractor_mock.stream_forward.call_args.kwargs + assert len(call_kwargs["intent_groups"]) == 1 + assert call_kwargs["intent_groups"][0]["intent"] == "get_current_weather" + + @pytest.mark.asyncio + async def test_intent_groups_excludes_already_collected_params(self) -> None: + """Already-collected params must not appear in intent_groups descriptions + forwarded to stream_forward.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " country?"], + _extraction( + {}, + ["countryIsoCode", "validFrom", "validTo", "station"], + "Which country and station?", + ), + ) + state_a = _make_state(_ENDPOINT_A, collected={"languageIsoCode": "ET"}) + state_b = _make_state(_ENDPOINT_B, collected={"languageIsoCode": "ET"}) + loop = _make_loop(extractor_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=[state_a, state_b], + turn_count=1, + ) + + extractor_mock.stream_forward.assert_awaited_once() + intent_groups = extractor_mock.stream_forward.call_args.kwargs["intent_groups"] + assert len(intent_groups) == 2 + # languageIsoCode is collected — its description must not appear + for group in intent_groups: + assert not any( + "language" in desc.lower() + for desc in group["missing_param_descriptions"] + ) + + @pytest.mark.asyncio + async def test_intent_groups_format_hints_stripped_before_forwarding(self) -> None: + """Format hints in descriptions must be stripped before being passed to stream_forward.""" + extractor_mock = _make_stream_extractor_mock( + ["Which", " dates?"], + _extraction( + {"languageIsoCode": "ET"}, + ["countryIsoCode", "validFrom", "validTo", "station"], + "Which dates and station?", + ), + ) + state_a = _make_state(_ENDPOINT_A, collected={"languageIsoCode": "ET"}) + state_b = _make_state(_ENDPOINT_B, collected={"languageIsoCode": "ET"}) + loop = _make_loop(extractor_mock) + + await loop.stream_run_turn( + chat_id=_CHAT_ID, + user_message="ET", + conversation_history=[], + endpoint_states=[state_a, state_b], + turn_count=1, + ) + + intent_groups = extractor_mock.stream_forward.call_args.kwargs["intent_groups"] + group_a = next(g for g in intent_groups if g["intent"] == "get_public_holidays") + for desc in group_a["missing_param_descriptions"]: + assert "YYYY" not in desc + assert "MM-DD" not in desc + + +# --------------------------------------------------------------------------- +# Conflict namespace fixtures — both endpoints define startDate/endDate +# --------------------------------------------------------------------------- + +_ENDPOINT_E: Dict[str, Any] = { + "name": "get_first_intent_data", + "url": "https://api.example.com/first", + "params": [ + { + "name": "startDate", + "type": "date", + "required": True, + "description": "Start date for first intent (YYYY-MM-DD)", + }, + { + "name": "endDate", + "type": "date", + "required": True, + "description": "End date for first intent (YYYY-MM-DD)", + }, + ], +} + +_ENDPOINT_F: Dict[str, Any] = { + "name": "get_second_intent_data", + "url": "https://api.example.com/second", + "params": [ + { + "name": "startDate", + "type": "date", + "required": True, + "description": "Start date for second intent (YYYY-MM-DD)", + }, + { + "name": "endDate", + "type": "date", + "required": True, + "description": "End date for second intent (YYYY-MM-DD)", + }, + ], +} + + +# --------------------------------------------------------------------------- +# TestSchemaNamespacing — conflicting param namespacing +# --------------------------------------------------------------------------- + + +class TestSchemaNamespacing: + """Unit tests verifying that conflicting params (same name in 2+ incomplete + endpoints) are namespaced as {name}__{ep_idx} in the merged schema.""" + + def test_build_merged_schema_produces_namespaced_keys_for_conflicting_params( + self, + ) -> None: + """startDate/endDate appear in both E and F ? namespaced keys in schema.""" + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + merged_schema, param_owners, namespace_map = loop._build_merged_schema( + [_make_state(_ENDPOINT_E), _make_state(_ENDPOINT_F)] + ) + + names = [p["name"] for p in merged_schema] + assert "startDate__0" in names + assert "startDate__1" in names + assert "endDate__0" in names + assert "endDate__1" in names + # Original names must NOT appear + assert "startDate" not in names + assert "endDate" not in names + + def test_namespace_map_populated_for_startdate_enddate_conflict(self) -> None: + """namespace_map maps each namespaced key to (ep_idx, original_name).""" + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, _, namespace_map = loop._build_merged_schema( + [_make_state(_ENDPOINT_E), _make_state(_ENDPOINT_F)] + ) + + assert namespace_map["startDate__0"] == (0, "startDate") + assert namespace_map["startDate__1"] == (1, "startDate") + assert namespace_map["endDate__0"] == (0, "endDate") + assert namespace_map["endDate__1"] == (1, "endDate") + + def test_param_owners_maps_each_namespaced_key_to_single_endpoint(self) -> None: + """param_owners for namespaced keys lists only one endpoint each.""" + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + _, param_owners, _ = loop._build_merged_schema( + [_make_state(_ENDPOINT_E), _make_state(_ENDPOINT_F)] + ) + + assert param_owners["startDate__0"] == [0] + assert param_owners["startDate__1"] == [1] + assert param_owners["endDate__0"] == [0] + assert param_owners["endDate__1"] == [1] + + def test_distribute_params_writes_original_names_to_each_endpoint(self) -> None: + """Namespaced keys in extracted_params are written under original names + to the correct endpoint's collected_params.""" + state_e = _make_state(_ENDPOINT_E) + state_f = _make_state(_ENDPOINT_F) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + _, param_owners, namespace_map = loop._build_merged_schema([state_e, state_f]) + loop._distribute_params( + { + "startDate__0": "2026-01-01", + "endDate__0": "2026-06-30", + "startDate__1": "2026-07-01", + "endDate__1": "2026-12-31", + }, + [state_e, state_f], + param_owners, + namespace_map, + ) + + assert state_e.collected_params["startDate"] == "2026-01-01" + assert state_e.collected_params["endDate"] == "2026-06-30" + assert state_f.collected_params["startDate"] == "2026-07-01" + assert state_f.collected_params["endDate"] == "2026-12-31" + + def test_build_intent_groups_shows_date_params_in_both_groups(self) -> None: + """With namespaced dates, both groups should list their own date descriptions.""" + state_e = _make_state(_ENDPOINT_E) + state_f = _make_state(_ENDPOINT_F) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + + _, _, namespace_map = loop._build_merged_schema([state_e, state_f]) + groups = loop._build_intent_groups( + [state_e, state_f], already_collected={}, namespace_map=namespace_map + ) + + assert len(groups) == 2 + for group in groups: + assert len(group["missing_param_descriptions"]) > 0 + + +# --------------------------------------------------------------------------- +# TestBuildNamespacedAlreadyCollected +# --------------------------------------------------------------------------- + + +class TestBuildNamespacedAlreadyCollected: + """Unit tests for _build_namespaced_already_collected().""" + + def test_fast_path_when_no_namespace_map(self) -> None: + """When namespace_map is empty, returns simple union of all collected_params.""" + state_a = _make_state(_ENDPOINT_A, collected={"languageIsoCode": "ET"}) + state_b = _make_state(_ENDPOINT_B, collected={"station": "Tallinn"}) + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + result = loop._build_namespaced_already_collected( + [state_a, state_b], namespace_map={} + ) + assert result["languageIsoCode"] == "ET" + assert result["station"] == "Tallinn" + + def test_conflicting_params_appear_under_namespaced_keys(self) -> None: + """When namespace_map is set, each endpoint's conflicting params are + returned under their namespaced keys.""" + state_e = _make_state( + _ENDPOINT_E, collected={"startDate": "2026-01-01", "endDate": "2026-06-30"} + ) + state_f = _make_state(_ENDPOINT_F, collected={"startDate": "2026-07-01"}) + + namespace_map = { + "startDate__0": (0, "startDate"), + "startDate__1": (1, "startDate"), + "endDate__0": (0, "endDate"), + "endDate__1": (1, "endDate"), + } + loop = _make_loop(_make_extractor_mock(_extraction({}, [], "none"))) + result = loop._build_namespaced_already_collected( + [state_e, state_f], namespace_map=namespace_map + ) + + assert result["startDate__0"] == "2026-01-01" + assert result["endDate__0"] == "2026-06-30" + assert result["startDate__1"] == "2026-07-01" + # endDate__1 not yet collected — absent + assert "endDate__1" not in result + # Original conflicting names must not appear at top level + assert "startDate" not in result + assert "endDate" not in result diff --git a/tests/test_multi_api_caller.py b/tests/test_multi_api_caller.py new file mode 100644 index 00000000..1a115100 --- /dev/null +++ b/tests/test_multi_api_caller.py @@ -0,0 +1,451 @@ +"""Unit tests for MultiAPICaller — parallel batch API execution.""" + +import asyncio + +import pytest + +from tool_classifier.api_caller import APICaller +from tool_classifier.constants import ( + CIRCUIT_BREAKER_OPEN_MESSAGES, + MULTI_API_BATCH_TIMEOUT, + MULTI_API_PARTIAL_FAILURE_MESSAGES, + SERVICE_TIMEOUT_MESSAGES, + SERVICE_UNAVAILABLE_MESSAGES, +) +from tool_classifier.models import APICallResult +from tool_classifier.multi_api_caller import MultiAPICaller + + +# --------------------------------------------------------------------------- +# Fixtures / helpers +# --------------------------------------------------------------------------- + +_WEATHER_URL = "https://publicapi.envir.ee/v1/combinedWeatherData" +_HOLIDAYS_URL = "https://openholidaysapi.org/PublicHolidays" + +_EP_A = {"url": _WEATHER_URL, "method": "GET", "call_params": {"station": "Tallinn"}} +_EP_B = { + "url": _HOLIDAYS_URL, + "method": "GET", + "call_params": { + "countryIsoCode": "EE", + "validFrom": "2026-01-01", + "validTo": "2026-12-31", + }, +} +_EP_C = {"url": _WEATHER_URL, "method": "GET", "call_params": {}} + +_WEATHER_PAYLOAD = { + "observations": { + "station": [ + { + "name": "Tallinn", + "phenomenon": "Cloudy", + "airtemperature": 12.5, + "windspeed": 4.2, + "relativehumidity": 78, + } + ] + } +} +_HOLIDAYS_PAYLOAD = [ + { + "startDate": "2026-01-01", + "endDate": "2026-01-01", + "type": "Public", + "name": [{"language": "ET", "text": "Uusaasta"}], + "nationwide": True, + }, + { + "startDate": "2026-02-24", + "endDate": "2026-02-24", + "type": "Public", + "name": [{"language": "ET", "text": "Eesti Vabariigi aastapäev"}], + "nationwide": True, + }, +] + + +def _ok(url: str = _WEATHER_URL) -> APICallResult: + payload: object = _HOLIDAYS_PAYLOAD if url == _HOLIDAYS_URL else _WEATHER_PAYLOAD + return APICallResult( + success=True, status_code=200, response_data=payload, error=None + ) + + +def _fail(status_code: int = 500, language: str = "en") -> APICallResult: + return APICallResult( + success=False, + status_code=status_code, + response_data="", + error=SERVICE_UNAVAILABLE_MESSAGES[language], + ) + + +def _timeout_result(language: str = "en") -> APICallResult: + return APICallResult( + success=False, + status_code=0, + response_data="", + error=SERVICE_TIMEOUT_MESSAGES[language], + ) + + +def _cb_open_result(language: str = "en") -> APICallResult: + return APICallResult( + success=False, + status_code=0, + response_data="", + error=CIRCUIT_BREAKER_OPEN_MESSAGES[language], + ) + + +def _make_caller_with_results(results_by_url: dict[str, APICallResult]) -> APICaller: + """Return an APICaller whose .call() returns the mapped result per URL.""" + caller = APICaller() + + async def _call( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: # noqa: ANN001 + return results_by_url[url] + + caller.call = _call # type: ignore[method-assign] + return caller + + +# --------------------------------------------------------------------------- +# All calls succeed +# --------------------------------------------------------------------------- + + +class TestAllSucceed: + @pytest.mark.asyncio + async def test_all_succeeded_true(self) -> None: + caller = _make_caller_with_results( + {_EP_A["url"]: _ok(_EP_A["url"]), _EP_B["url"]: _ok(_EP_B["url"])} + ) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert result.all_succeeded is True + assert len(result.results) == 2 + assert len(result.endpoints) == 2 + assert result.results[0].success is True + assert result.results[1].success is True + + @pytest.mark.asyncio + async def test_order_preserved(self) -> None: + ok_a = APICallResult( + success=True, status_code=200, response_data=_WEATHER_PAYLOAD, error=None + ) + ok_b = APICallResult( + success=True, status_code=200, response_data=_HOLIDAYS_PAYLOAD, error=None + ) + caller = _make_caller_with_results({_EP_A["url"]: ok_a, _EP_B["url"]: ok_b}) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert result.results[0].response_data == _WEATHER_PAYLOAD + assert result.results[1].response_data == _HOLIDAYS_PAYLOAD + + @pytest.mark.asyncio + async def test_successful_results_property(self) -> None: + caller = _make_caller_with_results({_EP_A["url"]: _ok(), _EP_B["url"]: _ok()}) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert len(result.successful_results) == 2 + assert len(result.failed_results) == 0 + + +# --------------------------------------------------------------------------- +# Partial failure (one call fails) +# --------------------------------------------------------------------------- + + +class TestPartialFailure: + @pytest.mark.asyncio + async def test_one_fails_all_succeeded_false(self) -> None: + caller = _make_caller_with_results({_EP_A["url"]: _ok(), _EP_B["url"]: _fail()}) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert result.all_succeeded is False + assert result.results[0].success is True + assert result.results[1].success is False + + @pytest.mark.asyncio + async def test_successful_and_failed_result_helpers(self) -> None: + caller = _make_caller_with_results({_EP_A["url"]: _ok(), _EP_B["url"]: _fail()}) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert len(result.successful_results) == 1 + assert len(result.failed_results) == 1 + ep_ok, res_ok = result.successful_results[0] + assert ep_ok == _EP_A + assert res_ok.success is True + ep_fail, res_fail = result.failed_results[0] + assert ep_fail == _EP_B + assert res_fail.success is False + + @pytest.mark.asyncio + async def test_5xx_error_kept_in_results(self) -> None: + caller = _make_caller_with_results( + { + _EP_A["url"]: _ok(), + _EP_B["url"]: _fail(500), + _EP_C["url"]: _ok(), + } + ) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B, _EP_C], language="en") + + assert len(result.results) == 3 + assert result.results[0].success is True + assert result.results[1].success is False + assert result.results[1].status_code == 500 + assert result.results[2].success is True + + +# --------------------------------------------------------------------------- +# Individual call timeout +# --------------------------------------------------------------------------- + + +class TestIndividualTimeout: + @pytest.mark.asyncio + async def test_individual_timeout_produces_failure_result(self) -> None: + caller = _make_caller_with_results( + {_EP_A["url"]: _ok(), _EP_B["url"]: _timeout_result()} + ) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert result.all_succeeded is False + assert result.results[1].success is False + assert result.results[1].status_code == 0 + + +# --------------------------------------------------------------------------- +# Circuit breaker open for one URL +# --------------------------------------------------------------------------- + + +class TestCircuitBreakerOpen: + @pytest.mark.asyncio + async def test_cb_open_url_gets_failure_others_succeed(self) -> None: + caller = _make_caller_with_results( + {_EP_A["url"]: _ok(), _EP_B["url"]: _cb_open_result()} + ) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert result.all_succeeded is False + assert result.results[0].success is True + assert result.results[1].success is False + assert result.results[1].status_code == 0 + + @pytest.mark.asyncio + async def test_cb_open_message_correct_language_et(self) -> None: + caller = _make_caller_with_results({_EP_A["url"]: _cb_open_result("et")}) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A], language="et") + + assert result.results[0].error == CIRCUIT_BREAKER_OPEN_MESSAGES["et"] + + +# --------------------------------------------------------------------------- +# Batch-level timeout +# --------------------------------------------------------------------------- + + +class TestBatchTimeout: + @pytest.mark.asyncio + async def test_batch_timeout_cancels_pending_returns_partial(self) -> None: + """Batch timeout fires; completed tasks are preserved, pending get failure result.""" + + async def _slow_call( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: + await asyncio.sleep(10) # never finishes in the test + return _ok(url) + + async def _fast_call( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: + return _ok(url) + + caller = APICaller() + + async def _dispatched( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: + if url == _EP_A["url"]: + return await _fast_call(url, method, params, language) + return await _slow_call(url, method, params, language) + + caller.call = _dispatched # type: ignore[method-assign] + + multi = MultiAPICaller(caller, batch_timeout=1) # 1-second batch timeout + result = await multi.call_all([_EP_A, _EP_B], language="en") + + assert len(result.results) == 2 + # EP_A completed fast; EP_B was cancelled → failure + assert result.results[0].success is True + assert result.results[1].success is False + assert result.results[1].status_code == 0 + assert result.results[1].error == MULTI_API_PARTIAL_FAILURE_MESSAGES["en"] + + @pytest.mark.asyncio + async def test_batch_timeout_all_cancelled(self) -> None: + """When both tasks are slow, both get failure results.""" + + async def _slow( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: + await asyncio.sleep(10) + return _ok(url) + + caller = APICaller() + caller.call = _slow # type: ignore[method-assign] + + multi = MultiAPICaller(caller, batch_timeout=1) + result = await multi.call_all([_EP_A, _EP_B], language="ru") + + assert len(result.results) == 2 + assert not any(r.success for r in result.results) + for r in result.results: + assert r.error == MULTI_API_PARTIAL_FAILURE_MESSAGES["ru"] + + +# --------------------------------------------------------------------------- +# Edge cases +# --------------------------------------------------------------------------- + + +class TestEdgeCases: + @pytest.mark.asyncio + async def test_empty_endpoints_returns_empty_result(self) -> None: + caller = APICaller() + multi = MultiAPICaller(caller) + result = await multi.call_all([], language="en") + + assert result.results == [] + assert result.endpoints == [] + assert result.all_succeeded is True + assert result.successful_results == [] + assert result.failed_results == [] + + @pytest.mark.asyncio + async def test_single_endpoint_success(self) -> None: + caller = _make_caller_with_results({_EP_A["url"]: _ok()}) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A], language="en") + + assert len(result.results) == 1 + assert result.all_succeeded is True + + @pytest.mark.asyncio + async def test_single_endpoint_failure(self) -> None: + caller = _make_caller_with_results({_EP_A["url"]: _fail()}) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A], language="en") + + assert len(result.results) == 1 + assert result.all_succeeded is False + + @pytest.mark.asyncio + async def test_exception_in_gather_converted_to_failure(self) -> None: + """A BaseException raised inside a task is caught and converted to APICallResult(success=False).""" + + async def _raises( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: + raise RuntimeError("unexpected internal error") + + caller = APICaller() + caller.call = _raises # type: ignore[method-assign] + + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A], language="en") + + assert len(result.results) == 1 + assert result.results[0].success is False + assert result.results[0].status_code == 0 + assert result.results[0].error == MULTI_API_PARTIAL_FAILURE_MESSAGES["en"] + + +# --------------------------------------------------------------------------- +# Per-URL circuit breaker isolation +# --------------------------------------------------------------------------- + + +class TestCircuitBreakerIsolation: + @pytest.mark.asyncio + async def test_cb_open_for_one_url_does_not_affect_others(self) -> None: + """Three endpoints; CB open only for B — A and C still execute normally.""" + caller = _make_caller_with_results( + { + _EP_A["url"]: _ok(), + _EP_B["url"]: _cb_open_result(), + _EP_C["url"]: _ok(), + } + ) + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A, _EP_B, _EP_C], language="en") + + assert result.results[0].success is True + assert result.results[1].success is False + assert result.results[2].success is True + assert len(result.failed_results) == 1 + assert len(result.successful_results) == 2 + + +# --------------------------------------------------------------------------- +# Multilingual error messages +# --------------------------------------------------------------------------- + + +class TestMultilingualMessages: + @pytest.mark.asyncio + @pytest.mark.parametrize("language", ["et", "ru", "en"]) + async def test_partial_failure_message_correct_language( + self, language: str + ) -> None: + """Batch timeout fills pending slots with the right language's message.""" + + async def _slow( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: + await asyncio.sleep(10) + return _ok(url) + + caller = APICaller() + caller.call = _slow # type: ignore[method-assign] + + multi = MultiAPICaller(caller, batch_timeout=1) + result = await multi.call_all([_EP_A], language=language) + + assert result.results[0].error == MULTI_API_PARTIAL_FAILURE_MESSAGES[language] + + @pytest.mark.asyncio + @pytest.mark.parametrize("language", ["et", "ru", "en"]) + async def test_exception_message_correct_language(self, language: str) -> None: + async def _raises( + url: str, method: str, params: dict, language: str = "et" + ) -> APICallResult: + raise ValueError("boom") + + caller = APICaller() + caller.call = _raises # type: ignore[method-assign] + + multi = MultiAPICaller(caller) + result = await multi.call_all([_EP_A], language=language) + + assert result.results[0].error == MULTI_API_PARTIAL_FAILURE_MESSAGES[language] + + @pytest.mark.asyncio + async def test_default_batch_timeout_constant_used(self) -> None: + multi = MultiAPICaller(APICaller()) + assert multi._batch_timeout == MULTI_API_BATCH_TIMEOUT diff --git a/tests/test_multi_response_formatter.py b/tests/test_multi_response_formatter.py new file mode 100644 index 00000000..92562168 --- /dev/null +++ b/tests/test_multi_response_formatter.py @@ -0,0 +1,947 @@ +"""Unit tests for MultiResponseFormatterModule — multi-API result synthesiser.""" + +from collections.abc import AsyncGenerator +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import dspy +import dspy.streaming +import pytest + +from tool_classifier.multi_response_formatter import ( + MultiResponseFormatterModule, + _MULTI_FORMATTER_ERROR_MESSAGES, + _MAX_TOTAL_RESPONSE_BYTES, +) + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture(autouse=True) +def mock_dspy_lm() -> Any: + """Mock DSPy LM to prevent 'No LM is loaded' errors during tests.""" + mock_lm = MagicMock() + mock_lm.history = [] + with patch("dspy.settings") as mock_settings: + mock_settings.lm = mock_lm + dspy.configure(lm=mock_lm) + yield mock_lm + + +def _make_mock_result(unified_answer: str) -> MagicMock: + """Build a mock DSPy Predict result with the unified_answer attribute.""" + mock_result = MagicMock() + mock_result.unified_answer = unified_answer + return mock_result + + +def _make_async_iter(*chunks: Any) -> AsyncMock: + """Return an async iterable that yields the given chunks then closes cleanly.""" + + async def _gen() -> AsyncGenerator[Any, None]: + for chunk in chunks: + yield chunk + + mock_stream = AsyncMock() + mock_stream.__aiter__ = lambda self: _gen() + mock_stream.aclose = AsyncMock() + return mock_stream + + +# --------------------------------------------------------------------------- +# Initialisation +# --------------------------------------------------------------------------- + + +class TestMultiResponseFormatterModuleInit: + """MultiResponseFormatterModule should initialise with the correct attributes.""" + + def test_module_has_formatter_attribute(self) -> None: + module = MultiResponseFormatterModule() + assert hasattr(module, "formatter") + + def test_formatter_is_dspy_predict(self) -> None: + module = MultiResponseFormatterModule() + assert isinstance(module.formatter, dspy.Predict) + + def test_default_custom_instructions_is_empty_string(self) -> None: + module = MultiResponseFormatterModule() + assert module._custom_instructions == "" + + def test_custom_instructions_stored_on_instance(self) -> None: + module = MultiResponseFormatterModule( + custom_instructions="Use formal language." + ) + assert module._custom_instructions == "Use formal language." + + +# --------------------------------------------------------------------------- +# _build_results_block +# --------------------------------------------------------------------------- + + +class TestBuildResultsBlock: + """_build_results_block() should serialise results into labeled text sections.""" + + def test_empty_list_returns_no_results_marker(self) -> None: + result = MultiResponseFormatterModule._build_results_block([]) + assert "NO RESULTS" in result + + def test_single_result_contains_endpoint_name(self) -> None: + block = MultiResponseFormatterModule._build_results_block( + [ + ( + "get_current_weather", + "Too praegused ja kombineeritud ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn", "airtemperature": 12.5}]}}', + {}, + ) + ] + ) + assert "get_current_weather" in block + assert "Too praegused ja kombineeritud ilmaandmed" in block + assert "12.5" in block + + def test_single_result_section_header_format(self) -> None: + block = MultiResponseFormatterModule._build_results_block( + [ + ( + "get_current_weather", + "Too praegused ja kombineeritud ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ] + ) + assert "Result 1: get_current_weather" in block + + def test_multiple_results_all_sections_present(self) -> None: + results = [ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn", "airtemperature": 12.5}]}}', + {}, + ), + ( + "get_public_holidays", + "Too riiklikud pühad konkreetse riigi kohta", + '[{"startDate": "2026-01-01", "name": [{"language": "ET", "text": "Uusaasta"}]}]', + {}, + ), + ( + "get_public_holidays_lv", + "Too Läti riiklikud pühad", + '[{"startDate": "2026-11-18", "name": [{"language": "LV", "text": "Latvijas Republikas proklamēšanas diena"}]}]', + {}, + ), + ] + block = MultiResponseFormatterModule._build_results_block(results) + assert "Result 1: get_current_weather" in block + assert "Result 2: get_public_holidays" in block + assert "Result 3: get_public_holidays_lv" in block + + def test_null_response_annotated_as_empty(self) -> None: + block = MultiResponseFormatterModule._build_results_block( + [("get_current_weather", "Too praegused ilmaandmed", "null", {})] + ) + assert "EMPTY RESPONSE" in block + + def test_empty_list_response_annotated(self) -> None: + block = MultiResponseFormatterModule._build_results_block( + [("get_public_holidays", "Too riiklikud pühad", "[]", {})] + ) + assert "EMPTY RESPONSE" in block + + def test_empty_dict_response_annotated(self) -> None: + block = MultiResponseFormatterModule._build_results_block( + [("get_current_weather", "Too praegused ilmaandmed", "{}", {})] + ) + assert "EMPTY RESPONSE" in block + + def test_dict_api_response_serialised_to_json(self) -> None: + block = MultiResponseFormatterModule._build_results_block( + [ + ( + "get_current_weather", + "Too praegused ilmaandmed", + { + "observations": { + "station": [{"name": "Tallinn", "airtemperature": 12.5}] + } + }, + {}, + ) + ] + ) + assert "Tallinn" in block + assert "12.5" in block + + def test_list_api_response_serialised_to_json(self) -> None: + block = MultiResponseFormatterModule._build_results_block( + [ + ( + "get_public_holidays", + "Too riiklikud pühad", + [ + {"startDate": "2026-01-01", "name": [{"text": "Uusaasta"}]}, + { + "startDate": "2026-02-24", + "name": [{"text": "Eesti Vabariigi aastapäev"}], + }, + ], + {}, + ) + ] + ) + assert '"startDate"' in block + + def test_large_combined_response_truncated(self) -> None: + """Combined block exceeding _MAX_TOTAL_RESPONSE_BYTES should be truncated.""" + big_response = "x" * (_MAX_TOTAL_RESPONSE_BYTES + 10_000) + results = [ + ("get_current_weather", "Too praegused ilmaandmed", big_response, {}), + ("get_public_holidays", "Too riiklikud pühad", '{"small": true}', {}), + ] + block = MultiResponseFormatterModule._build_results_block(results) + block_bytes = len(block.encode("utf-8")) + # Allow small overshoot from truncation note text + assert block_bytes <= _MAX_TOTAL_RESPONSE_BYTES + 200 + + def test_large_combined_response_has_truncation_note(self) -> None: + big_response = "x" * (_MAX_TOTAL_RESPONSE_BYTES + 10_000) + block = MultiResponseFormatterModule._build_results_block( + [("get_current_weather", "Too praegused ilmaandmed", big_response, {})] + ) + assert "truncated" in block.lower() + + def test_per_result_truncation_applied(self) -> None: + """Individual responses over 50KB should be truncated by _truncate_if_needed.""" + large_items = [{"id": i, "data": "x" * 200} for i in range(600)] + block = MultiResponseFormatterModule._build_results_block( + [("get_public_holidays", "Too riiklikud pühad", large_items, {})] + ) + assert "NOTE" in block + + +# --------------------------------------------------------------------------- +# forward() — basic field mapping +# --------------------------------------------------------------------------- + + +class TestForwardFieldMapping: + """forward() must pass the correct keyword arguments to the DSPy predictor.""" + + def test_predictor_called_with_all_fields(self) -> None: + module = MultiResponseFormatterModule() + mock_result = _make_mock_result("Combined answer.") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + module.forward( + user_query="What is the current weather and upcoming public holidays in Estonia?", + api_results=[ + ( + "get_current_weather", + "Too praegused ja kombineeritud ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn", "airtemperature": 12.5}]}}', + {}, + ), + ( + "get_public_holidays", + "Too riiklikud pühad konkreetse riigi kohta", + '[{"startDate": "2026-01-01", "name": [{"language": "ET", "text": "Uusaasta"}]}]', + {}, + ), + ], + ) + + mock_formatter.assert_called_once() + call_kwargs = mock_formatter.call_args.kwargs + assert "user_query" in call_kwargs + assert "api_results_block" in call_kwargs + assert "response_language" in call_kwargs + assert "custom_instructions" in call_kwargs + assert "num_results" in call_kwargs + + def test_num_results_matches_input_length(self) -> None: + module = MultiResponseFormatterModule() + mock_result = _make_mock_result("Answer") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + module.forward( + user_query="Query", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ), + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01"}]', + {}, + ), + ( + "get_public_holidays_lv", + "Too Läti riiklikud pühad", + '[{"startDate": "2026-11-18"}]', + {}, + ), + ], + ) + + call_kwargs = mock_formatter.call_args.kwargs + assert call_kwargs["num_results"] == "3" + + def test_returns_unified_answer(self) -> None: + module = MultiResponseFormatterModule() + expected = "This is the synthesised answer." + mock_result = _make_mock_result(expected) + + with patch.object(module, "formatter", return_value=mock_result): + result = module.forward( + user_query="Query", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn", "airtemperature": 12.5}]}}', + {}, + ) + ], + ) + + assert result == expected + + def test_empty_api_results_still_calls_predictor(self) -> None: + module = MultiResponseFormatterModule() + mock_result = _make_mock_result("No results.") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + result = module.forward( + user_query="Anything", + api_results=[], + ) + + mock_formatter.assert_called_once() + assert result == "No results." + + +# --------------------------------------------------------------------------- +# forward() — language mapping +# --------------------------------------------------------------------------- + + +class TestForwardLanguageMapping: + """forward() must map ISO codes to display names.""" + + @pytest.mark.parametrize( + "language_code, expected_display", + [ + ("en", "English"), + ("et", "Estonian"), + ("ru", "Russian"), + ], + ) + def test_language_code_mapped_to_display_name( + self, language_code: str, expected_display: str + ) -> None: + module = MultiResponseFormatterModule() + mock_result = _make_mock_result("Answer") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + module.forward( + user_query="Test", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + detected_language=language_code, + ) + + assert mock_formatter.call_args.kwargs["response_language"] == expected_display + + def test_unknown_language_defaults_to_english(self) -> None: + module = MultiResponseFormatterModule() + mock_result = _make_mock_result("Answer") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + module.forward( + user_query="Test", + api_results=[ + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01"}]', + {}, + ) + ], + detected_language="fr", + ) + + assert mock_formatter.call_args.kwargs["response_language"] == "English" + + def test_default_language_is_english(self) -> None: + module = MultiResponseFormatterModule() + mock_result = _make_mock_result("Answer") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + module.forward( + user_query="Test", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + ) + + assert mock_formatter.call_args.kwargs["response_language"] == "English" + + +# --------------------------------------------------------------------------- +# forward() — custom instructions +# --------------------------------------------------------------------------- + + +class TestForwardCustomInstructions: + """Custom instructions must be forwarded verbatim to the predictor.""" + + def test_custom_instructions_passed_to_predictor(self) -> None: + module = MultiResponseFormatterModule( + custom_instructions="Always respond in Estonian." + ) + mock_result = _make_mock_result("Answer") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + module.forward( + user_query="Query", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + ) + + assert ( + mock_formatter.call_args.kwargs["custom_instructions"] + == "Always respond in Estonian." + ) + + def test_empty_custom_instructions_by_default(self) -> None: + module = MultiResponseFormatterModule() + mock_result = _make_mock_result("Answer") + + with patch.object( + module, "formatter", return_value=mock_result + ) as mock_formatter: + module.forward( + user_query="Query", + api_results=[ + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01"}]', + {}, + ) + ], + ) + + assert mock_formatter.call_args.kwargs["custom_instructions"] == "" + + +# --------------------------------------------------------------------------- +# forward() — error handling +# --------------------------------------------------------------------------- + + +class TestForwardErrorHandling: + """forward() must return a localized fallback if the predictor raises.""" + + @pytest.mark.parametrize("language_code", ["en", "et", "ru"]) + def test_predictor_exception_returns_localized_error( + self, language_code: str + ) -> None: + module = MultiResponseFormatterModule() + + with patch.object( + module, "formatter", side_effect=RuntimeError("LLM unavailable") + ): + result = module.forward( + user_query="Test", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + detected_language=language_code, + ) + + assert isinstance(result, str) + assert len(result) > 0 + assert result == _MULTI_FORMATTER_ERROR_MESSAGES[language_code] + + def test_unknown_language_exception_falls_back_to_english_error(self) -> None: + module = MultiResponseFormatterModule() + + with patch.object(module, "formatter", side_effect=RuntimeError("failure")): + result = module.forward( + user_query="Test", + api_results=[], + detected_language="fr", + ) + + assert result == _MULTI_FORMATTER_ERROR_MESSAGES["en"] + + +# --------------------------------------------------------------------------- +# stream_forward_multi — token streaming (Tier 1) +# --------------------------------------------------------------------------- + + +class TestStreamForwardMultiTokenStreaming: + """stream_forward_multi() should yield StreamResponse tokens for unified_answer.""" + + @pytest.mark.asyncio + async def test_stream_response_tokens_are_yielded(self) -> None: + module = MultiResponseFormatterModule() + + token1 = MagicMock(spec=dspy.streaming.StreamResponse) + token1.signature_field_name = "unified_answer" + token1.chunk = "Weather is " + + token2 = MagicMock(spec=dspy.streaming.StreamResponse) + token2.signature_field_name = "unified_answer" + token2.chunk = "sunny." + + mock_predictor = MagicMock(return_value=_make_async_iter(token1, token2)) + + with patch("dspy.streamify", return_value=mock_predictor): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="What's the current weather in Tallinn?", + api_results=[ + ( + "get_current_weather", + "Too praegused ja kombineeritud ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn", "airtemperature": 25}]}}', + {}, + ) + ], + ) + ] + + assert tokens == ["Weather is ", "sunny."] + + @pytest.mark.asyncio + async def test_tokens_for_wrong_field_not_yielded(self) -> None: + """Tokens for a different signature field must not be yielded.""" + module = MultiResponseFormatterModule() + + wrong_token = MagicMock(spec=dspy.streaming.StreamResponse) + wrong_token.signature_field_name = "other_field" + wrong_token.chunk = "should not appear" + + right_token = MagicMock(spec=dspy.streaming.StreamResponse) + right_token.signature_field_name = "unified_answer" + right_token.chunk = "correct" + + mock_predictor = MagicMock( + return_value=_make_async_iter(wrong_token, right_token) + ) + + with patch("dspy.streamify", return_value=mock_predictor): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Query", + api_results=[ + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01"}]', + {}, + ) + ], + ) + ] + + assert tokens == ["correct"] + + +# --------------------------------------------------------------------------- +# stream_forward_multi — Prediction fallback (Tier 2) +# --------------------------------------------------------------------------- + + +class TestStreamForwardMultiPredictionFallback: + """When no StreamResponse tokens arrive, fall back to Prediction.unified_answer.""" + + @pytest.mark.asyncio + async def test_prediction_fallback_yields_full_answer(self) -> None: + module = MultiResponseFormatterModule() + + prediction = MagicMock(spec=dspy.Prediction) + prediction.unified_answer = "Full synthesised answer." + + mock_predictor = MagicMock(return_value=_make_async_iter(prediction)) + + with patch("dspy.streamify", return_value=mock_predictor): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Query", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + ) + ] + + assert tokens == ["Full synthesised answer."] + + @pytest.mark.asyncio + async def test_prediction_without_unified_answer_falls_to_blocking(self) -> None: + """A Prediction with no unified_answer attribute triggers the blocking fallback.""" + module = MultiResponseFormatterModule() + + prediction = MagicMock(spec=dspy.Prediction) + prediction.unified_answer = None + + mock_predictor = MagicMock(return_value=_make_async_iter(prediction)) + + with patch("dspy.streamify", return_value=mock_predictor): + with patch.object(module, "forward", return_value="blocking result"): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Query", + api_results=[ + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01"}]', + {}, + ) + ], + ) + ] + + assert tokens == ["blocking result"] + + +# --------------------------------------------------------------------------- +# stream_forward_multi — blocking forward() fallback (Tier 3) +# --------------------------------------------------------------------------- + + +class TestStreamForwardMultiBlockingFallback: + """When streamify yields nothing, fall back to the blocking forward().""" + + @pytest.mark.asyncio + async def test_blocking_forward_used_when_no_output(self) -> None: + module = MultiResponseFormatterModule() + + mock_predictor = MagicMock(return_value=_make_async_iter()) + + with patch("dspy.streamify", return_value=mock_predictor): + with patch.object( + module, "forward", return_value="Blocking fallback." + ) as mock_fwd: + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Fallback test", + api_results=[("ep", "desc", '{"x": 1}', {})], + detected_language="en", + ) + ] + + mock_fwd.assert_called_once() + assert tokens == ["Blocking fallback."] + + @pytest.mark.asyncio + async def test_blocking_forward_receives_correct_args(self) -> None: + module = MultiResponseFormatterModule() + api_results = [ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn", "airtemperature": 12.5}]}}', + {}, + ), + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01", "name": [{"language": "ET", "text": "Uusaasta"}]}]', + {}, + ), + ] + + mock_predictor = MagicMock(return_value=_make_async_iter()) + + with patch("dspy.streamify", return_value=mock_predictor): + with patch.object(module, "forward", return_value="answer") as mock_fwd: + async for _ in module.stream_forward_multi( + user_query="Test query", + api_results=api_results, + detected_language="et", + ): + pass + + mock_fwd.assert_called_once_with( + user_query="Test query", + api_results=api_results, + detected_language="et", + ) + + +# --------------------------------------------------------------------------- +# stream_forward_multi — localized error fallback (Tier 4) +# --------------------------------------------------------------------------- + + +class TestStreamForwardMultiErrorFallback: + """On exception, stream_forward_multi must yield a localized error string.""" + + @pytest.mark.asyncio + async def test_exception_yields_localized_error(self) -> None: + module = MultiResponseFormatterModule() + + mock_predictor = MagicMock(side_effect=RuntimeError("stream broken")) + + with patch("dspy.streamify", return_value=mock_predictor): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Anything", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + detected_language="en", + ) + ] + + assert len(tokens) == 1 + assert isinstance(tokens[0], str) + assert tokens[0] == _MULTI_FORMATTER_ERROR_MESSAGES["en"] + + @pytest.mark.asyncio + @pytest.mark.parametrize("language_code", ["en", "et", "ru"]) + async def test_exception_yields_correct_language_error( + self, language_code: str + ) -> None: + module = MultiResponseFormatterModule() + + mock_predictor = MagicMock(side_effect=RuntimeError("failure")) + + with patch("dspy.streamify", return_value=mock_predictor): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Test", + api_results=[ + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01"}]', + {}, + ), + ], + detected_language=language_code, + ) + ] + + assert tokens[0] == _MULTI_FORMATTER_ERROR_MESSAGES[language_code] + + @pytest.mark.asyncio + async def test_unknown_language_exception_uses_english_error(self) -> None: + module = MultiResponseFormatterModule() + + mock_predictor = MagicMock(side_effect=RuntimeError("failure")) + + with patch("dspy.streamify", return_value=mock_predictor): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Test", + api_results=[], + detected_language="fr", + ) + ] + + assert tokens[0] == _MULTI_FORMATTER_ERROR_MESSAGES["en"] + + +# --------------------------------------------------------------------------- +# stream_forward_multi — custom instructions threading +# --------------------------------------------------------------------------- + + +class TestStreamForwardMultiCustomInstructions: + """stream_forward_multi must thread custom_instructions to the stream predictor.""" + + @pytest.mark.asyncio + async def test_custom_instructions_forwarded(self) -> None: + module = MultiResponseFormatterModule( + custom_instructions="Always respond in Estonian." + ) + + captured: dict[str, Any] = {} + + def fake_stream_predictor(**kwargs: Any) -> Any: + captured.update(kwargs) + return _make_async_iter() + + mock_predictor = MagicMock(side_effect=fake_stream_predictor) + + with patch("dspy.streamify", return_value=mock_predictor): + with patch.object(module, "forward", return_value="fallback"): + async for _ in module.stream_forward_multi( + user_query="Query", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + ): + pass + + assert captured.get("custom_instructions") == "Always respond in Estonian." + + @pytest.mark.asyncio + async def test_empty_custom_instructions_forwarded(self) -> None: + module = MultiResponseFormatterModule() + + captured: dict[str, Any] = {} + + def fake_stream_predictor(**kwargs: Any) -> Any: + captured.update(kwargs) + return _make_async_iter() + + mock_predictor = MagicMock(side_effect=fake_stream_predictor) + + with patch("dspy.streamify", return_value=mock_predictor): + with patch.object(module, "forward", return_value="fallback"): + async for _ in module.stream_forward_multi( + user_query="Query", + api_results=[ + ( + "get_public_holidays", + "Too riiklikud pühad", + '[{"startDate": "2026-01-01"}]', + {}, + ) + ], + ): + pass + + assert captured.get("custom_instructions") == "" + + +# --------------------------------------------------------------------------- +# stream_forward_multi — stream cleanup +# --------------------------------------------------------------------------- + + +class TestStreamForwardMultiCleanup: + """stream_forward_multi must call aclose() on the output stream.""" + + @pytest.mark.asyncio + async def test_aclose_called_on_stream(self) -> None: + module = MultiResponseFormatterModule() + mock_stream = _make_async_iter() + + mock_predictor = MagicMock(return_value=mock_stream) + + with patch("dspy.streamify", return_value=mock_predictor): + with patch.object(module, "forward", return_value="fallback"): + async for _ in module.stream_forward_multi( + user_query="Query", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + ): + pass + + mock_stream.aclose.assert_called_once() + + @pytest.mark.asyncio + async def test_aclose_called_even_on_exception(self) -> None: + """Exception during stream_forward_multi must not propagate — error is yielded.""" + module = MultiResponseFormatterModule() + + # Force the exception at the build_results_block step to test finally + with patch.object( + module, + "_build_results_block", + side_effect=RuntimeError("build error"), + ): + tokens = [ + t + async for t in module.stream_forward_multi( + user_query="Query", + api_results=[ + ( + "get_current_weather", + "Too praegused ilmaandmed", + '{"observations": {"station": [{"name": "Tallinn"}]}}', + {}, + ) + ], + detected_language="en", + ) + ] + + # Exception is caught — a localized error is yielded instead of propagating + assert len(tokens) == 1 + assert tokens[0] == _MULTI_FORMATTER_ERROR_MESSAGES["en"] diff --git a/tests/test_param_extractor.py b/tests/test_param_extractor.py index 9f1c0321..f47d2450 100644 --- a/tests/test_param_extractor.py +++ b/tests/test_param_extractor.py @@ -8,9 +8,9 @@ import dspy.streaming import pytest -from src.tool_classifier.param_extractor import ( +from tool_classifier.param_extractor import ( ParamExtractionModule, - _strip_format_hints, + strip_format_hints, ) @@ -745,12 +745,12 @@ def test_default_custom_instructions_is_empty_string(self) -> None: # --------------------------------------------------------------------------- -# _strip_format_hints helper +# strip_format_hints helper # --------------------------------------------------------------------------- class TestStripFormatHints: - """_strip_format_hints() should remove format hints from descriptions.""" + """strip_format_hints() should remove format hints from descriptions.""" @pytest.mark.parametrize( "description, expected", @@ -776,21 +776,19 @@ class TestStripFormatHints: ], ) def test_strips_known_patterns(self, description: str, expected: str) -> None: - assert _strip_format_hints(description) == expected + assert strip_format_hints(description) == expected def test_preserves_unrelated_parentheses(self) -> None: """Parentheses that don't contain format-like keywords are kept.""" desc = "City name (required)" # "required" does not match the keyword list, so it should be preserved - result = _strip_format_hints(desc) + result = strip_format_hints(desc) assert "required" in result def test_idempotent(self) -> None: """Calling twice produces the same result.""" desc = "Start date (YYYY-MM-DD)" - assert _strip_format_hints(_strip_format_hints(desc)) == _strip_format_hints( - desc - ) + assert strip_format_hints(strip_format_hints(desc)) == strip_format_hints(desc) # --------------------------------------------------------------------------- diff --git a/tests/test_qdrant_manager.py b/tests/test_qdrant_manager.py index 58b96778..73554999 100644 --- a/tests/test_qdrant_manager.py +++ b/tests/test_qdrant_manager.py @@ -103,8 +103,6 @@ def _make_qdrant_client( """Build a mock QdrantClient.""" client = MagicMock() - col_mock = MagicMock() - col_mock.name = "some_collection" collections_result = MagicMock() collections_result.collections = [MagicMock(name=n) for n in collection_names] # Fix: MagicMock(name=n) doesn't work as expected — set attribute explicitly diff --git a/tests/test_tool_classifier.py b/tests/test_tool_classifier.py index a9158b42..f1d3edbe 100644 --- a/tests/test_tool_classifier.py +++ b/tests/test_tool_classifier.py @@ -12,6 +12,8 @@ - Qdrant timeout during classification → fallback """ +from __future__ import annotations + from typing import Any, Dict, List, Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -21,6 +23,7 @@ from models.request_models import OrchestrationRequest from tool_classifier.classifier import ToolClassifier from tool_classifier.enums import WorkflowType +from tool_classifier.intent_decomposer import IntentDecomposerModule # --------------------------------------------------------------------------- @@ -280,7 +283,6 @@ async def test_active_session_returns_api_tool_calling(self) -> None: ): result = await classifier.classify( query="EE", - conversation_history=[], language="en", request=request, ) @@ -305,7 +307,6 @@ async def test_no_active_session_does_not_short_circuit(self) -> None: ): result = await classifier.classify( query="some query", - conversation_history=[], language="en", request=request, ) @@ -342,7 +343,6 @@ async def test_different_endpoint_match_deletes_old_session(self) -> None: ): result = await classifier.classify( query="What is the weather in Tallinn?", - conversation_history=[], language="en", request=request, ) @@ -378,7 +378,6 @@ async def test_api_tool_match_when_service_disabled(self) -> None: ): result = await classifier.classify( query="public holidays", - conversation_history=[], language="en", request=_make_request("public holidays"), ) @@ -406,7 +405,6 @@ async def test_context_fallback_when_service_disabled_and_no_api_match( ): result = await classifier.classify( query="tell me a joke", - conversation_history=[], language="en", request=_make_request("tell me a joke"), ) @@ -433,7 +431,6 @@ async def test_embedding_failure_falls_back_to_context(self) -> None: ): result = await classifier.classify( query="public holidays", - conversation_history=[], language="en", ) @@ -471,8 +468,286 @@ async def test_qdrant_timeout_falls_back_to_context(self) -> None: classifier.api_tool_searcher.search = AsyncMock(return_value=[]) result = await classifier.classify( query="public holidays", + language="en", + ) + + assert result.workflow == WorkflowType.CONTEXT + + +# --------------------------------------------------------------------------- +# IntentDecomposerModule +# --------------------------------------------------------------------------- + + +class TestIntentDecomposer: + """Unit tests for IntentDecomposerModule.forward() and .decompose(). + + The DSPy predictor is replaced with a MagicMock so no real LLM is called. + """ + + def _make_module_with_prediction( + self, + mode: str, + sub_queries: str, + ) -> IntentDecomposerModule: + module = IntentDecomposerModule() + mock_pred = MagicMock() + mock_pred.mode = mode + mock_pred.sub_queries = sub_queries + module.predictor = MagicMock(return_value=mock_pred) + return module + + def test_forward_single_mode_returns_single(self) -> None: + from tool_classifier.intent_decomposer import DecompositionResult + + module = self._make_module_with_prediction("single", "[]") + result = module.forward("What are the public holidays in Estonia?") + + assert isinstance(result, DecompositionResult) + assert result.mode == "single" + assert result.sub_queries == [] + + def test_forward_parallel_mode_returns_sub_queries(self) -> None: + from tool_classifier.intent_decomposer import DecompositionResult + + module = self._make_module_with_prediction( + "parallel", + '["public holidays in Estonia", "weather in Tallinn"]', + ) + result = module.forward( + "What are the public holidays in Estonia AND the weather in Tallinn?" + ) + + assert isinstance(result, DecompositionResult) + assert result.mode == "parallel" + assert len(result.sub_queries) == 2 + assert "public holidays in Estonia" in result.sub_queries + + def test_forward_caps_sub_queries_at_max_endpoints(self) -> None: + """sub_queries exceeding MULTI_API_MAX_ENDPOINTS are truncated.""" + from tool_classifier.intent_decomposer import ( + DecompositionResult, + IntentDecomposerModule, + ) + from tool_classifier.constants import MULTI_API_MAX_ENDPOINTS + + module = IntentDecomposerModule() + # Build a prediction with more sub-queries than the cap allows + over_cap = ["query " + str(i) for i in range(MULTI_API_MAX_ENDPOINTS + 2)] + import json + + mock_pred = MagicMock() + mock_pred.mode = "parallel" + mock_pred.sub_queries = json.dumps(over_cap) + module.predictor = MagicMock(return_value=mock_pred) + + result = module.forward("many intents query") + + assert isinstance(result, DecompositionResult) + assert result.mode == "parallel" + assert len(result.sub_queries) == MULTI_API_MAX_ENDPOINTS + + def test_forward_unexpected_mode_falls_back_to_single(self) -> None: + from tool_classifier.intent_decomposer import DecompositionResult + + module = self._make_module_with_prediction("unknown_value", "[]") + result = module.forward("some query") + + assert isinstance(result, DecompositionResult) + assert result.mode == "single" + assert result.sub_queries == [] + + def test_forward_parallel_with_fewer_than_2_sub_queries_falls_back(self) -> None: + """mode=parallel but only 1 sub-query parsed → conservative fallback to single.""" + from tool_classifier.intent_decomposer import DecompositionResult + + module = self._make_module_with_prediction("parallel", '["only one query"]') + result = module.forward("something") + + assert isinstance(result, DecompositionResult) + assert result.mode == "single" + + def test_forward_predictor_exception_falls_back_to_single(self) -> None: + """Any exception from the DSPy predictor → conservative single fallback.""" + from tool_classifier.intent_decomposer import ( + DecompositionResult, + IntentDecomposerModule, + ) + + module = IntentDecomposerModule() + module.predictor = MagicMock(side_effect=RuntimeError("LLM unavailable")) + + result = module.forward("multi intent query") + + assert isinstance(result, DecompositionResult) + assert result.mode == "single" + + def test_forward_parallel_with_invalid_json_falls_back_to_single(self) -> None: + """Invalid JSON in sub_queries → falls back to mode=single.""" + from tool_classifier.intent_decomposer import DecompositionResult + + module = self._make_module_with_prediction("parallel", "not valid json") + result = module.forward("holidays and weather") + + assert isinstance(result, DecompositionResult) + assert result.mode == "single" + + def test_forward_markdown_fenced_json_parsed_correctly(self) -> None: + """sub_queries wrapped in markdown code fences are unwrapped before JSON parse.""" + from tool_classifier.intent_decomposer import DecompositionResult + + fenced = '```json\n["query A", "query B"]\n```' + module = self._make_module_with_prediction("parallel", fenced) + result = module.forward("query") + + assert isinstance(result, DecompositionResult) + assert result.mode == "parallel" + assert result.sub_queries == ["query A", "query B"] + + @pytest.mark.asyncio + async def test_decompose_async_wraps_forward(self) -> None: + """.decompose() is the async wrapper — result matches .forward() output.""" + from tool_classifier.intent_decomposer import DecompositionResult + + module = self._make_module_with_prediction("parallel", '["sub A", "sub B"]') + + with patch( + "tool_classifier.intent_decomposer.asyncio.to_thread", + new_callable=AsyncMock, + ) as mock_thread: + mock_thread.return_value = DecompositionResult( + mode="parallel", sub_queries=["sub A", "sub B"] + ) + result = await module.decompose("two intents") + + assert isinstance(result, DecompositionResult) + assert result.mode == "parallel" + assert result.sub_queries == ["sub A", "sub B"] + + +# --------------------------------------------------------------------------- +# classify() — MULTI_INTENT_ENABLED feature flag toggle +# --------------------------------------------------------------------------- + + +class TestClassifyMultiIntentFeatureFlag: + """Verify the MULTI_INTENT_ENABLED flag gates the parallel decomposition path.""" + + @pytest.mark.asyncio + async def test_multi_intent_disabled_suppresses_hint_result(self) -> None: + """When MULTI_INTENT_ENABLED=False a multi_intent_hint result is suppressed + and the classifier falls through to CONTEXT/RAG.""" + from tool_classifier.api_semantic_searcher import APIToolSearchResult + + svc = _make_orchestration_service(session_store=None) + classifier = _make_classifier(svc) + + # Build a result with multi_intent_hint=True (disambiguator rejected all) + hint_result = MagicMock(spec=APIToolSearchResult) + hint_result.endpoint_id = "ep-holidays" + hint_result.name = "get_public_holidays" + hint_result.description = "Returns public holidays" + hint_result.method = "GET" + hint_result.url = "https://openholidaysapi.org/PublicHolidays" + hint_result.params = [] + hint_result.cosine_score = 0.55 + hint_result.rrf_score = 0.01 + hint_result.confidence = "medium" + hint_result.llm_validated = False + hint_result.multi_intent_hint = True + hint_result.to_dict.return_value = {"endpoint_id": "ep-holidays"} + + classifier.api_tool_searcher.search = AsyncMock(return_value=[hint_result]) + + with ( + patch( + "tool_classifier.classifier.FeatureFlags.API_TOOL_CALLING_WORKFLOW_ENABLED", + True, + ), + patch( + "tool_classifier.classifier.FeatureFlags.MULTI_INTENT_ENABLED", + False, + ), + patch( + "tool_classifier.classifier.FeatureFlags.SERVICE_WORKFLOW_ENABLED", + False, + ), + ): + result = await classifier.classify( + query="public holidays AND weather", conversation_history=[], language="en", + request=_make_request("public holidays AND weather"), ) assert result.workflow == WorkflowType.CONTEXT + + @pytest.mark.asyncio + async def test_multi_intent_enabled_triggers_decomposer_on_ambiguous_band( + self, + ) -> None: + """MULTI_INTENT_ENABLED=True + cosine in ambiguous band → IntentDecomposer runs.""" + from tool_classifier.api_semantic_searcher import APIToolSearchResult + from tool_classifier.intent_decomposer import DecompositionResult + + svc = _make_orchestration_service(session_store=None) + classifier = _make_classifier(svc) + + # A result in the ambiguous band (not llm_validated, not multi_intent_hint) + ambiguous_result = MagicMock(spec=APIToolSearchResult) + ambiguous_result.endpoint_id = "ep-holidays" + ambiguous_result.name = "get_public_holidays" + ambiguous_result.description = "Returns public holidays" + ambiguous_result.method = "GET" + ambiguous_result.url = "https://openholidaysapi.org/PublicHolidays" + ambiguous_result.params = [] + ambiguous_result.cosine_score = 0.50 # in [0.40, 0.60) band + ambiguous_result.rrf_score = 0.01 + ambiguous_result.confidence = "medium" + ambiguous_result.llm_validated = False + ambiguous_result.multi_intent_hint = False + ambiguous_result.to_dict.return_value = { + "endpoint_id": "ep-holidays", + "name": "get_public_holidays", + "description": "Returns public holidays", + "method": "GET", + "url": "https://openholidaysapi.org/PublicHolidays", + "params": [], + "cosine_score": 0.50, + "rrf_score": 0.01, + "confidence": "medium", + } + + classifier.api_tool_searcher.search = AsyncMock(return_value=[ambiguous_result]) + + # IntentDecomposer returns single → falls through to single-endpoint path + decomposer_result = DecompositionResult(mode="single", sub_queries=[]) + classifier.intent_decomposer.decompose = AsyncMock( + return_value=decomposer_result + ) + + with ( + patch( + "tool_classifier.classifier.FeatureFlags.API_TOOL_CALLING_WORKFLOW_ENABLED", + True, + ), + patch( + "tool_classifier.classifier.FeatureFlags.MULTI_INTENT_ENABLED", + True, + ), + patch( + "tool_classifier.classifier.FeatureFlags.SERVICE_WORKFLOW_ENABLED", + False, + ), + ): + result = await classifier.classify( + query="public holidays AND weather", + conversation_history=[], + language="en", + request=_make_request("public holidays AND weather"), + ) + + # Decomposer was consulted + classifier.intent_decomposer.decompose.assert_awaited_once() + # Single mode → normal API_TOOL_CALLING result + assert result.workflow == WorkflowType.API_TOOL_CALLING diff --git a/uv.lock b/uv.lock index a5204c8d..3cbac326 100644 --- a/uv.lock +++ b/uv.lock @@ -203,6 +203,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/83/7b/5652771e24fff12da9dde4c20ecf4682e606b104f26419d139758cc935a6/azure_identity-1.25.1-py3-none-any.whl", hash = "sha256:e9edd720af03dff020223cd269fa3a61e8f345ea75443858273bcb44844ab651", size = 191317, upload-time = "2025-10-06T20:30:04.251Z" }, ] +[[package]] +name = "azure-storage-blob" +version = "12.28.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "azure-core" }, + { name = "cryptography" }, + { name = "isodate" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/71/24/072ba8e27b0e2d8fec401e9969b429d4f5fc4c8d4f0f05f4661e11f7234a/azure_storage_blob-12.28.0.tar.gz", hash = "sha256:e7d98ea108258d29aa0efbfd591b2e2075fa1722a2fae8699f0b3c9de11eff41", size = 604225, upload-time = "2026-01-06T23:48:57.282Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d8/3a/6ef2047a072e54e1142718d433d50e9514c999a58f51abfff7902f3a72f8/azure_storage_blob-12.28.0-py3-none-any.whl", hash = "sha256:00fb1db28bf6a7b7ecaa48e3b1d5c83bfadacc5a678b77826081304bd87d6461", size = 431499, upload-time = "2026-01-06T23:48:58.995Z" }, +] + [[package]] name = "backoff" version = "2.2.1" @@ -977,6 +992,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/cb/b1/3846dd7f199d53cb17f49cba7e651e9ce294d8497c8c150530ed11865bb8/iniconfig-2.3.0-py3-none-any.whl", hash = "sha256:f631c04d2c48c52b84d0d0549c99ff3859c98df65b3101406327ecc7d53fbf12", size = 7484, upload-time = "2025-10-18T21:55:41.639Z" }, ] +[[package]] +name = "isodate" +version = "0.7.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/54/4d/e940025e2ce31a8ce1202635910747e5a87cc3a6a6bb2d00973375014749/isodate-0.7.2.tar.gz", hash = "sha256:4cd1aa0f43ca76f4a6c6c0292a85f40b35ec2e43e315b59f06e6d32171a953e6", size = 29705, upload-time = "2024-10-08T23:04:11.5Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/15/aa/0aca39a37d3c7eb941ba736ede56d689e7be91cab5d9ca846bde3999eba6/isodate-0.7.2-py3-none-any.whl", hash = "sha256:28009937d8031054830160fce6d409ed342816b543597cece116d966c6d99e15", size = 22320, upload-time = "2024-10-08T23:04:09.501Z" }, +] + [[package]] name = "jinja2" version = "3.1.6" @@ -2226,6 +2250,7 @@ source = { virtual = "." } dependencies = [ { name = "anthropic" }, { name = "azure-identity" }, + { name = "azure-storage-blob" }, { name = "boto3" }, { name = "deepeval" }, { name = "deepteam" }, @@ -2261,6 +2286,7 @@ dependencies = [ requires-dist = [ { name = "anthropic", specifier = ">=0.69.0" }, { name = "azure-identity", specifier = ">=1.24.0" }, + { name = "azure-storage-blob", specifier = ">=12.24.0" }, { name = "boto3", specifier = ">=1.40.25" }, { name = "deepeval", specifier = ">=3.6.0" }, { name = "deepteam", specifier = ">=0.2.5" }, @@ -2939,4 +2965,4 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ff/8d/0309daffea4fcac7981021dbf21cdb2e3427a9e76bafbcdbdf5392ff99a4/zstandard-0.25.0-cp312-cp312-win32.whl", hash = "sha256:23ebc8f17a03133b4426bcc04aabd68f8236eb78c3760f12783385171b0fd8bd", size = 436922, upload-time = "2025-09-14T22:17:24.398Z" }, { url = "https://files.pythonhosted.org/packages/79/3b/fa54d9015f945330510cb5d0b0501e8253c127cca7ebe8ba46a965df18c5/zstandard-0.25.0-cp312-cp312-win_amd64.whl", hash = "sha256:ffef5a74088f1e09947aecf91011136665152e0b4b359c42be3373897fb39b01", size = 506276, upload-time = "2025-09-14T22:17:21.429Z" }, { url = "https://files.pythonhosted.org/packages/ea/6b/8b51697e5319b1f9ac71087b0af9a40d8a6288ff8025c36486e0c12abcc4/zstandard-0.25.0-cp312-cp312-win_arm64.whl", hash = "sha256:181eb40e0b6a29b3cd2849f825e0fa34397f649170673d385f3598ae17cca2e9", size = 462679, upload-time = "2025-09-14T22:17:23.147Z" }, -] +] \ No newline at end of file diff --git a/vault-init.sh b/vault-init.sh index 63db07a2..44af6bbf 100644 --- a/vault-init.sh +++ b/vault-init.sh @@ -7,6 +7,88 @@ INIT_FLAG="/vault/data/.initialized" echo "=== Vault Initialization Script ===" +# --------------------------------------------------------------------------- +# Helpers (used by the SUBSEQUENT DEPLOYMENT branch) +# --------------------------------------------------------------------------- + +# Ensure a role_id file exists on disk; fetch from Vault if missing. +# Usage: ensure_role_id +ensure_role_id() { + role="$1"; rid_file="$2" + if [ -f "$rid_file" ] && [ -s "$rid_file" ]; then + return 0 + fi + echo "Fetching role_id for $role..." + rid=$(wget -q -O- \ + --header="X-Vault-Token: $ROOT_TOKEN" \ + "$VAULT_ADDR/v1/auth/approle/role/$role/role-id" | \ + grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') + echo "$rid" > "$rid_file" + chmod 640 "$rid_file" +} + +# Return 0 if the on-disk role_id + secret_id still authenticate, 1 otherwise. +# Usage: validate_secret_id +validate_secret_id() { + rid_file="$1"; sid_file="$2" + [ -f "$rid_file" ] && [ -f "$sid_file" ] || return 1 + rid=$(cat "$rid_file"); sid=$(cat "$sid_file") + [ -n "$rid" ] && [ -n "$sid" ] || return 1 + # wget returns non-zero on HTTP 400 (invalid creds); also confirm a token came back. + resp=$(wget -q -O- \ + --post-data="{\"role_id\":\"$rid\",\"secret_id\":\"$sid\"}" \ + --header='Content-Type: application/json' \ + "$VAULT_ADDR/v1/auth/approle/login" 2>/dev/null) || return 1 + echo "$resp" | grep -q '"client_token"' || return 1 + return 0 +} + +# Mint a fresh secret_id for a role and write it to disk. +# Usage: mint_secret_id +mint_secret_id() { + role="$1"; sid_file="$2" + sid=$(wget -q -O- --post-data='' \ + --header="X-Vault-Token: $ROOT_TOKEN" \ + "$VAULT_ADDR/v1/auth/approle/role/$role/secret-id" | \ + grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') + echo "$sid" > "$sid_file" + chmod 640 "$sid_file" +} + +# Reuse the existing secret_id if it still authenticates; otherwise mint a new one. +# Usage: reconcile_secret_id +reconcile_secret_id() { + role="$1"; rid_file="$2"; sid_file="$3" + ensure_role_id "$role" "$rid_file" + if validate_secret_id "$rid_file" "$sid_file"; then + echo "$role: existing secret_id still valid - reusing" + else + echo "$role: secret_id invalid or missing - minting a new one" + mint_secret_id "$role" "$sid_file" + fi +} + +# Create or update an AppRole that issues a PERIODIC token (no max_ttl): the +# agent renews it forever and never re-runs approle/login in steady state. +# secret_id_ttl=0 + secret_id_num_uses=0 keep the secret_id valid across +# restarts. Idempotent: does not invalidate existing secret_ids, safe per run. +# Usage: upsert_approle +upsert_approle() { + role="$1"; policy="$2"; period="$3" + wget -q -O- --post-data='{"token_policies":["'"$policy"'"],"token_period":"'"$period"'","token_num_uses":0,"secret_id_ttl":"0","secret_id_num_uses":0,"bind_secret_id":true}' \ + --header="X-Vault-Token: $ROOT_TOKEN" \ + --header='Content-Type: application/json' \ + "$VAULT_ADDR/v1/auth/approle/role/$role" >/dev/null +} + +# Apply the current AppRole definitions for all three services. +ensure_approles() { + echo "Ensuring AppRole configs (periodic tokens)..." + upsert_approle "gui-service" "gui-policy" "20m" + upsert_approle "cron-manager-service" "cron-manager-policy" "30m" + upsert_approle "llm-orchestration-service" "llm-orchestration-policy" "1h" +} + # Wait for Vault to be ready echo "Waiting for Vault..." for i in $(seq 1 30); do @@ -106,6 +188,8 @@ path "secret/metadata/llm/connections/*" { capabilities = ["read", "list"] } path "secret/data/embeddings/connections/*" { capabilities = ["read", "list"] } path "secret/metadata/embeddings/connections/*" { capabilities = ["read", "list"] } path "secret/data/encryption/*" { capabilities = ["deny"] } +path "secret/data/langfuse/*" { capabilities = ["read"] } +path "secret/metadata/langfuse/*" { capabilities = ["read", "list"] } path "auth/token/lookup-self" { capabilities = ["read"] }' LLM_POLICY_JSON=$(echo "$LLM_POLICY" | jq -Rs '{"policy":.}') @@ -114,27 +198,9 @@ path "auth/token/lookup-self" { capabilities = ["read"] }' --header='Content-Type: application/json' \ "$VAULT_ADDR/v1/sys/policies/acl/llm-orchestration-policy" >/dev/null - # Create GUI AppRole - echo "Creating gui-service AppRole..." - wget -q -O- --post-data='{"token_policies":["gui-policy"],"token_no_default_policy":true,"token_ttl":"15m","token_max_ttl":"1h","secret_id_ttl":"24h","secret_id_num_uses":0,"bind_secret_id":true}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/auth/approle/role/gui-service" >/dev/null - - # Create CronManager AppRole - echo "Creating cron-manager-service AppRole..." - wget -q -O- --post-data='{"token_policies":["cron-manager-policy"],"token_no_default_policy":true,"token_ttl":"30m","token_max_ttl":"8h","secret_id_ttl":"24h","secret_id_num_uses":0,"bind_secret_id":true}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/auth/approle/role/cron-manager-service" >/dev/null - - # Create LLM Orchestration AppRole - echo "Creating llm-orchestration-service AppRole..." - wget -q -O- --post-data='{"token_policies":["llm-orchestration-policy"],"token_no_default_policy":true,"token_ttl":"1h","token_max_ttl":"24h","secret_id_ttl":"24h","secret_id_num_uses":0,"bind_secret_id":true}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/auth/approle/role/llm-orchestration-service" >/dev/null - + # Create the three AppRoles (periodic tokens - see upsert_approle). + ensure_approles + # Ensure credentials directory exists mkdir -p /agent/credentials @@ -229,13 +295,6 @@ path "auth/token/lookup-self" { capabilities = ["read"] }' rm -rf "$TEMP_KEY_DIR" echo "RSA keypair generated and stored successfully" - # Store test LLM credentials for testing - echo "Creating test LLM credentials..." - wget -q -O- --post-data='{"data":{"access_key":"TEST_AWS_ACCESS_KEY","secret_key":"TEST_AWS_SECRET_KEY","environment":"production","model":"claude-3"}}' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - --header='Content-Type: application/json' \ - "$VAULT_ADDR/v1/secret/data/llm/connections/aws_bedrock/production/claude-3" >/dev/null - # Mark as initialized touch "$INIT_FLAG" echo "=== First time setup complete ===" @@ -276,65 +335,22 @@ else # Get root token ROOT_TOKEN=$(grep -o '"root_token":"[^"]*"' "$UNSEAL_KEYS_FILE" | cut -d':' -f2 | tr -d '"') export VAULT_TOKEN="$ROOT_TOKEN" - + + # Re-apply AppRole definitions so config changes (e.g. periodic tokens) + # take effect on redeploy without re-initializing Vault. Idempotent and + # does not invalidate existing secret_ids. + ensure_approles + # Ensure credentials directory exists mkdir -p /agent/credentials - # Always regenerate all secret_ids on restart - echo "Regenerating GUI secret_id..." - GUI_SECRET_ID=$(wget -q -O- --post-data='' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/gui-service/secret-id" | \ - grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$GUI_SECRET_ID" > /agent/credentials/gui_secret_id - - echo "Regenerating CronManager secret_id..." - CRON_SECRET_ID=$(wget -q -O- --post-data='' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/cron-manager-service/secret-id" | \ - grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$CRON_SECRET_ID" > /agent/credentials/cron_secret_id - - echo "Regenerating LLM secret_id..." - LLM_SECRET_ID=$(wget -q -O- --post-data='' \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/llm-orchestration-service/secret-id" | \ - grep -o '"secret_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$LLM_SECRET_ID" > /agent/credentials/llm_secret_id - - # Set permissions - chmod 640 /agent/credentials/*_secret_id - - # Ensure role_ids exist - if [ ! -f /agent/credentials/gui_role_id ]; then - echo "Copying GUI role_id..." - GUI_ROLE_ID=$(wget -q -O- \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/gui-service/role-id" | \ - grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$GUI_ROLE_ID" > /agent/credentials/gui_role_id - chmod 640 /agent/credentials/gui_role_id - fi - - if [ ! -f /agent/credentials/cron_role_id ]; then - echo "Copying CronManager role_id..." - CRON_ROLE_ID=$(wget -q -O- \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/cron-manager-service/role-id" | \ - grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$CRON_ROLE_ID" > /agent/credentials/cron_role_id - chmod 640 /agent/credentials/cron_role_id - fi - - if [ ! -f /agent/credentials/llm_role_id ]; then - echo "Copying LLM role_id..." - LLM_ROLE_ID=$(wget -q -O- \ - --header="X-Vault-Token: $ROOT_TOKEN" \ - "$VAULT_ADDR/v1/auth/approle/role/llm-orchestration-service/role-id" | \ - grep -o '"role_id":"[^"]*"' | cut -d':' -f2 | tr -d '"') - echo "$LLM_ROLE_ID" > /agent/credentials/llm_role_id - chmod 640 /agent/credentials/llm_role_id - fi + # Reconcile secret_ids: reuse the existing one if it still authenticates, + # mint a new one only if invalid or missing - keeps one stable secret_id + # across restarts instead of rotating every boot. reconcile_secret_id also + # ensures the role_id file exists first (validation needs both). + reconcile_secret_id "gui-service" /agent/credentials/gui_role_id /agent/credentials/gui_secret_id + reconcile_secret_id "cron-manager-service" /agent/credentials/cron_role_id /agent/credentials/cron_secret_id + reconcile_secret_id "llm-orchestration-service" /agent/credentials/llm_role_id /agent/credentials/llm_secret_id fi echo "=== Vault init complete ===" \ No newline at end of file diff --git a/vault/agents/cron/cron-agent.hcl b/vault/agents/cron/cron-agent.hcl index f2db227e..9454c9b7 100644 --- a/vault/agents/cron/cron-agent.hcl +++ b/vault/agents/cron/cron-agent.hcl @@ -2,7 +2,9 @@ # This agent provides CronManager with access to encryption keys and write access to secrets vault { - address = "http://vault:8200" + # Local testing: use rag-vault, not bare "vault" — that name collides with the + # ckb stack on the shared bykstack network and authenticates the wrong Vault. + address = "http://rag-vault:8200" retry { num_retries = 5 } @@ -42,6 +44,4 @@ listener "tcp" { # API proxy configuration api_proxy { use_auto_auth_token = true - enforce_consistency = "always" - when_inconsistent = "forward" } diff --git a/vault/agents/gui/gui-agent.hcl b/vault/agents/gui/gui-agent.hcl index a28db871..672d6d4d 100644 --- a/vault/agents/gui/gui-agent.hcl +++ b/vault/agents/gui/gui-agent.hcl @@ -2,7 +2,9 @@ # This agent provides GUI with access to public encryption key only vault { - address = "http://vault:8200" + # Local testing: use rag-vault, not bare "vault" — that name collides with the + # ckb stack on the shared bykstack network and authenticates the wrong Vault. + address = "http://rag-vault:8200" retry { num_retries = 5 } @@ -42,6 +44,4 @@ listener "tcp" { # API proxy configuration api_proxy { use_auto_auth_token = true - enforce_consistency = "always" - when_inconsistent = "forward" } diff --git a/vault/agents/llm/agent.hcl b/vault/agents/llm/agent.hcl index d7237be7..1a575260 100644 --- a/vault/agents/llm/agent.hcl +++ b/vault/agents/llm/agent.hcl @@ -1,5 +1,7 @@ vault { - address = "http://vault:8200" + # Local testing: use rag-vault, not bare "vault" — that name collides with the + # ckb stack on the shared bykstack network and authenticates the wrong Vault. + address = "http://rag-vault:8200" retry { num_retries = 5 } @@ -34,6 +36,4 @@ listener "tcp" { api_proxy { use_auto_auth_token = true - enforce_consistency = "always" - when_inconsistent = "forward" } diff --git a/vault/config/vault.hcl b/vault/config/vault.hcl index eaef415a..64ab325e 100644 --- a/vault/config/vault.hcl +++ b/vault/config/vault.hcl @@ -1,22 +1,27 @@ # HashiCorp Vault Server Configuration -# Production-ready configuration for LLM Orchestration Service +# Single-node Raft for the RAG-Module services -# Storage backend - Raft for high availability +# Storage backend - Raft storage "raft" { path = "/vault/file" node_id = "vault-node-1" - - # Retry join configuration for clustering (single node for now) - retry_join { - leader_api_addr = "http://vault:8200" - } + + # NOTE: No retry_join for a single node. A lone node self-bootstraps. + # A retry_join pointing at itself causes repeated + # "failed to get raft challenge ... Vault is sealed" errors and a + # messy double Raft init on every boot. Add retry_join back only when + # you actually have peer nodes to join. } -# HTTP listener configuration +# HTTP API listener. +# Vault automatically uses the next port up (8201) as its internal +# cluster port, so do NOT define a separate listener on 8201 — that +# collides with the cluster listener ("bind: address already in use") +# and degrades the login/request-forwarding path the agents rely on. listener "tcp" { - address = "0.0.0.0:8200" - tls_disable = true - + address = "0.0.0.0:8200" + tls_disable = true + # Enable CORS for web UI access cors_enabled = true cors_allowed_origins = [ @@ -25,14 +30,9 @@ listener "tcp" { ] } -# Cluster listener for HA (required even for single node) -listener "tcp" { - address = "0.0.0.0:8201" - cluster_addr = "http://0.0.0.0:8201" - tls_disable = true -} - -# API and cluster addresses +# API and cluster addresses. +# cluster_addr tells Vault where its internal cluster port (8201) is +# reachable; Vault binds that port itself — no listener block needed. api_addr = "http://vault:8200" cluster_addr = "http://vault:8201" @@ -46,9 +46,5 @@ default_lease_ttl = "168h" # 7 days max_lease_ttl = "720h" # 30 days # Logging configuration -log_level = "INFO" +log_level = "INFO" log_format = "json" - -# Development settings (remove in production) -# Note: In production, you should not use dev mode -# and should properly initialize and unseal the vault \ No newline at end of file