mirror of
https://github.com/dataelement/bisheng.git
synced 2026-08-29 01:22:31 +08:00
fix(mark): prevent invalid assignees from breaking task list
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user