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.1"
version = "0.3.2"
description = "A utlity package for RAG operations"
authors = ["petercat.ai <antd.antgroup@gmail.com>"]
readme = "README.md"
Expand Down
52 changes: 52 additions & 0 deletions src/whiskerrag_types/interface/db_engine_plugin_interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
TaskStatus,
Tenant,
)
from whiskerrag_types.model.artifact_index import ArtifactIndex, ArtifactIndexCreate
from whiskerrag_types.model.chunk import Chunk
from whiskerrag_types.model.retrieval import (
RetrievalByKnowledgeRequest,
Expand Down Expand Up @@ -351,6 +352,57 @@ async def delete_tagging_by_id(
) -> Union[Tagging, None]:
pass

# =================== artifacts ===================
"""
ArtifactIndex 仓储层抽象接口
"""

@abstractmethod
async def add_artifact_list(
self, artifact_list: List[ArtifactIndexCreate]
) -> List[ArtifactIndex]:
"""
批量添加 artifact 记录
返回实际插入成功的 ArtifactIndex 记录列表
"""
pass

@abstractmethod
async def get_artifact_list(
self, page_params: PageQueryParams[ArtifactIndex]
) -> PageResponse[ArtifactIndex]:
"""
分页获取 artifact 列表
"""
pass

@abstractmethod
async def get_artifact_by_id(self, artifact_id: str) -> Union[ArtifactIndex, None]:
"""
根据 artifact_id 获取 ArtifactIndex 详情
"""
pass

@abstractmethod
async def delete_artifact_by_id(
self, artifact_id: str
) -> Union[ArtifactIndex, None]:
"""
根据 artifact_id 删除 ArtifactIndex 记录
返回被删除的记录(如果存在的话)
"""
pass

@abstractmethod
async def update_artifact_space_id(
self, artifact_id: str, new_space_id: str
) -> Union[ArtifactIndex, None]:
"""
更新 artifact 的 space_id 绑定
返回更新后的 ArtifactIndex 记录(如果存在的话)
"""
pass

# =================== dashboard ===================
# TODO: add dashboard related methods

Expand Down
80 changes: 80 additions & 0 deletions src/whiskerrag_types/model/artifact_index.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
from datetime import datetime, timezone
from typing import Any, Dict, Optional
from uuid import UUID, uuid4

from pydantic import BaseModel, Field, field_validator, model_validator

from whiskerrag_types.model.timeStampedModel import TimeStampedModel


class ArtifactIndexCreate(BaseModel):
"""
创建 whisker_artifact_index 条目的入参模型
"""

ecosystem: str = Field(
...,
max_length=32,
description="制品来源生态系统(pypi / npm / maven / go / php)",
)
name: str = Field(
...,
max_length=255,
description="制品名(构建产物名,如 requests / @company/sdk)",
)
version: Optional[str] = Field(
default=None, max_length=64, description="版本号(可为空)"
)
space_id: str = Field(
...,
max_length=255,
pattern=r"^[A-Za-z0-9._@/-]{1,255}$",
description="关联的 whisker_space.space_id",
)
extra: Dict[str, Any] = Field(
default={}, description="额外元数据信息,扩展用,如构建参数、标签、扫描信息等"
)

@field_validator("space_id")
@classmethod
def check_forbidden_sequences(cls, v: str) -> str:
if ".." in v:
raise ValueError("space_id cannot contain consecutive dots '..'")
if "//" in v:
raise ValueError("space_id cannot contain consecutive slashes '//'")
return v

@field_validator("ecosystem")
@classmethod
def normalize_ecosystem(cls, v: str) -> str:
if isinstance(v, str):
return v.strip().lower()
return v

@field_validator("name")
@classmethod
def normalize_name(cls, v: str) -> str:
return v.strip() if isinstance(v, str) else v


class ArtifactIndex(ArtifactIndexCreate, TimeStampedModel):
"""
whisker_artifact_index 模型
"""

artifact_id: str = Field(
default_factory=lambda: str(uuid4()), description="制品索引表主键(UUID字符串)"
)

def update(self, **kwargs: Any) -> "ArtifactIndex":
for key, value in kwargs.items():
setattr(self, key, value)
self.updated_at = datetime.now(timezone.utc)
return self

@model_validator(mode="before")
def preprocess(cls, data: dict) -> dict:
for field, value in list(data.items()):
if isinstance(value, UUID):
data[field] = str(value)
return data
17 changes: 15 additions & 2 deletions src/whiskerrag_types/model/space.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
from typing import Any, Dict, Optional
from uuid import UUID, uuid4

from pydantic import BaseModel, Field, model_validator
from pydantic import BaseModel, Field, field_validator, model_validator

from whiskerrag_types.model.timeStampedModel import TimeStampedModel

Expand All @@ -23,7 +23,7 @@ class SpaceCreate(BaseModel):
space_id: Optional[str] = Field(
default=None,
description="space id, e.g. petercat/bot-group",
pattern=r"^([a-zA-Z0-9-]{1,39}/)?[A-Za-z0-9_.-]{1,100}(/[A-Za-z0-9_.-]{1,100})?(@[A-Za-z0-9_.\-\\/]+)?$",
pattern=r"^[A-Za-z0-9._@/-]{1,255}$",
max_length=255,
)
description: str = Field(..., max_length=255, description="descrition of the space")
Expand All @@ -32,6 +32,19 @@ class SpaceCreate(BaseModel):
description="metadata of the space resource",
)

