Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from dbgpt_app.openapi.api_view_model import (
ConversationVo,
Result,
resolve_dialogue_user_name,
)
from dbgpt_serve.datasource.manages import ConnectorManager
from dbgpt_serve.utils.auth import UserRequest, get_user_from_headers
Expand Down Expand Up @@ -4069,7 +4070,9 @@ async def chat_react_agent(
dialogue.select_param,
dialogue.model_name,
)
dialogue.user_name = user_token.user_id if user_token else dialogue.user_name
dialogue.user_name = resolve_dialogue_user_name(
dialogue.user_name, user_token.user_id if user_token else None
)
headers = {
"Content-Type": "text/event-stream",
"Cache-Control": "no-cache",
Expand Down
11 changes: 8 additions & 3 deletions packages/dbgpt-app/src/dbgpt_app/openapi/api_v1/api_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@
ConversationVo,
MessageVo,
Result,
resolve_dialogue_user_name,
)
from dbgpt_app.scene import BaseChat, ChatFactory, ChatParam, ChatScene
from dbgpt_serve.agent.db.gpts_app import UserRecentAppsDao, adapt_native_app_model
Expand Down Expand Up @@ -512,15 +513,17 @@ async def chat_prepare(
):
logger.info(json.dumps(dialogue.__dict__))
# dialogue.model_name = CFG.LLM_MODEL
dialogue.user_name = user_token.user_id if user_token else dialogue.user_name
dialogue.user_name = resolve_dialogue_user_name(
dialogue.user_name, user_token.user_id if user_token else None
)
Comment on lines +516 to +518
logger.info(f"chat_prepare:{dialogue}")
## check conv_uid
chat: BaseChat = await get_chat_instance(dialogue)

await chat.prepare()

# Refresh messages
return Result.succ(get_hist_messages(dialogue.conv_uid, user_token.user_id))
return Result.succ(get_hist_messages(dialogue.conv_uid, dialogue.user_name))


@router.post("/v1/chat/completions")
Expand All @@ -533,7 +536,9 @@ async def chat_completions(
f"chat_completions:{dialogue.chat_mode},{dialogue.select_param},"
f"{dialogue.model_name}, timestamp={int(time.time() * 1000)}"
)
dialogue.user_name = user_token.user_id if user_token else dialogue.user_name
dialogue.user_name = resolve_dialogue_user_name(
dialogue.user_name, user_token.user_id if user_token else None
)
dialogue = adapt_native_app_model(dialogue)

# Handle knowledge space selection from ext_info for normal chat mode
Expand Down
7 changes: 7 additions & 0 deletions packages/dbgpt-app/src/dbgpt_app/openapi/api_view_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,13 @@
T = TypeVar("T")


def resolve_dialogue_user_name(
request_user_name: Optional[str], token_user_id: Optional[str]
) -> Optional[str]:
"""Prefer explicit dialogue user_name and fall back to the auth token user."""
return request_user_name or token_user_id


class Result(BaseModel, Generic[T]):
success: bool
err_code: Optional[str] = None
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
import pytest

from dbgpt_app.openapi.api_view_model import resolve_dialogue_user_name
from dbgpt_serve.utils.auth import UserRequest


def test_resolve_dialogue_user_name_preserves_explicit_request_user():
assert resolve_dialogue_user_name("request_user", "token_user") == "request_user"


def test_resolve_dialogue_user_name_falls_back_to_authenticated_user():
assert resolve_dialogue_user_name(None, "token_user") == "token_user"


@pytest.mark.asyncio
async def test_chat_prepare_refreshes_history_with_resolved_dialogue_user(monkeypatch):
from dbgpt_app.openapi.api_v1 import api_v1
from dbgpt_app.openapi.api_view_model import ConversationVo

class FakeChat:
async def prepare(self):
return None

captured = {}

async def fake_get_chat_instance(dialogue):
captured["chat_user_name"] = dialogue.user_name
return FakeChat()

def fake_get_hist_messages(conv_uid, user_name=None):
captured["history_conv_uid"] = conv_uid
captured["history_user_name"] = user_name
return ["history"]

monkeypatch.setattr(api_v1, "get_chat_instance", fake_get_chat_instance)
monkeypatch.setattr(api_v1, "get_hist_messages", fake_get_hist_messages)

result = await api_v1.chat_prepare(
ConversationVo(conv_uid="conv-1", user_name="request_user"),
UserRequest(user_id="token_user"),
)

assert result.data == ["history"]
assert captured == {
"chat_user_name": "request_user",
"history_conv_uid": "conv-1",
"history_user_name": "request_user",
}
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ def to_response(self, entity: ServeEntity) -> ServerResponse:
conv_uid=entity.conv_uid,
user_input=entity.summary,
chat_mode=entity.chat_mode,
user_name="",
user_name=entity.user_name,
sys_code=entity.sys_code,
gmt_created=gmt_created,
gmt_modified=gmt_modified,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,21 @@ def test_entity_create(default_entity_dict):
session.add(entity)


def test_to_response_preserves_user_name(dao):
entity = ServeEntity(
conv_uid="test_conv_uid",
summary="hello",
chat_mode="chat_normal",
user_name="request_user",
sys_code="dbgpt",
app_code="chat_normal",
)

response = dao.to_response(entity)

assert response.user_name == "request_user"


def test_entity_unique_key(default_entity_dict):
# TODO: implement your test case
pass
Expand Down
Loading