260 lines
8.8 KiB
Python
260 lines
8.8 KiB
Python
"""Database-backed lifecycle operations for published workspace artifacts."""
|
||
from __future__ import annotations
|
||
|
||
import hashlib
|
||
import mimetypes
|
||
import os
|
||
import shutil
|
||
from datetime import datetime, timezone
|
||
from pathlib import Path
|
||
from uuid import UUID, uuid4
|
||
|
||
from sqlalchemy import or_, select, text
|
||
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
||
|
||
from .artifacts import ARTIFACT_TRASH_DIR, ArtifactRef
|
||
from .storage import session_scope
|
||
from .storage.models import Artifact
|
||
|
||
|
||
def _rel(root: Path, path: Path) -> str:
|
||
return Path(path).resolve().relative_to(Path(root).resolve()).as_posix()
|
||
|
||
|
||
def _hash_file(path: Path) -> str:
|
||
digest = hashlib.sha256()
|
||
with Path(path).open("rb") as stream:
|
||
for chunk in iter(lambda: stream.read(1024 * 1024), b""):
|
||
digest.update(chunk)
|
||
return digest.hexdigest()
|
||
|
||
|
||
def artifact_subtree_clause(path: str):
|
||
"""精确匹配文件/目录自身及其子树,并转义 LIKE 元字符。"""
|
||
escaped = path.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||
return or_(
|
||
Artifact.current_path == path,
|
||
Artifact.current_path.like(escaped + "/%", escape="\\"),
|
||
)
|
||
|
||
|
||
def register_published_artifacts(
|
||
*,
|
||
user_id: UUID,
|
||
task_id: UUID,
|
||
user_root: Path,
|
||
working_dir: Path,
|
||
refs: tuple[dict, ...],
|
||
preserve_existing_labels: bool = False,
|
||
software_job_id: UUID | None = None,
|
||
) -> tuple[dict, ...]:
|
||
"""Upsert active artifact identities and return version-2 message refs."""
|
||
root = Path(user_root).resolve()
|
||
wd = Path(working_dir).resolve()
|
||
output: list[dict] = []
|
||
with session_scope() as session:
|
||
for ref in refs:
|
||
task_path = str(ref.get("path") or "")
|
||
path = (wd / Path(task_path)).resolve()
|
||
path.relative_to(wd)
|
||
if not path.is_file():
|
||
continue
|
||
current_path = _rel(root, path)
|
||
label = str(ref.get("label") or "")
|
||
media_type = str(ref.get("media_type") or "") or mimetypes.guess_type(path.name)[0]
|
||
size_bytes = path.stat().st_size
|
||
content_sha256 = _hash_file(path)
|
||
update_values = {
|
||
"media_type": media_type,
|
||
"size_bytes": size_bytes,
|
||
"content_sha256": content_sha256,
|
||
"updated_at": datetime.now(timezone.utc),
|
||
}
|
||
if not preserve_existing_labels:
|
||
update_values["label"] = label
|
||
if software_job_id is not None:
|
||
update_values["software_job_id"] = software_job_id
|
||
statement = pg_insert(Artifact).values(
|
||
user_id=user_id,
|
||
origin_task_id=task_id,
|
||
software_job_id=software_job_id,
|
||
current_path=current_path,
|
||
label=label,
|
||
media_type=media_type,
|
||
size_bytes=size_bytes,
|
||
content_sha256=content_sha256,
|
||
).on_conflict_do_update(
|
||
index_elements=[Artifact.user_id, Artifact.current_path],
|
||
# 必须与 0028 的部分唯一索引谓词按 SQL 字面量一致。若用
|
||
# ``Artifact.status == "active"``,SQLAlchemy 会编译成绑定参数,
|
||
# PostgreSQL 无法据此推断 ON CONFLICT 对应哪个部分索引。
|
||
index_where=text("status = 'active'"),
|
||
set_=update_values,
|
||
).returning(Artifact.artifact_id)
|
||
artifact_id = session.execute(statement).scalar_one()
|
||
output.append(ArtifactRef(
|
||
path=task_path,
|
||
label=label,
|
||
artifact_id=artifact_id,
|
||
version=2,
|
||
).as_dict())
|
||
return tuple(output)
|
||
|
||
|
||
def register_workspace_artifact(
|
||
*,
|
||
user_id: UUID,
|
||
task_id: UUID,
|
||
user_root: Path,
|
||
working_dir: Path,
|
||
path: str,
|
||
label: str = "",
|
||
) -> dict:
|
||
"""Register one task-relative file for an internal platform workflow.
|
||
|
||
Unlike ``publish_artifacts``, this only creates or refreshes the stable artifact
|
||
identity. It does not attach the file to an assistant message or present it as a
|
||
final deliverable.
|
||
"""
|
||
refs = register_published_artifacts(
|
||
user_id=user_id,
|
||
task_id=task_id,
|
||
user_root=user_root,
|
||
working_dir=working_dir,
|
||
refs=({"path": path, "label": label},),
|
||
preserve_existing_labels=True,
|
||
)
|
||
if len(refs) != 1 or not refs[0].get("artifact_id"):
|
||
raise ValueError("workspace artifact could not be registered")
|
||
return refs[0]
|
||
|
||
|
||
def rename_active_artifacts(
|
||
*,
|
||
user_id: UUID,
|
||
user_root: Path,
|
||
old_path: Path,
|
||
new_path: Path,
|
||
) -> int:
|
||
"""Rewrite active artifact paths for one file or directory subtree."""
|
||
root = Path(user_root).resolve()
|
||
old_rel = _rel(root, old_path)
|
||
new_rel = _rel(root, new_path)
|
||
with session_scope() as session:
|
||
rows = session.execute(
|
||
select(Artifact).where(
|
||
Artifact.user_id == user_id,
|
||
Artifact.status == "active",
|
||
artifact_subtree_clause(old_rel),
|
||
)
|
||
).scalars().all()
|
||
changed = 0
|
||
for row in rows:
|
||
if row.current_path == old_rel:
|
||
suffix = ""
|
||
elif row.current_path.startswith(old_rel + "/"):
|
||
suffix = row.current_path[len(old_rel):]
|
||
else:
|
||
continue
|
||
row.current_path = new_rel + suffix
|
||
changed += 1
|
||
return changed
|
||
|
||
|
||
def copy_active_artifacts(
|
||
*,
|
||
user_id: UUID,
|
||
user_root: Path,
|
||
source: Path,
|
||
target: Path,
|
||
) -> int:
|
||
"""Create independent artifact identities for copied files."""
|
||
root = Path(user_root).resolve()
|
||
source_rel = _rel(root, source)
|
||
target_rel = _rel(root, target)
|
||
with session_scope() as session:
|
||
sources = session.execute(
|
||
select(Artifact).where(
|
||
Artifact.user_id == user_id,
|
||
Artifact.status == "active",
|
||
artifact_subtree_clause(source_rel),
|
||
)
|
||
).scalars().all()
|
||
created = 0
|
||
for original in sources:
|
||
if original.current_path == source_rel:
|
||
suffix = ""
|
||
elif original.current_path.startswith(source_rel + "/"):
|
||
suffix = original.current_path[len(source_rel):]
|
||
else:
|
||
continue
|
||
copied_path = target_rel + suffix
|
||
session.add(Artifact(
|
||
user_id=user_id,
|
||
origin_task_id=original.origin_task_id,
|
||
copied_from_artifact_id=original.artifact_id,
|
||
software_job_id=original.software_job_id,
|
||
current_path=copied_path,
|
||
label=(original.label + " 副本").strip(),
|
||
media_type=original.media_type,
|
||
size_bytes=original.size_bytes,
|
||
content_sha256=original.content_sha256,
|
||
))
|
||
created += 1
|
||
return created
|
||
|
||
|
||
def trash_active_artifacts(
|
||
*,
|
||
user_id: UUID,
|
||
user_root: Path,
|
||
target: Path,
|
||
) -> int:
|
||
"""Move matching active artifacts to hidden trash and mark them deleted."""
|
||
root = Path(user_root).resolve()
|
||
source = Path(target).resolve()
|
||
source_rel = _rel(root, source)
|
||
entry = (
|
||
root / ARTIFACT_TRASH_DIR
|
||
/ datetime.now(timezone.utc).strftime("%Y/%m/%d")
|
||
/ uuid4().hex
|
||
)
|
||
moved: list[tuple[Path, Path]] = []
|
||
try:
|
||
with session_scope() as session:
|
||
rows = session.execute(
|
||
select(Artifact).where(
|
||
Artifact.user_id == user_id,
|
||
Artifact.status == "active",
|
||
artifact_subtree_clause(source_rel),
|
||
).with_for_update()
|
||
).scalars().all()
|
||
matches = [
|
||
row for row in rows
|
||
if row.current_path == source_rel
|
||
or row.current_path.startswith(source_rel + "/")
|
||
]
|
||
if not matches:
|
||
return 0
|
||
for row in matches:
|
||
original = root / Path(row.current_path)
|
||
if not original.is_file():
|
||
continue
|
||
destination = entry / "files" / Path(row.current_path)
|
||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||
os.replace(original, destination)
|
||
moved.append((original, destination))
|
||
row.status = "deleted"
|
||
row.deleted_at = datetime.now(timezone.utc)
|
||
row.trash_path = _rel(root, destination)
|
||
return len(moved)
|
||
except Exception:
|
||
for original, destination in reversed(moved):
|
||
try:
|
||
original.parent.mkdir(parents=True, exist_ok=True)
|
||
os.replace(destination, original)
|
||
except OSError:
|
||
pass
|
||
shutil.rmtree(entry, ignore_errors=True)
|
||
raise
|