mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-25 18:03:29 +08:00
fix: prevent session user_id spoofing via request body (#15077)
### What problem does this PR solve? Closes #15076 Two endpoints in `api/apps/restful_apis/chat_api.py` accepted a `user_id` field from the request body and used it directly when creating a session: ```python # before (vulnerable) "user_id": req.get("user_id", current_user.id) # create_session conv = await _create_session_for_completion(chat_id, dia, req.get("user_id", current_user.id)) # session_completion ``` Any authenticated caller could supply an arbitrary `user_id` and have the new session attributed to a different user — effectively spoofing session ownership. Both call sites are now fixed to always use `current_user.id`, which is set by the authentication middleware and cannot be tampered with via the request payload. ### Type of change - [x] Bug Fix (non-breaking change which fixes an issue) ### Changes | File | Change | |------|--------| | `api/apps/restful_apis/chat_api.py` | Remove `req.get("user_id", ...)` fallback in `create_session` and `session_completion`; always use `current_user.id` | | `test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py` | Add `test_create_session_user_id_not_spoofable` and `test_session_completion_user_id_not_spoofable` (both `@pytest.mark.p2`) | ### Testing Two new unit tests assert that a `user_id` value supplied in the request body is silently ignored and the session is always owned by the authenticated user: ``` test_create_session_user_id_not_spoofable test_session_completion_user_id_not_spoofable ``` Run with: ```bash uv run pytest test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py -k "spoofable" -v ```
This commit is contained in:
@@ -662,6 +662,7 @@ async def bulk_delete_chats():
|
||||
@manager.route("/chats/<chat_id>/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):
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user