fix(api): remove unsafe Bedrock test endpoint (#40742)

This commit is contained in:
WH-2099 2026-08-13 19:47:13 +00:00 committed by GitHub
parent 8238684334
commit c1bde25bcc
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
13 changed files with 2 additions and 506 deletions

View File

@ -335,22 +335,6 @@ class CeleryConfig(DatabaseConfig):
return self.CELERY_BROKER_URL.startswith("rediss://") if self.CELERY_BROKER_URL else False
class InternalTestConfig(BaseSettings):
"""
Configuration settings for Internal Test
"""
AWS_SECRET_ACCESS_KEY: str | None = Field(
description="Internal test AWS secret access key",
default=None,
)
AWS_ACCESS_KEY_ID: str | None = Field(
description="Internal test AWS access key ID",
default=None,
)
class DatasetQueueMonitorConfig(BaseSettings):
"""
Configuration settings for Dataset Queue Monitor
@ -416,7 +400,6 @@ class MiddlewareConfig(
WeaviateConfig,
ElasticsearchConfig,
CouchbaseConfig,
InternalTestConfig,
VikingDBConfig,
UpstashConfig,
TidbOnQdrantConfig,

View File

@ -42,7 +42,6 @@ from services.enterprise import rbac_service as enterprise_rbac_service
from services.entities.external_knowledge_entities.external_knowledge_entities import ExternalDatasetCreatePayload
from services.external_knowledge_service import ExternalDatasetService
from services.hit_testing_service import HitTestingService
from services.knowledge_service import BedrockRetrievalSetting, ExternalDatasetTestService
class ExternalKnowledgeApiPayload(BaseModel):
@ -56,12 +55,6 @@ class ExternalHitTestingPayload(BaseModel):
metadata_filtering_conditions: dict[str, Any] | None = None
class BedrockRetrievalPayload(BaseModel):
retrieval_setting: BedrockRetrievalSetting
query: str
knowledge_id: str
class ExternalApiTemplateListQuery(BaseModel):
page: int = Field(default=1, description="Page number")
limit: int = Field(default=20, description="Number of items per page")
@ -137,23 +130,11 @@ class ExternalHitTestingResponse(ResponseModel):
records: list[ExternalHitTestingRecordResponse]
class BedrockRetrievalRecordResponse(ResponseModel):
metadata: dict[str, Any] | None = None
score: float
title: str | None = None
content: str | None = None
class BedrockRetrievalResponse(ResponseModel):
records: list[BedrockRetrievalRecordResponse]
register_schema_models(
console_ns,
ExternalKnowledgeApiPayload,
ExternalDatasetCreatePayload,
ExternalHitTestingPayload,
BedrockRetrievalPayload,
ExternalApiTemplateListQuery,
)
register_response_schema_models(
@ -166,8 +147,6 @@ register_response_schema_models(
ExternalHitTestingQueryResponse,
ExternalHitTestingRecordResponse,
ExternalHitTestingResponse,
BedrockRetrievalRecordResponse,
BedrockRetrievalResponse,
)
@ -451,20 +430,3 @@ class ExternalKnowledgeHitTestingApi(Resource):
return dump_response(ExternalHitTestingResponse, response)
except Exception as e:
raise InternalServerError(str(e))
@console_ns.route("/test/retrieval")
class BedrockRetrievalApi(Resource):
# this api is only for internal testing
@console_ns.doc("bedrock_retrieval_test")
@console_ns.doc(description="Bedrock retrieval test (internal use only)")
@console_ns.expect(console_ns.models[BedrockRetrievalPayload.__name__])
@console_ns.response(200, "Bedrock retrieval test completed", console_ns.models[BedrockRetrievalResponse.__name__])
@model_validate(BedrockRetrievalPayload)
def post(self, req_data: BedrockRetrievalPayload):
# Call the knowledge retrieval service
result = ExternalDatasetTestService.knowledge_retrieval(
req_data.retrieval_setting, req_data.query, req_data.knowledge_id
)
return dump_response(BedrockRetrievalResponse, result), 200

View File

@ -9615,21 +9615,6 @@ Remove one or more tag bindings from a target.
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [TagResponse](#tagresponse)<br> |
### [POST] /test/retrieval
Bedrock retrieval test (internal use only)
#### Request Body
| Required | Schema |
| -------- | ------ |
| Yes | **application/json**: [BedrockRetrievalPayload](#bedrockretrievalpayload)<br> |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Bedrock retrieval test completed | **application/json**: [BedrockRetrievalResponse](#bedrockretrievalresponse)<br> |
### [GET] /trial-apps/{app_id}
**Get app detail**
@ -15905,38 +15890,6 @@ ExporleBanner status
| ---- | ---- | ----------- | -------- |
| upload_file_id | string | | Yes |
#### BedrockRetrievalPayload
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| knowledge_id | string | | Yes |
| query | string | | Yes |
| retrieval_setting | [BedrockRetrievalSetting](#bedrockretrievalsetting) | | Yes |
#### BedrockRetrievalRecordResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| content | string | | No |
| metadata | object | | No |
| score | number | | Yes |
| title | string | | No |
#### BedrockRetrievalResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| records | [ [BedrockRetrievalRecordResponse](#bedrockretrievalrecordresponse) ] | | Yes |
#### BedrockRetrievalSetting
Retrieval settings for Amazon Bedrock knowledge base queries.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| score_threshold | number | Minimum relevance score threshold | No |
| top_k | integer | Maximum number of results to retrieve | No |
#### BillingInvoiceResponse
| Name | Type | Description | Required |

View File

@ -1,53 +0,0 @@
import boto3
from pydantic import BaseModel, Field
from configs import dify_config
class BedrockRetrievalSetting(BaseModel):
"""Retrieval settings for Amazon Bedrock knowledge base queries."""
top_k: int | None = Field(default=None, description="Maximum number of results to retrieve")
score_threshold: float = Field(default=0.0, description="Minimum relevance score threshold")
class ExternalDatasetTestService:
# this service is only for internal testing
@staticmethod
def knowledge_retrieval(retrieval_setting: BedrockRetrievalSetting, query: str, knowledge_id: str):
# get bedrock client
client = boto3.client(
"bedrock-agent-runtime",
aws_secret_access_key=dify_config.AWS_SECRET_ACCESS_KEY,
aws_access_key_id=dify_config.AWS_ACCESS_KEY_ID,
# example: us-east-1
region_name="us-east-1",
)
# fetch external knowledge retrieval
response = client.retrieve(
knowledgeBaseId=knowledge_id,
retrievalConfiguration={
"vectorSearchConfiguration": {
"numberOfResults": retrieval_setting.top_k,
"overrideSearchType": "HYBRID",
}
},
retrievalQuery={"text": query},
)
# parse response
results = []
if response.get("ResponseMetadata") and response.get("ResponseMetadata").get("HTTPStatusCode") == 200:
if response.get("retrievalResults"):
retrieval_results = response.get("retrievalResults")
for retrieval_result in retrieval_results:
# filter out results with score less than threshold
if retrieval_result.get("score") < retrieval_setting.score_threshold:
continue
result = {
"metadata": retrieval_result.get("metadata"),
"score": retrieval_result.get("score"),
"title": retrieval_result.get("metadata").get("x-amz-bedrock-kb-source-uri"),
"content": retrieval_result.get("content").get("text"),
}
results.append(result)
return {"records": results}

View File

@ -12,8 +12,6 @@ import services
from controllers.console import console_ns
from controllers.console.datasets.error import DatasetNameDuplicateError
from controllers.console.datasets.external import (
BedrockRetrievalApi,
BedrockRetrievalPayload,
ExternalApiTemplateApi,
ExternalApiTemplateListApi,
ExternalApiTemplateListQuery,
@ -28,7 +26,6 @@ from services.dataset_service import DatasetService
from services.entities.external_knowledge_entities.external_knowledge_entities import ExternalDatasetCreatePayload
from services.external_knowledge_service import ExternalDatasetService
from services.hit_testing_service import HitTestingService
from services.knowledge_service import ExternalDatasetTestService
@pytest.fixture
@ -539,52 +536,6 @@ class TestExternalKnowledgeHitTestingApi(_UsesSQLiteSession):
)
class TestBedrockRetrievalApi:
def test_bedrock_retrieval(self, app: Flask):
api = BedrockRetrievalApi()
method = inspect.unwrap(api.post)
payload = {
"retrieval_setting": {"top_k": 5, "score_threshold": 0.72},
"query": "hello bedrock",
"knowledge_id": "knowledge-base-1",
}
retrieval_response = {
"records": [
{
"metadata": {"source": "bedrock", "uri": "s3://bucket/doc.txt"},
"score": 0.8,
"title": "doc",
"content": "answer",
},
{
"metadata": {"source": "bedrock", "uri": "s3://bucket/other.txt"},
"score": 0.65,
"title": None,
"content": None,
},
]
}
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
patch.object(
ExternalDatasetTestService,
"knowledge_retrieval",
return_value=retrieval_response,
) as knowledge_retrieval,
):
resp, status = method(api, BedrockRetrievalPayload.model_validate(payload))
assert status == 200
assert resp == retrieval_response
retrieval_setting, query, knowledge_id = knowledge_retrieval.call_args.args
assert retrieval_setting.model_dump() == payload["retrieval_setting"]
assert query == "hello bedrock"
assert knowledge_id == "knowledge-base-1"
class TestExternalApiTemplateListApiAdvanced(_UsesSQLiteSession):
def test_post_duplicate_name_error(self, app: Flask, current_user: Account):
api = ExternalApiTemplateListApi()
@ -733,26 +684,3 @@ class TestExternalKnowledgeHitTestingApiAdvanced(_UsesSQLiteSession):
external_retrieval_model={"type": "bm25"},
metadata_filtering_conditions={"status": "active"},
)
class TestBedrockRetrievalApiAdvanced:
def test_bedrock_retrieval_with_invalid_setting(self, app: Flask):
api = BedrockRetrievalApi()
method = inspect.unwrap(api.post)
payload = {
"retrieval_setting": {},
"query": "test",
"knowledge_id": "k-1",
}
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", payload),
patch(
"controllers.console.datasets.external.ExternalDatasetTestService.knowledge_retrieval",
side_effect=ValueError("Invalid settings"),
),
):
with pytest.raises(ValueError):
method(api, BedrockRetrievalPayload.model_validate(payload))

View File

@ -180,6 +180,8 @@ def test_openapi_json_endpoints_render(monkeypatch: pytest.MonkeyPatch):
assert "paths" in payload
assert "schemas" in payload["components"]
assert isinstance(payload["components"]["schemas"], dict)
if route == "/console/api/openapi.json":
assert "/test/retrieval" not in payload["paths"]
missing_refs = _schema_refs(payload) - set(payload["components"]["schemas"])
assert not missing_refs
get_request_body_paths = [path for path, operation in _get_operations(payload) if "requestBody" in operation]

View File

@ -858,7 +858,6 @@ project-excludes = [
"services/test_human_input_delivery_test_service.py",
"services/test_human_input_file_upload_service.py",
"services/test_knowledge_retrieval_inner_service.py",
"services/test_knowledge_service.py",
"services/test_message_service.py",
"services/test_messages_clean_service.py",
"services/test_metadata_nullable_bug.py",

View File

@ -1,155 +0,0 @@
from typing import Any, cast
from unittest.mock import MagicMock, patch
import pytest
from services.knowledge_service import BedrockRetrievalSetting, ExternalDatasetTestService
class TestKnowledgeService:
"""Test suite for ExternalDatasetTestService"""
# ===== Happy Path Tests =====
@patch("services.knowledge_service.boto3.client")
@patch("services.knowledge_service.dify_config")
def test_knowledge_retrieval_should_succeed_with_valid_results(
self, mock_dify_config: MagicMock, mock_boto_client: MagicMock
):
"""Test that knowledge_retrieval successfully parses results from Bedrock"""
# Arrange
mock_dify_config.AWS_SECRET_ACCESS_KEY = "dummy_secret"
mock_dify_config.AWS_ACCESS_KEY_ID = "dummy_id"
mock_client = MagicMock()
mock_boto_client.return_value = mock_client
retrieval_setting = BedrockRetrievalSetting(top_k=4, score_threshold=0.5)
query = "test query"
knowledge_id = "kb-123"
# Mock successful response
mock_client.retrieve.return_value = {
"ResponseMetadata": {"HTTPStatusCode": 200},
"retrievalResults": [
{
"score": 0.9,
"metadata": {"x-amz-bedrock-kb-source-uri": "s3://bucket/doc1.pdf"},
"content": {"text": "content from doc1"},
},
{
"score": 0.4, # Below threshold
"metadata": {"x-amz-bedrock-kb-source-uri": "s3://bucket/doc2.pdf"},
"content": {"text": "content from doc2"},
},
],
}
# Act
result = cast(
dict[str, Any], ExternalDatasetTestService.knowledge_retrieval(retrieval_setting, query, knowledge_id)
)
# Assert
assert len(result["records"]) == 1
record = result["records"][0]
assert record["score"] == 0.9
assert record["title"] == "s3://bucket/doc1.pdf"
assert record["content"] == "content from doc1"
# verify retrieve called correctly
mock_client.retrieve.assert_called_once_with(
knowledgeBaseId=knowledge_id,
retrievalConfiguration={
"vectorSearchConfiguration": {
"numberOfResults": 4,
"overrideSearchType": "HYBRID",
}
},
retrievalQuery={"text": query},
)
# NEW: verify boto3.client created with proper service name and config values
mock_boto_client.assert_called_once_with(
"bedrock-agent-runtime",
aws_secret_access_key="dummy_secret",
aws_access_key_id="dummy_id",
region_name="us-east-1",
)
@patch("services.knowledge_service.boto3.client")
def test_knowledge_retrieval_should_return_empty_when_no_results(self, mock_boto: MagicMock):
"""Test that knowledge_retrieval returns empty records when Bedrock returns nothing"""
# Arrange
mock_client = MagicMock()
mock_boto.return_value = mock_client
mock_client.retrieve.return_value = {"ResponseMetadata": {"HTTPStatusCode": 200}, "retrievalResults": []}
# Act
result = cast(
dict[str, Any],
ExternalDatasetTestService.knowledge_retrieval(BedrockRetrievalSetting(top_k=1), "query", "kb"),
)
# Assert
assert result["records"] == []
# ===== Error Handling Tests =====
@patch("services.knowledge_service.boto3.client")
def test_knowledge_retrieval_should_return_empty_on_http_error(self, mock_boto: MagicMock):
"""Test that knowledge_retrieval returns empty records if Bedrock returns non-200 status"""
# Arrange
mock_client = MagicMock()
mock_boto.return_value = mock_client
mock_client.retrieve.return_value = {"ResponseMetadata": {"HTTPStatusCode": 500}}
# Act
result = cast(
dict[str, Any],
ExternalDatasetTestService.knowledge_retrieval(BedrockRetrievalSetting(top_k=1), "query", "kb"),
)
# Assert
assert result["records"] == []
def test_knowledge_retrieval_should_raise_when_boto_client_creation_fails(self):
"""Test that exceptions from boto3.client propagate (e.g., network/credentials issues)"""
with patch("services.knowledge_service.boto3.client") as mock_boto:
mock_boto.side_effect = Exception("client init failed")
with pytest.raises(Exception) as exc_info:
ExternalDatasetTestService.knowledge_retrieval(BedrockRetrievalSetting(top_k=1), "query", "kb")
assert "client init failed" in str(exc_info.value)
# ===== Edge Cases =====
@patch("services.knowledge_service.boto3.client")
def test_knowledge_retrieval_should_handle_missing_threshold_in_settings(self, mock_boto: MagicMock):
"""Test that knowledge_retrieval uses 0.0 as default threshold if not provided"""
# Arrange
mock_client = MagicMock()
mock_boto.return_value = mock_client
mock_client.retrieve.return_value = {
"ResponseMetadata": {"HTTPStatusCode": 200},
"retrievalResults": [
{
"score": 0.1,
"metadata": {"x-amz-bedrock-kb-source-uri": "uri"},
"content": {"text": "text"},
}
],
}
# Act
# retrieval_setting missing "score_threshold"
result = cast(
dict[str, Any],
ExternalDatasetTestService.knowledge_retrieval(BedrockRetrievalSetting(top_k=1), "query", "kb"),
)
# Assert
assert len(result["records"]) == 1
assert result["records"][0]["score"] == 0.1

View File

@ -70,7 +70,6 @@ export const contractLoaders = {
import('./system-features/orpc.gen').then(({ systemFeatures }) => ({ systemFeatures })),
tagBindings: () => import('./tag-bindings/orpc.gen').then(({ tagBindings }) => ({ tagBindings })),
tags: () => import('./tags/orpc.gen').then(({ tags }) => ({ tags })),
test: () => import('./test/orpc.gen').then(({ test }) => ({ test })),
trialApps: () => import('./trial-apps/orpc.gen').then(({ trialApps }) => ({ trialApps })),
trialModels: () => import('./trial-models/orpc.gen').then(({ trialModels }) => ({ trialModels })),
version: () => import('./version/orpc.gen').then(({ version }) => ({ version })),

View File

@ -46,7 +46,6 @@ import { spec } from './spec/orpc.gen'
import { systemFeatures } from './system-features/orpc.gen'
import { tagBindings } from './tag-bindings/orpc.gen'
import { tags } from './tags/orpc.gen'
import { test } from './test/orpc.gen'
import { trialApps } from './trial-apps/orpc.gen'
import { trialModels } from './trial-models/orpc.gen'
import { version } from './version/orpc.gen'
@ -102,7 +101,6 @@ const communityContract = {
systemFeatures,
tagBindings,
tags,
test,
trialApps,
trialModels,
version,

View File

@ -1,32 +0,0 @@
// This file is auto-generated by @hey-api/openapi-ts
import { oc } from '@orpc/contract'
import * as z from 'zod'
import { zPostTestRetrievalBody, zPostTestRetrievalResponse } from './zod.gen'
/**
* Bedrock retrieval test (internal use only)
*/
export const post = oc
.route({
description: 'Bedrock retrieval test (internal use only)',
inputStructure: 'detailed',
method: 'POST',
operationId: 'postTestRetrieval',
path: '/test/retrieval',
tags: ['console'],
})
.input(z.object({ body: zPostTestRetrievalBody }))
.output(zPostTestRetrievalResponse)
export const retrieval = {
post,
}
export const test = {
retrieval,
}
export const contract = {
test,
}

View File

@ -1,42 +0,0 @@
// This file is auto-generated by @hey-api/openapi-ts
export type ClientOptions = {
baseUrl: `${string}://${string}/console/api` | (string & {})
}
export type BedrockRetrievalPayload = {
knowledge_id: string
query: string
retrieval_setting: BedrockRetrievalSetting
}
export type BedrockRetrievalResponse = {
records: Array<BedrockRetrievalRecordResponse>
}
export type BedrockRetrievalSetting = {
score_threshold?: number
top_k?: number | null
}
export type BedrockRetrievalRecordResponse = {
content?: string | null
metadata?: {
[key: string]: unknown
} | null
score: number
title?: string | null
}
export type PostTestRetrievalData = {
body: BedrockRetrievalPayload
path?: never
query?: never
url: '/test/retrieval'
}
export type PostTestRetrievalResponses = {
200: BedrockRetrievalResponse
}
export type PostTestRetrievalResponse = PostTestRetrievalResponses[keyof PostTestRetrievalResponses]

View File

@ -1,46 +0,0 @@
// This file is auto-generated by @hey-api/openapi-ts
import * as z from 'zod'
/**
* BedrockRetrievalSetting
*
* Retrieval settings for Amazon Bedrock knowledge base queries.
*/
export const zBedrockRetrievalSetting = z.object({
score_threshold: z.number().optional().default(0),
top_k: z.int().nullish(),
})
/**
* BedrockRetrievalPayload
*/
export const zBedrockRetrievalPayload = z.object({
knowledge_id: z.string(),
query: z.string(),
retrieval_setting: zBedrockRetrievalSetting,
})
/**
* BedrockRetrievalRecordResponse
*/
export const zBedrockRetrievalRecordResponse = z.object({
content: z.string().nullish(),
metadata: z.record(z.string(), z.unknown()).nullish(),
score: z.number(),
title: z.string().nullish(),
})
/**
* BedrockRetrievalResponse
*/
export const zBedrockRetrievalResponse = z.object({
records: z.array(zBedrockRetrievalRecordResponse),
})
export const zPostTestRetrievalBody = zBedrockRetrievalPayload
/**
* Bedrock retrieval test completed
*/
export const zPostTestRetrievalResponse = zBedrockRetrievalResponse