fix(mark): prevent invalid assignees from breaking task list

This commit is contained in:
GuoQing Zhang
2026-07-16 20:28:21 +08:00
parent 57a7585c34
commit bbbaeda973
5 changed files with 176 additions and 51 deletions
+27 -18
View File
@@ -1,5 +1,4 @@
from collections import deque
from typing import Optional
from fastapi import APIRouter, Depends, Request
from loguru import logger
@@ -7,7 +6,6 @@ from loguru import logger
from bisheng.api.v1.schema.mark_schema import MarkData, MarkTaskCreate
from bisheng.api.v1.schemas import resp_200, resp_500
from bisheng.common.dependencies.user_deps import UserPayload
from bisheng.database.models.mark_app_user import MarkAppUser, MarkAppUserDao
from bisheng.database.models.mark_record import MarkRecord, MarkRecordDao
from bisheng.database.models.mark_task import MarkTask, MarkTaskDao, MarkTaskRead, MarkTaskStatus
from bisheng.database.models.message import ChatMessageDao
@@ -19,9 +17,22 @@ from bisheng.utils.linked_list import DoubleLinkList
router = APIRouter(prefix='/mark', tags=['Mark'])
def _parse_process_user_ids(process_users: str, task_id: int) -> list[int]:
user_ids = []
for raw_user_id in (process_users or '').split(','):
raw_user_id = raw_user_id.strip()
if not raw_user_id:
continue
try:
user_ids.append(int(raw_user_id))
except ValueError:
logger.warning('Skipping invalid mark task process user: task_id={} user_id={}', task_id, raw_user_id)
return user_ids
@router.get('/list')
def list(request: Request,
status: Optional[int] = None,
status: int | None = None,
page_size: int = 10,
page_num: int = 1,
login_user: UserPayload = Depends(UserPayload.get_login_user)):
@@ -43,13 +54,16 @@ def list(request: Request,
process_list = []
user_count = {}
for c in task.process_users.split(","):
user = UserDao.get_user(int(c))
process_count = "{}:{}".format(user.user_name, 0)
user_count[int(c)] = process_count
for user_id in _parse_process_user_ids(task.process_users, task.id):
user = UserDao.get_user(user_id)
if not user:
logger.warning('Skipping missing mark task process user: task_id={} user_id={}', task.id, user_id)
continue
process_count = f"{user.user_name}:0"
user_count[user_id] = process_count
for c in record:
process_count = "{}:{}".format(c.create_user, c.user_count)
process_count = f"{c.create_user}:{c.user_count}"
user_count[c.create_id] = process_count
for c in user_count:
@@ -85,14 +99,9 @@ async def create(task_create: MarkTaskCreate, login_user: UserPayload = Depends(
task = MarkTask(create_id=login_user.user_id,
create_user=login_user.user_name,
app_id=",".join(task_create.app_list),
process_users=",".join(task_create.user_list)
process_users=",".join(str(user_id) for user_id in task_create.user_list)
)
MarkTaskDao.create_task(task)
user_app = [MarkAppUser(task_id=task.id, create_id=login_user.user_id, app_id=app, user_id=user) for app in
task_create.app_list for user in task_create.user_list]
MarkAppUserDao.create_task(user_app)
MarkTaskDao.create_task_with_assignments(task, task_create.app_list, task_create.user_list)
return resp_200(data="ok")
@@ -110,8 +119,8 @@ async def get_user(task_id: int):
task = MarkTaskDao.get_task_byid(task_id)
user_list = []
for u in task.process_users.split(","):
user = UserDao.get_user(int(u))
for user_id in _parse_process_user_ids(task.process_users, task.id):
user = UserDao.get_user(user_id)
if not user:
continue
user_list.append({"user_id": user.user_id, "user_name": user.user_name})
@@ -185,7 +194,7 @@ async def get_record(chat_id: str, task_id: int):
async def pre_or_next(chat_id: str, action: str, task_id: int,
login_user: UserPayload = Depends(UserPayload.get_login_user)):
"""
prev or next
prev or next
"""
if action not in ["prev", "next"]:
@@ -1,26 +1,13 @@
from typing import List, Optional, Any
from pydantic import BaseModel, Field, field_validator
from pydantic import BaseModel, Field
class MarkTaskCreate(BaseModel):
app_list: List[str] = Field(max_length=30)
user_list: List[str]
@field_validator('user_list', mode='before')
@classmethod
def convert_user_list(cls, v: Any):
ret = []
for one in v:
if isinstance(one, str):
ret.append(one)
else:
ret.append(str(one))
return ret
app_list: list[str] = Field(max_length=30)
user_list: list[int]
class MarkData(BaseModel):
session_id: str
task_id: int
status: int
flow_type: Optional[int] = None
flow_type: int | None = None
@@ -1,17 +1,14 @@
from datetime import datetime
from enum import Enum
from typing import List, Optional
# if TYPE_CHECKING:
from sqlalchemy import Column, DateTime, Integer, and_, delete, func, or_, text
from sqlmodel import Field, select, update
from bisheng.common.models.base import SQLModelSerializable
from bisheng.core.database import get_sync_db_session
from bisheng.core.database.dialect_helpers import UPDATE_TIME_SERVER_DEFAULT
class MarkTaskStatus(Enum):
DEFAULT = 1
DONE = 2
@@ -23,26 +20,26 @@ class MarkTaskBase(SQLModelSerializable):
create_id: int = Field(index=True)
app_id: str = Field(index=False, max_length=2048)
process_users: str = Field(index=False) # 23,2323
mark_user: Optional[str] = Field(default=None, index=True, nullable=True)
status: Optional[int] = Field(index=False, default=1)
tenant_id: Optional[int] = Field(
mark_user: str | None = Field(default=None, index=True, nullable=True)
status: int | None = Field(index=False, default=1)
tenant_id: int | None = Field(
default=None,
sa_column=Column(Integer, nullable=False, server_default=text('1'),
index=True, comment='Tenant ID'),
)
create_time: Optional[datetime] = Field(default=None, sa_column=Column(
create_time: datetime | None = Field(default=None, sa_column=Column(
DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP')))
update_time: Optional[datetime] = Field(default=None, sa_column=Column(
update_time: datetime | None = Field(default=None, sa_column=Column(
DateTime, nullable=False, server_default=UPDATE_TIME_SERVER_DEFAULT))
class MarkTask(MarkTaskBase, table=True):
id: Optional[int] = Field(default=None, primary_key=True)
id: int | None = Field(default=None, primary_key=True)
class MarkTaskRead(MarkTaskBase):
id: Optional[int] = None
mark_process: Optional[List[str]] = None
id: int | None = None
mark_process: list[str] | None = None
class MarkTaskDao(MarkTaskBase):
@@ -55,6 +52,33 @@ class MarkTaskDao(MarkTaskBase):
session.refresh(task_info)
return task_info
@classmethod
def create_task_with_assignments(
cls,
task_info: MarkTask,
app_ids: list[str],
user_ids: list[int],
) -> MarkTask:
from bisheng.database.models.mark_app_user import MarkAppUser
with get_sync_db_session() as session:
session.add(task_info)
session.flush()
assignments = [
MarkAppUser(
task_id=task_info.id,
create_id=task_info.create_id,
app_id=app_id,
user_id=user_id,
)
for app_id in app_ids
for user_id in user_ids
]
session.add_all(assignments)
session.commit()
session.refresh(task_info)
return task_info
@classmethod
def delete_task(cls, task_id: int):
with get_sync_db_session() as session:
@@ -71,7 +95,7 @@ class MarkTaskDao(MarkTaskBase):
@classmethod
def get_task(cls, user_id: int) -> MarkTask:
with get_sync_db_session() as session:
statement = select(MarkTask).where(MarkTask.process_users.like('%{}%'.format(user_id)))
statement = select(MarkTask).where(MarkTask.process_users.like(f'%{user_id}%'))
return session.exec(statement).first()
@classmethod
@@ -105,8 +129,8 @@ class MarkTaskDao(MarkTaskBase):
def get_task_list(
cls,
status: int,
create_id: Optional[int],
user_id: Optional[int],
create_id: int | None,
user_id: int | None,
page_size: int = 10,
page_num: int = 1,
):
@@ -0,0 +1,105 @@
from contextlib import contextmanager
from unittest.mock import Mock
import pytest
from pydantic import ValidationError
from bisheng.api.v1 import mark_task as mark_task_api
from bisheng.api.v1.schema.mark_schema import MarkTaskCreate
from bisheng.database.models import mark_task as mark_task_model
from bisheng.database.models.mark_task import MarkTask, MarkTaskDao
def test_mark_task_create_rejects_non_numeric_user_ids():
with pytest.raises(ValidationError):
MarkTaskCreate(app_list=['api-contract'], user_list=['api-contract'])
def test_mark_task_create_accepts_numeric_user_ids():
payload = MarkTaskCreate(app_list=['flow-id'], user_list=['686', 687])
assert payload.user_list == [686, 687]
def test_parse_process_user_ids_skips_legacy_invalid_values():
assert mark_task_api._parse_process_user_ids('api-contract,686,,invalid,687', task_id=111) == [686, 687]
def test_list_mark_tasks_keeps_rows_with_legacy_invalid_user_ids(monkeypatch):
task = MarkTask(
id=111,
create_id=1,
create_user='contract-probe',
app_id='api-contract',
process_users='api-contract',
)
login_user = Mock(user_id=3)
login_user.is_admin.return_value = True
monkeypatch.setattr(mark_task_api.UserGroupDao, 'get_user_admin_group', Mock(return_value=[]))
monkeypatch.setattr(MarkTaskDao, 'get_task_list', Mock(return_value=([task], 1)))
monkeypatch.setattr(mark_task_api.MarkRecordDao, 'get_count', Mock(return_value=[]))
get_user = Mock()
monkeypatch.setattr(mark_task_api.UserDao, 'get_user', get_user)
response = mark_task_api.list(request=Mock(), login_user=login_user)
assert response.status_code == 200
assert response.data['total'] == 1
assert response.data['list'][0].id == 111
assert response.data['list'][0].mark_process == []
get_user.assert_not_called()
@pytest.mark.asyncio
async def test_create_mark_task_uses_atomic_assignment_write(monkeypatch):
create_task = Mock()
monkeypatch.setattr(MarkTaskDao, 'create_task_with_assignments', create_task)
login_user = Mock(user_id=3, user_name='admin')
payload = MarkTaskCreate(app_list=['flow-id'], user_list=[686])
response = await mark_task_api.create(payload, login_user)
task, app_ids, user_ids = create_task.call_args.args
assert task.process_users == '686'
assert app_ids == ['flow-id']
assert user_ids == [686]
assert response.status_code == 200
def test_task_and_assignments_share_one_transaction(monkeypatch):
events = []
assignments = []
class FakeSession:
def add(self, task):
events.append('add_task')
self.task = task
def flush(self):
events.append('flush_task')
self.task.id = 42
def add_all(self, values):
events.append('add_assignments')
assignments.extend(values)
def commit(self):
events.append('commit')
def refresh(self, _task):
events.append('refresh_task')
@contextmanager
def fake_session():
yield FakeSession()
monkeypatch.setattr(mark_task_model, 'get_sync_db_session', fake_session)
task = MarkTask(create_id=3, create_user='admin', app_id='flow-a,flow-b', process_users='686')
MarkTaskDao.create_task_with_assignments(task, ['flow-a', 'flow-b'], [686])
assert events == ['add_task', 'flush_task', 'add_assignments', 'commit', 'refresh_task']
assert [(item.task_id, item.app_id, item.user_id) for item in assignments] == [
(42, 'flow-a', 686),
(42, 'flow-b', 686),
]
@@ -155,7 +155,7 @@ export async function getMarksApi({ status, pageSize, page }): Promise<{}> {
}
// 创建标注任务
export async function createMarkApi(data: { app_list: string[], user_list: string[] }) {
export async function createMarkApi(data: { app_list: string[], user_list: number[] }) {
return await axios.post('/api/v1/mark/create_task', data)
}