mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-09 00:47:59 +08:00
fix: reject duplicate chat parameter keys (#17977)
This commit is contained in:
@@ -161,6 +161,24 @@ def _validate_name(name, *, required=True):
|
||||
return name, None
|
||||
|
||||
|
||||
def _validate_prompt_parameters(prompt_config):
|
||||
parameters = prompt_config.get("parameters")
|
||||
if not isinstance(parameters, list):
|
||||
return None
|
||||
|
||||
keys = []
|
||||
for parameter in parameters:
|
||||
if not isinstance(parameter, dict):
|
||||
continue
|
||||
key = parameter.get("key")
|
||||
if key is None:
|
||||
continue
|
||||
if key in keys:
|
||||
return f"`parameters` contains duplicate key: {key}"
|
||||
keys.append(key)
|
||||
return None
|
||||
|
||||
|
||||
def _build_session_response(conv: dict) -> dict:
|
||||
conv = dict(conv)
|
||||
conv["chat_id"] = conv.pop("dialog_id", conv.get("chat_id"))
|
||||
@@ -428,6 +446,9 @@ async def create():
|
||||
if "prompt_config" in req:
|
||||
if not isinstance(req["prompt_config"], dict):
|
||||
return get_data_error_result(message="`prompt_config` should be an object.")
|
||||
err = _validate_prompt_parameters(req["prompt_config"])
|
||||
if err:
|
||||
return get_data_error_result(message=err)
|
||||
# err = _validate_prompt_config(req["prompt_config"])
|
||||
# if err:
|
||||
# return get_data_error_result(message=err)
|
||||
@@ -602,6 +623,9 @@ async def update_chat(chat_id):
|
||||
if "prompt_config" in req:
|
||||
if not isinstance(req["prompt_config"], dict):
|
||||
return get_data_error_result(message="`prompt_config` should be an object.")
|
||||
err = _validate_prompt_parameters(req["prompt_config"])
|
||||
if err:
|
||||
return get_data_error_result(message=err)
|
||||
# err = _validate_prompt_config(req["prompt_config"])
|
||||
# if err:
|
||||
# return get_data_error_result(message=err)
|
||||
@@ -687,6 +711,9 @@ async def patch_chat(chat_id):
|
||||
if "prompt_config" in req:
|
||||
if not isinstance(req["prompt_config"], dict):
|
||||
return get_data_error_result(message="`prompt_config` should be an object.")
|
||||
err = _validate_prompt_parameters(req["prompt_config"])
|
||||
if err:
|
||||
return get_data_error_result(message=err)
|
||||
prompt_config = deepcopy(current_chat.get("prompt_config", {}))
|
||||
prompt_config.update(req["prompt_config"])
|
||||
req["prompt_config"] = prompt_config
|
||||
|
||||
@@ -256,9 +256,13 @@ func (s *ChatService) Create(ctx context.Context, userID string, req map[string]
|
||||
}
|
||||
|
||||
if promptConfigValue, ok := req["prompt_config"]; ok {
|
||||
if _, ok := mapFromValue(promptConfigValue); !ok {
|
||||
promptConfig, ok := mapFromValue(promptConfigValue)
|
||||
if !ok {
|
||||
return nil, common.CodeDataError, errors.New("`prompt_config` should be an object")
|
||||
}
|
||||
if err := validatePromptConfigParameters(promptConfig); err != nil {
|
||||
return nil, common.CodeDataError, err
|
||||
}
|
||||
}
|
||||
|
||||
if metaDataFilterValue, ok := req["meta_data_filter"]; ok && metaDataFilterValue != nil {
|
||||
@@ -915,6 +919,9 @@ func (s *ChatService) updateChatREST(ctx context.Context, userID, chatID string,
|
||||
if !ok {
|
||||
return nil, errors.New("`prompt_config` should be an object")
|
||||
}
|
||||
if err := validatePromptConfigParameters(promptConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if patch {
|
||||
req["prompt_config"] = mergeJSONMap(currentChat.PromptConfig, promptConfig)
|
||||
} else {
|
||||
@@ -980,6 +987,30 @@ func (s *ChatService) updateChatREST(ctx context.Context, userID, chatID string,
|
||||
return s.buildRESTChatResponse(ctx, updatedChat), nil
|
||||
}
|
||||
|
||||
func validatePromptConfigParameters(promptConfig map[string]interface{}) error {
|
||||
parameters, ok := promptConfig["parameters"].([]interface{})
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
seen := make(map[string]struct{}, len(parameters))
|
||||
for _, value := range parameters {
|
||||
parameter, ok := mapFromValue(value)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
key, ok := parameter["key"].(string)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[key]; exists {
|
||||
return fmt.Errorf("`parameters` contains duplicate key: %s", key)
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateRESTChatName(value interface{}, required bool) (string, bool, error) {
|
||||
if value == nil {
|
||||
if required {
|
||||
|
||||
@@ -489,6 +489,24 @@ func TestChatServiceCreateRejectsInvalidPromptConfig(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatServiceUpdateRejectsDuplicatePromptParameterKeys(t *testing.T) {
|
||||
db := setupChatRESTUpdateServiceTestDB(t)
|
||||
createChatRESTUpdateServiceTestChat(t, db, "chat-1", "user-1")
|
||||
|
||||
svc := NewChatService()
|
||||
_, err := svc.UpdateChat(t.Context(), "user-1", "chat-1", map[string]interface{}{
|
||||
"prompt_config": map[string]interface{}{
|
||||
"parameters": []interface{}{
|
||||
map[string]interface{}{"key": "knowledge"},
|
||||
map[string]interface{}{"key": "knowledge"},
|
||||
},
|
||||
},
|
||||
})
|
||||
if err == nil || err.Error() != "`parameters` contains duplicate key: knowledge" {
|
||||
t.Fatalf("expected duplicate parameter key error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatServiceCreatePromptDefaultsContract(t *testing.T) {
|
||||
setupChatRESTUpdateServiceTestDB(t)
|
||||
|
||||
|
||||
@@ -1957,6 +1957,42 @@ def test_chat_update_mapping_and_validation_branches_p2(rest_client, clear_chats
|
||||
target_payload = target_res.json()
|
||||
assert target_payload["code"] == 0, target_payload
|
||||
chat_id = target_payload["data"]["id"]
|
||||
original_parameters = target_payload["data"]["prompt_config"]["parameters"]
|
||||
|
||||
duplicate_create_res = rest_client.post(
|
||||
"/chats",
|
||||
json={
|
||||
"name": "restful_chat_update_mapping_duplicate_parameters",
|
||||
"dataset_ids": [],
|
||||
"prompt_config": {"parameters": [{"key": "knowledge"}, {"key": "knowledge"}]},
|
||||
},
|
||||
)
|
||||
assert duplicate_create_res.status_code == 200
|
||||
duplicate_create_payload = duplicate_create_res.json()
|
||||
assert duplicate_create_payload["code"] == 102, duplicate_create_payload
|
||||
assert duplicate_create_payload["message"] == "`parameters` contains duplicate key: knowledge", duplicate_create_payload
|
||||
|
||||
duplicate_parameters_res = rest_client.put(
|
||||
f"/chats/{chat_id}",
|
||||
json={"prompt_config": {"parameters": [{"key": "knowledge"}, {"key": "knowledge"}]}},
|
||||
)
|
||||
assert duplicate_parameters_res.status_code == 200
|
||||
duplicate_parameters_payload = duplicate_parameters_res.json()
|
||||
assert duplicate_parameters_payload["code"] == 102, duplicate_parameters_payload
|
||||
assert duplicate_parameters_payload["message"] == "`parameters` contains duplicate key: knowledge", duplicate_parameters_payload
|
||||
get_after_put_res = rest_client.get(f"/chats/{chat_id}")
|
||||
assert get_after_put_res.json()["data"]["prompt_config"]["parameters"] == original_parameters
|
||||
|
||||
duplicate_parameters_patch_res = rest_client.patch(
|
||||
f"/chats/{chat_id}",
|
||||
json={"prompt_config": {"parameters": [{"key": "knowledge"}, {"key": "knowledge"}]}},
|
||||
)
|
||||
assert duplicate_parameters_patch_res.status_code == 200
|
||||
duplicate_parameters_patch_payload = duplicate_parameters_patch_res.json()
|
||||
assert duplicate_parameters_patch_payload["code"] == 102, duplicate_parameters_patch_payload
|
||||
assert duplicate_parameters_patch_payload["message"] == "`parameters` contains duplicate key: knowledge", duplicate_parameters_patch_payload
|
||||
get_after_patch_res = rest_client.get(f"/chats/{chat_id}")
|
||||
assert get_after_patch_res.json()["data"]["prompt_config"]["parameters"] == original_parameters
|
||||
|
||||
unauthorized = rest_client.patch("/chats/invalid-chat-id", json={"name": "anything"})
|
||||
assert unauthorized.status_code == 200
|
||||
|
||||
Reference in New Issue
Block a user