@field_validator("space_id")
@classmethod
def check_forbidden_sequences(cls, v: str) -> str:
if v is None:
return v
# 禁止连续 ..
if ".." in v:
raise ValueError("space_id cannot contain consecutive dots '..'")
# 禁止连续 //
if "//" in v:
raise ValueError("space_id cannot contain consecutive slashes '//'")
return v


class Space(SpaceCreate, TimeStampedModel):
space_id: str = Field(default_factory=lambda: str(uuid4()), description="space id")
Expand Down
69 changes: 69 additions & 0 deletions tests/whiskerrag_types/test_artifact_index.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
from uuid import UUID, uuid4

import pytest
from pydantic import ValidationError

from whiskerrag_types.model.artifact_index import ArtifactIndex, ArtifactIndexCreate


def test_artifact_create_with_valid_data():
"""正常创建 ArtifactIndexCreate,ecosystem 应该标准化为小写"""
create_data = ArtifactIndexCreate(
ecosystem="PYPI", name=" Requests ", version="1.0.0", space_id="my-space"
)
assert create_data.ecosystem == "pypi"
assert create_data.name == "Requests" # name 去掉首尾空格
assert create_data.version == "1.0.0"
assert create_data.space_id == "my-space"


def test_space_id_validation():
"""space_id 禁止包含 .. 和 //"""
with pytest.raises(ValidationError) as exc_info:
ArtifactIndexCreate(ecosystem="npm", name="left-pad", space_id="invalid..id")
assert "space_id cannot contain consecutive dots" in str(exc_info.value)

with pytest.raises(ValidationError):
ArtifactIndexCreate(ecosystem="npm", name="left-pad", space_id="invalid//id")


def test_artifact_id_auto_uuid():
"""ArtifactIndex 自动生成 artifact_id"""
art = ArtifactIndex(ecosystem="go", name="gin", space_id="my-space")
# 检查 artifact_id 是否为 UUIDv4
UUID(art.artifact_id, version=4)


def test_update_method_and_updated_at():
"""update 方法应该更新字段并刷新更新时间"""
art = ArtifactIndex(ecosystem="php", name="laravel/laravel", space_id="my-space")
old_updated_at = art.updated_at
art.update(version="9.0.0")
assert art.version == "9.0.0"
assert art.updated_at > old_updated_at


def test_model_validator_uuid_conversion():
"""model_validator 应该把 UUID 类型的输入转成 str"""
art = ArtifactIndex(
artifact_id=uuid4(), ecosystem="pypi", name="requests", space_id="my-space"
)
assert isinstance(art.artifact_id, str)
UUID(art.artifact_id, version=4)


def test_artifact_extra_field():
art = ArtifactIndexCreate(
ecosystem="pypi",
name="requests",
version="2.31.0",
space_id="my-space",
extra={"build_env": "linux", "tag": "release"},
)
assert art.extra["build_env"] == "linux"
assert art.extra["tag"] == "release"


def test_extra_field_defaults_to_empty_dict():
art = ArtifactIndexCreate(ecosystem="npm", name="lodash", space_id="npm-space")
assert art.extra == {}