diff --git a/api/apps/restful_apis/chat_api.py b/api/apps/restful_apis/chat_api.py index baba98ae28..33fb74acf9 100644 --- a/api/apps/restful_apis/chat_api.py +++ b/api/apps/restful_apis/chat_api.py @@ -662,6 +662,7 @@ async def bulk_delete_chats(): @manager.route("/chats//sessions", methods=["POST"]) # noqa: F821 @login_required async def create_session(chat_id): + """Create a new conversation session for the given chat, owned by the authenticated user.""" if not await _ensure_owned_chat(chat_id): return get_json_result(data=False, message="No authorization.", code=RetCode.AUTHENTICATION_ERROR) try: @@ -678,7 +679,7 @@ async def create_session(chat_id): "dialog_id": chat_id, "name": name, "message": [{"role": "assistant", "content": dia.prompt_config.get("prologue", "")}], - "user_id": req.get("user_id", current_user.id), + "user_id": current_user.id, "reference": [], } ConversationService.save(**conv) @@ -1058,6 +1059,7 @@ async def recommendation(): @login_required @validate_request("messages") async def session_completion(chat_id_in_arg=""): + """Handle chat completion requests, streaming or non-streaming, scoped to the authenticated user.""" req = await get_request_json() msg = [] for m in req["messages"]: @@ -1100,7 +1102,7 @@ async def session_completion(chat_id_in_arg=""): if conv.dialog_id != chat_id: return get_data_error_result(message="Session does not belong to this chat!") else: - conv = await _create_session_for_completion(chat_id, dia, req.get("user_id", current_user.id)) + conv = await _create_session_for_completion(chat_id, dia, current_user.id) session_id = conv.id conv.message = deepcopy(req["messages"]) else: @@ -1124,12 +1126,14 @@ async def session_completion(chat_id_in_arg=""): stream_mode = req.pop("stream", True) def _format_answer(ans): + """Wrap a raw answer dict with session and chat identifiers.""" formatted = structure_answer(conv, ans, message_id, session_id) if chat_id: formatted["chat_id"] = chat_id return formatted async def stream(): + """Yield SSE-formatted chunks from the async chat generator.""" nonlocal dia, msg, req, conv try: async for ans in async_chat(dia, msg, True, **req): diff --git a/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py b/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py index 889de4ba1f..33f977a4c2 100644 --- a/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py +++ b/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py @@ -646,6 +646,7 @@ def _load_session_module(monkeypatch): dialog_service_mod = ModuleType("api.db.services.dialog_service") dialog_service_mod.DialogService = SimpleNamespace( + model=SimpleNamespace(_meta=SimpleNamespace(fields=[])), query=lambda **_kwargs: [], get_by_id=lambda *_args, **_kwargs: (False, None), ) @@ -1998,3 +1999,236 @@ def test_build_reference_chunks_metadata_matrix_unit(monkeypatch): assert res[1]["document_metadata"] == {"author": "bob"} assert "document_metadata" not in res[2] assert "document_metadata" not in res[3] + + +# --------------------------------------------------------------------------- +# chat_api unit tests — session user-id spoofing fix +# --------------------------------------------------------------------------- + +def _load_chat_api_module(monkeypatch): + """Load api/apps/restful_apis/chat_api.py with all heavy dependencies mocked.""" + repo_root = Path(__file__).resolve().parents[4] + + from enum import Enum, StrEnum + + class _RetCode(int, Enum): + SUCCESS = 0 + DATA_ERROR = 102 + AUTHENTICATION_ERROR = 109 + SERVER_ERROR = 500 + + class _StatusEnum(str, Enum): + VALID = "1" + + class _LLMType(StrEnum): + CHAT = "chat" + TTS = "tts" + SPEECH2TEXT = "speech2text" + RERANK = "rerank" + + quart_mod = ModuleType("quart") + quart_mod.request = SimpleNamespace(args=_Args(), headers={}, files=_AwaitableValue({}), method="POST") + quart_mod.Response = _StubResponse + quart_mod.jsonify = lambda payload: payload + quart_mod.current_app = SimpleNamespace() + monkeypatch.setitem(sys.modules, "quart", quart_mod) + + api_apps_mod = ModuleType("api.apps") + api_apps_mod.__path__ = [str(repo_root / "api" / "apps")] + api_apps_mod.current_user = SimpleNamespace(id="authenticated-user") + api_apps_mod.login_required = lambda func: func + monkeypatch.setitem(sys.modules, "api.apps", api_apps_mod) + + common_pkg = ModuleType("common") + common_pkg.__path__ = [str(repo_root / "common")] + monkeypatch.setitem(sys.modules, "common", common_pkg) + + common_constants_mod = ModuleType("common.constants") + common_constants_mod.LLMType = _LLMType + common_constants_mod.RetCode = _RetCode + common_constants_mod.StatusEnum = _StatusEnum + monkeypatch.setitem(sys.modules, "common.constants", common_constants_mod) + + common_settings_mod = ModuleType("common.settings") + common_settings_mod.STORAGE_IMPL = SimpleNamespace(rm=lambda *_a, **_k: None) + common_pkg.settings = common_settings_mod + monkeypatch.setitem(sys.modules, "common.settings", common_settings_mod) + + async def _thread_pool_exec(func, *args, **kwargs): + """Run func synchronously, standing in for the real async thread-pool executor.""" + return func(*args, **kwargs) + + common_misc_mod = ModuleType("common.misc_utils") + common_misc_mod.get_uuid = lambda: "test-uuid" + common_misc_mod.thread_pool_exec = _thread_pool_exec + monkeypatch.setitem(sys.modules, "common.misc_utils", common_misc_mod) + + joint_pkg = ModuleType("api.db.joint_services") + joint_pkg.__path__ = [] + monkeypatch.setitem(sys.modules, "api.db.joint_services", joint_pkg) + + tenant_model_svc = ModuleType("api.db.joint_services.tenant_model_service") + tenant_model_svc.get_model_config_by_type_and_name = lambda *_a, **_k: {} + tenant_model_svc.get_tenant_default_model_by_type = lambda *_a, **_k: {} + monkeypatch.setitem(sys.modules, "api.db.joint_services.tenant_model_service", tenant_model_svc) + + chunk_feedback_mod = ModuleType("api.db.services.chunk_feedback_service") + chunk_feedback_mod.ChunkFeedbackService = SimpleNamespace() + monkeypatch.setitem(sys.modules, "api.db.services.chunk_feedback_service", chunk_feedback_mod) + + def _make_conv(conv_id): + """Return a minimal conversation stub for ConversationService.get_by_id mocks.""" + return SimpleNamespace( + id=conv_id, + dialog_id="chat-1", + message=[], + reference=[], + user_id="authenticated-user", + name="test", + to_dict=lambda: { + "id": conv_id, + "dialog_id": "chat-1", + "message": [], + "reference": [], + "user_id": "authenticated-user", + "name": "test", + }, + ) + + conv_svc_mod = ModuleType("api.db.services.conversation_service") + conv_svc_mod.ConversationService = SimpleNamespace( + save=lambda **_k: True, + get_by_id=lambda _id: (True, _make_conv(_id)), + query=lambda **_k: [], + get_list=lambda *_a, **_k: [], + ) + conv_svc_mod.structure_answer = lambda *_a, **_k: {} + monkeypatch.setitem(sys.modules, "api.db.services.conversation_service", conv_svc_mod) + + dialog_svc_mod = ModuleType("api.db.services.dialog_service") + dialog_svc_mod.DialogService = SimpleNamespace( + model=SimpleNamespace(_meta=SimpleNamespace(fields=[])), + query=lambda **_k: [SimpleNamespace(id="chat-1", icon="")], + get_by_id=lambda _id: (True, SimpleNamespace( + prompt_config={"prologue": ""}, + tenant_id="tenant-1", + llm_id="model", + kb_ids=[], + id=_id, + )), + ) + dialog_svc_mod.async_chat = lambda *_a, **_k: None + dialog_svc_mod.gen_mindmap = lambda *_a, **_k: None + monkeypatch.setitem(sys.modules, "api.db.services.dialog_service", dialog_svc_mod) + + kb_svc_mod = ModuleType("api.db.services.knowledgebase_service") + kb_svc_mod.KnowledgebaseService = SimpleNamespace(query=lambda **_k: [], accessible=lambda **_k: True) + monkeypatch.setitem(sys.modules, "api.db.services.knowledgebase_service", kb_svc_mod) + + class _FakeLLMBundle: + def __init__(self, *_a, **_k): + pass + + llm_svc_mod = ModuleType("api.db.services.llm_service") + llm_svc_mod.LLMBundle = _FakeLLMBundle + monkeypatch.setitem(sys.modules, "api.db.services.llm_service", llm_svc_mod) + + search_svc_mod = ModuleType("api.db.services.search_service") + search_svc_mod.SearchService = SimpleNamespace() + monkeypatch.setitem(sys.modules, "api.db.services.search_service", search_svc_mod) + + tenant_llm_svc_mod = ModuleType("api.db.services.tenant_llm_service") + tenant_llm_svc_mod.TenantLLMService = SimpleNamespace(get_api_key=lambda **_k: None) + monkeypatch.setitem(sys.modules, "api.db.services.tenant_llm_service", tenant_llm_svc_mod) + + user_svc_mod = ModuleType("api.db.services.user_service") + user_svc_mod.TenantService = SimpleNamespace( + get_by_id=lambda _id: (True, SimpleNamespace(id=_id)), + get_joined_tenants_by_user_id=lambda _id: [], + ) + user_svc_mod.UserTenantService = SimpleNamespace(query=lambda **_k: []) + monkeypatch.setitem(sys.modules, "api.db.services.user_service", user_svc_mod) + + api_utils_mod = ModuleType("api.utils.api_utils") + api_utils_mod.check_duplicate_ids = lambda ids, _kind: (ids, []) + api_utils_mod.get_data_error_result = lambda message="Error", code=_RetCode.DATA_ERROR: {"code": code, "message": message} + api_utils_mod.get_json_result = lambda code=_RetCode.SUCCESS, message="success", data=None: {"code": code, "message": message, "data": data} + api_utils_mod.get_request_json = lambda: _AwaitableValue({}) + api_utils_mod.server_error_response = lambda e: {"code": _RetCode.SERVER_ERROR, "message": str(e)} + api_utils_mod.validate_request = lambda *_a, **_k: (lambda func: func) + monkeypatch.setitem(sys.modules, "api.utils.api_utils", api_utils_mod) + + tenant_utils_mod = ModuleType("api.utils.tenant_utils") + tenant_utils_mod.ensure_tenant_model_id_for_params = lambda _tenant_id, req: req + monkeypatch.setitem(sys.modules, "api.utils.tenant_utils", tenant_utils_mod) + + rag_gen_mod = ModuleType("rag.prompts.generator") + rag_gen_mod.chunks_format = lambda chunks: chunks + monkeypatch.setitem(sys.modules, "rag.prompts.generator", rag_gen_mod) + + rag_tmpl_mod = ModuleType("rag.prompts.template") + rag_tmpl_mod.load_prompt = lambda *_a, **_k: "" + monkeypatch.setitem(sys.modules, "rag.prompts.template", rag_tmpl_mod) + + module_path = repo_root / "api" / "apps" / "restful_apis" / "chat_api.py" + spec = importlib.util.spec_from_file_location("test_chat_api_unit_module", module_path) + module = importlib.util.module_from_spec(spec) + module.manager = _DummyManager() + monkeypatch.setitem(sys.modules, "test_chat_api_unit_module", module) + spec.loader.exec_module(module) + return module + + +@pytest.mark.p2 +def test_create_session_user_id_not_spoofable(monkeypatch): + """create_session must attribute the new session to current_user, ignoring any user_id in the request body.""" + module = _load_chat_api_module(monkeypatch) + + saved_kwargs = {} + + def _capture_save(**kwargs): + """Capture ConversationService.save arguments for later assertion.""" + saved_kwargs.update(kwargs) + return True + + monkeypatch.setattr(module.ConversationService, "save", _capture_save) + monkeypatch.setattr( + module, + "get_request_json", + lambda: _AwaitableValue({"name": "my session", "user_id": "attacker-id"}), + ) + + res = _run(inspect.unwrap(module.create_session)("chat-1")) + + assert res["code"] == 0, res + assert saved_kwargs["user_id"] == module.current_user.id + assert saved_kwargs["user_id"] != "attacker-id" + + +@pytest.mark.p2 +def test_session_completion_user_id_not_spoofable(monkeypatch): + """session_completion must pass current_user.id to _create_session_for_completion, ignoring user_id in the request.""" + module = _load_chat_api_module(monkeypatch) + + captured_user_ids = [] + + async def _spy_create_session(_chat_id, _dialog, user_id): + """Record the user_id passed to _create_session_for_completion, then abort early.""" + captured_user_ids.append(user_id) + raise RuntimeError("stop-sentinel") + + monkeypatch.setattr(module, "_create_session_for_completion", _spy_create_session) + monkeypatch.setattr( + module, + "get_request_json", + lambda: _AwaitableValue({ + "messages": [{"role": "user", "content": "hello"}], + "chat_id": "chat-1", + "user_id": "attacker-id", + "stream": False, + }), + ) + + _run(inspect.unwrap(module.session_completion)()) + + assert captured_user_ids == [module.current_user.id]