Skip to content
Merged
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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[tool.poetry]
name = "whiskerrag"
version = "0.3.5"
version = "0.3.6"
description = "A utlity package for RAG operations"
authors = ["petercat.ai <antd.antgroup@gmail.com>"]
readme = "README.md"
Expand Down
22 changes: 8 additions & 14 deletions src/whiskerrag_types/model/page.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,25 +31,19 @@ class FilterGroup(BaseModel):
TAGGING_ALLOWED_FIELDS = {"tag_name", "tag_id"}


class TagFilter(BaseModel):
class TagFilter(FilterGroup):
"""
针对 Tagging 表的过滤条件。
注意:object_id / object_type 不作为用户输入过滤字段。
针对 Tagging 表的过滤条件(只允许使用 tag_name / tag_id 作为字段)
"""

advanced_filter: Optional[FilterGroup] = Field(
default=None, description="标签过滤条件,只允许 tag_name 和 tag_id"
)

@model_validator(mode="after")
def validate_tag_fields(self) -> "TagFilter":
if self.advanced_filter:
invalid_fields = self._validate_tag_filter_group(self.advanced_filter)
if invalid_fields:
raise ValueError(
f"Invalid tag_filter fields: {invalid_fields}; "
f"only {TAGGING_ALLOWED_FIELDS} are supported"
)
invalid_fields = self._validate_tag_filter_group(self)
if invalid_fields:
raise ValueError(
f"Invalid tag_filter fields: {invalid_fields}; "
f"only {TAGGING_ALLOWED_FIELDS} are supported"
)
return self

def _validate_tag_filter_group(self, filter_group: FilterGroup) -> set[str]:
Expand Down
74 changes: 74 additions & 0 deletions tests/whiskerrag_types/test_page.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,3 +142,77 @@ def test_alias_field_validation_fixed(self):
with pytest.raises(ValueError) as exc_info:
PageQueryParams[DummyModel](eq_conditions={"invalid_field": "value"})
assert "Invalid fields found: {'invalid_field'}" in str(exc_info.value)

def test_tag_filter_valid(self):
"""标签过滤器:合法字段"""
from whiskerrag_types.model.page import Condition, Operator, TagFilter

# 直接使用允许的字段
tf = TagFilter(
operator=Operator.AND,
conditions=[
Condition(field="tag_name", operator="eq", value="科技"),
Condition(field="tag_id", operator="neq", value="123"),
],
)
assert isinstance(tf, TagFilter)
assert tf.conditions[0].field == "tag_name"
assert tf.conditions[1].field == "tag_id"

def test_tag_filter_invalid_field(self):
"""标签过滤器:非法字段应报错"""
import pytest

from whiskerrag_types.model.page import Condition, Operator, TagFilter

with pytest.raises(ValueError) as exc_info:
TagFilter(
operator=Operator.AND,
conditions=[
Condition(field="not_allowed_field", operator="eq", value="x")
],
)
assert "only {'tag_name', 'tag_id'} are supported" in str(exc_info.value)

def test_tag_filter_nested_groups(self):
"""标签过滤器:嵌套组的校验"""
import pytest

from whiskerrag_types.model.page import (
Condition,
FilterGroup,
Operator,
TagFilter,
)

nested = FilterGroup(
operator=Operator.OR,
conditions=[
Condition(field="tag_name", operator="eq", value="科技"),
Condition(field="invalid_field", operator="eq", value="bad"),
],
)

with pytest.raises(ValueError) as exc_info:
TagFilter(operator=Operator.AND, conditions=[nested])
assert "invalid_field" in str(exc_info.value)

def test_query_params_with_tag_filter(self):
"""QueryParams 支持 TagFilter 并进行校验"""
from whiskerrag_types.model.page import (
Condition,
Operator,
QueryParams,
TagFilter,
)

tf = TagFilter(
operator=Operator.AND,
conditions=[
Condition(field="tag_name", operator="eq", value="科技"),
],
)
# 使用 DummyModel 作为泛型参数
params = QueryParams[DummyModel](tag_filter=tf)
assert isinstance(params.tag_filter, TagFilter)
assert params.tag_filter.conditions[0].field == "tag_name"
Loading