Compare commits

..

12 Commits

Author SHA1 Message Date
Jann Stute 92b40a72a9 refactor: move set_version to DBMigrations 2026-07-24 16:13:36 +02:00
Jann Stute ccd07ee0ff fix: add override decorators 2026-07-24 16:02:28 +02:00
Jann Stute c799450072 refactor: condense imports 2026-07-24 16:00:31 +02:00
Jann Stute 200a285de6 fix: allow set_version to fail, but don't commit in that case 2026-07-23 22:24:19 +02:00
Jann Stute 75c586db79 fix: some syntax errors had slipped through 2026-07-23 22:09:38 +02:00
Jann Stute eef0da6901 refactor: package each migration in a class 2026-07-23 21:52:26 +02:00
Jann Stute 92809cc225 refactor: move migrations to different file 2026-07-23 21:10:59 +02:00
Jann Stute a22482df9e refactor: don't require setting library_dir to create a backup 2026-07-23 20:33:38 +02:00
Jann Stute c5402e6bda refactor: don't blindly create all tables in the beginning
The only table that has been added since DB version 6 (the earliest supported version), is the versions table in DB version 101.
This commit removes the "create all tables" statement, and instead creates the versions table in the 101 migration.
See 12e074b71d.
2026-07-23 20:23:21 +02:00
Jann Stute 24f9f27e63 refactor: remove unnecessary assurance
Bumping the auto increment value has been done since the original sql PR, so it doesn't need to be done on migrations.
See e5e7b8afc6.
2026-07-23 19:37:29 +02:00
Jann Stute 0d1597d520 refactor: inline make_tables + minor cleanup 2026-07-23 19:27:48 +02:00
Jann Stute 617ff710e3 fix: backup library before making any changes 2026-07-23 19:20:40 +02:00
26 changed files with 1717 additions and 1619 deletions
+2 -2
View File
@@ -88,7 +88,7 @@ filterwarnings = [
ignore = [
".venv/**",
"src/tagstudio/core/library/json/",
"src/tagstudio/renderers/vendored/pydub/",
"src/tagstudio/qt/previews/vendored/pydub/",
]
include = ["src/tagstudio", "tests"]
# Reference for the settings here: https://github.com/microsoft/pyright/blob/main/docs/configuration.md
@@ -117,7 +117,7 @@ ignore = ["D100", "D101", "D102", "D103", "D104", "D105", "D106", "D107"]
[tool.ruff.lint.per-file-ignores]
"tests/**" = ["D", "E402"]
"src/tagstudio/renderers/vendored/**" = ["B", "E", "N", "UP", "SIM115"]
"src/tagstudio/qt/previews/vendored/**" = ["B", "E", "N", "UP", "SIM115"]
[tool.ruff.lint.pydocstyle]
convention = "google"
@@ -4,6 +4,11 @@
from sqlalchemy import text
from tagstudio.core.library.alchemy.fields import (
DatetimeFieldTemplate,
TextFieldTemplate,
)
SQL_FILENAME: str = "ts_library.sqlite"
JSON_FILENAME: str = "ts_library.json"
@@ -32,3 +37,15 @@ WITH RECURSIVE ChildTags AS (
)
SELECT tag_id FROM ChildTags;
""")
DEFAULT_FIELD_TEMPLATES = (
TextFieldTemplate(name="Title"),
TextFieldTemplate(name="Author"),
TextFieldTemplate(name="Artist"),
TextFieldTemplate(name="URL"),
TextFieldTemplate(name="Description", is_multiline=True),
TextFieldTemplate(name="Notes", is_multiline=True),
TextFieldTemplate(name="Comments", is_multiline=True),
DatetimeFieldTemplate(name="Date"),
)
+1 -39
View File
@@ -6,12 +6,9 @@ from pathlib import Path
from typing import override
import structlog
from sqlalchemy import Dialect, Engine, String, TypeDecorator, create_engine, text
from sqlalchemy.exc import OperationalError
from sqlalchemy import Dialect, String, TypeDecorator
from sqlalchemy.orm import DeclarativeBase
from tagstudio.core.constants import RESERVED_TAG_END
logger = structlog.getLogger(__name__)
@@ -34,38 +31,3 @@ class PathType(TypeDecorator):
class Base(DeclarativeBase):
type_annotation_map = {Path: PathType}
def make_engine(connection_string: str) -> Engine:
return create_engine(connection_string)
def make_tables(engine: Engine) -> None:
logger.info("[Library] Creating DB tables...")
with engine.connect() as conn:
# TODO: this should instead be migrations that create the exact tables that were added in
# the respective DB versions
Base.metadata.create_all(conn)
conn.commit()
# TODO: this needs to be a migration
# tag IDs < 1000 are reserved
# create tag and delete it to bump the autoincrement sequence
# TODO - find a better way
# is this the better way?
result = conn.execute(text("SELECT SEQ FROM sqlite_sequence WHERE name='tags'"))
autoincrement_val = result.scalar()
if not autoincrement_val or autoincrement_val <= RESERVED_TAG_END:
try:
conn.execute(
text(
"INSERT INTO tags "
"(id, name, color_namespace, color_slug, is_category, is_hidden) VALUES "
f"({RESERVED_TAG_END}, 'temp', NULL, NULL, false, false)"
)
)
conn.execute(text(f"DELETE FROM tags WHERE id = {RESERVED_TAG_END}"))
conn.commit()
except OperationalError as e:
logger.error("Could not initialize built-in tags", error=e)
conn.rollback()
+55 -452
View File
@@ -2,11 +2,6 @@
# SPDX-License-Identifier: GPL-3.0-only
# NOTE: This file contains necessary use of deprecated first-party code until that
# code is removed in a future version (prefs).
# pyright: reportDeprecated=false
import re
import shutil
import sys
@@ -21,7 +16,6 @@ from typing import TYPE_CHECKING
import sqlalchemy
import structlog
import ujson
from humanfriendly import format_timespan # pyright: ignore[reportUnknownVariableType]
from sqlalchemy import (
URL,
@@ -44,7 +38,7 @@ from sqlalchemy import (
update,
)
from sqlalchemy.dialects import sqlite
from sqlalchemy.exc import IntegrityError
from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy.orm import (
InstanceState,
Session,
@@ -72,16 +66,13 @@ from tagstudio.core.library.alchemy.constants import (
DB_VERSION,
DB_VERSION_CURRENT_KEY,
DB_VERSION_INITIAL_KEY,
DEFAULT_FIELD_TEMPLATES,
JSON_FILENAME,
SQL_FILENAME,
TAG_CHILDREN_QUERY,
)
from tagstudio.core.library.alchemy.db import make_tables
from tagstudio.core.library.alchemy.enums import (
MAX_SQL_VARIABLES,
BrowsingState,
SortingModeEnum,
)
from tagstudio.core.library.alchemy.db import Base as ModelBase
from tagstudio.core.library.alchemy.enums import MAX_SQL_VARIABLES, BrowsingState, SortingModeEnum
from tagstudio.core.library.alchemy.fields import (
LEGACY_FIELD_MAP,
BaseField,
@@ -92,6 +83,7 @@ from tagstudio.core.library.alchemy.fields import (
TextFieldTemplate,
)
from tagstudio.core.library.alchemy.joins import TagEntry, TagParent
from tagstudio.core.library.alchemy.migrations import DBMigrations, MigrationError
from tagstudio.core.library.alchemy.models import (
Entry,
Namespace,
@@ -104,7 +96,6 @@ from tagstudio.core.library.alchemy.visitors import SQLBoolExpressionBuilder
from tagstudio.core.library.ignore import migrate_ext_list
from tagstudio.core.library.json.library import Library as JsonLibrary
from tagstudio.core.utils.types import unwrap
from tagstudio.qt.translations import Translations
if TYPE_CHECKING:
from sqlalchemy import Select
@@ -170,20 +161,6 @@ def get_default_tags() -> tuple[Tag, ...]:
return archive_tag, favorite_tag, meta_tag
def get_default_field_templates() -> tuple[BaseFieldTemplate, ...]:
"""Return the default field templates for a new TagStudio library."""
title = TextFieldTemplate(name="Title")
author = TextFieldTemplate(name="Author")
artist = TextFieldTemplate(name="Artist")
url = TextFieldTemplate(name="URL")
description = TextFieldTemplate(name="Description", is_multiline=True)
notes = TextFieldTemplate(name="Notes", is_multiline=True)
comments = TextFieldTemplate(name="Comments", is_multiline=True)
date = DatetimeFieldTemplate(name="Date")
return title, author, artist, url, description, notes, comments, date
# The difference in the number of default JSON tags vs default tags in the current version.
DEFAULT_TAG_DIFF: int = len(get_default_tags()) - len([TAG_ARCHIVED, TAG_FAVORITE])
@@ -431,21 +408,41 @@ class Library:
self, library_dir: Path, in_memory: bool, sql_filename: str = SQL_FILENAME
) -> LibraryStatus:
self.engine = self.__get_engine(library_dir, in_memory, sql_filename)
loaded_db_version: int = 0
logger.info(
"[Library] Opening SQLite Library",
library_dir=library_dir,
)
logger.info(f"[Library] Library DB version: {loaded_db_version}")
make_tables(self.engine)
logger.info("[Library] Creating DB tables...")
with self.engine.connect() as conn:
ModelBase.metadata.create_all(conn)
conn.commit()
# TODO - find a better way
# is this the better way?
# Could we perhaps update the row we are reading from here?
result = conn.execute(text("SELECT SEQ FROM sqlite_sequence WHERE name='tags'"))
autoincrement_val = result.scalar()
if not autoincrement_val or autoincrement_val <= RESERVED_TAG_END:
try:
conn.execute(
text(
"INSERT INTO tags "
"(id, name, color_namespace, color_slug, is_category, is_hidden) "
f"VALUES ({RESERVED_TAG_END}, 'temp', NULL, NULL, false, false)"
)
)
conn.execute(text(f"DELETE FROM tags WHERE id = {RESERVED_TAG_END}"))
conn.commit()
except OperationalError as e:
logger.error("Could not initialize built-in tags", error=e)
conn.rollback()
with Session(self.engine) as session:
# Add default tag color namespaces.
namespaces = default_color_groups.namespaces()
# TODO: are all of these commits necessary?
session.add_all(namespaces)
session.flush()
@@ -465,7 +462,7 @@ class Library:
session.flush()
# Add default field templates
for template in get_default_field_templates():
for template in DEFAULT_FIELD_TEMPLATES:
session.add(template)
session.flush()
@@ -506,414 +503,25 @@ class Library:
def open_sqlite_library(
self, library_dir: Path, in_memory: bool, sql_filename: str = SQL_FILENAME
) -> LibraryStatus:
logger.info("[Library] Opening SQLite Library", library_dir=library_dir)
self.engine = self.__get_engine(library_dir, in_memory, sql_filename)
loaded_db_version: int = 0
initial_db_version: int = DB_VERSION
logger.info(
"[Library] Opening SQLite Library",
library_dir=library_dir,
)
try:
migrations = DBMigrations(library_dir, self.engine)
# Don't check DB version when creating new library
loaded_db_version = self.get_version(DB_VERSION_CURRENT_KEY)
initial_db_version = self.get_version(DB_VERSION_INITIAL_KEY)
# save backup if patches will be applied
if migrations.required:
Library.save_library_backup_to_disk(library_dir)
# ======================== Library Database Version Checking =======================
# DB_VERSION 6 is the first supported SQLite DB version.
# If the DB_VERSION is >= 100, that means it's a compound major + minor version.
# - Dividing by 100 and flooring gives the major (breaking changes) version.
# - If a DB has major version higher than the current program, don't load it.
# - If only the minor version is higher, it's still allowed to load.
if loaded_db_version < 6 or (
loaded_db_version >= 100 and loaded_db_version // 100 > DB_VERSION // 100
):
mismatch_text = Translations["status.library_version_mismatch"]
found_text = Translations["status.library_version_found"]
expected_text = Translations["status.library_version_expected"]
return LibraryStatus(
success=False,
message=(
f"{mismatch_text}\n"
f"{found_text} v{loaded_db_version}, "
f"{expected_text} v{DB_VERSION}"
),
)
logger.info(f"[Library] Library DB version: {loaded_db_version}")
# TODO: this is very sketchy; blindly creating all tables the newest DB version should have
# without considering what version the DB is currently on and then doing all of the
# migrations after that seems like it could cause problems in some scenarios.
# instead only have this on creation and create new tables as part of migrations
# Note: this actually produces an error and fails to initialise built-in tags when opening
# a library that doesn't yet have the is_hidden property on the tags table
make_tables(self.engine)
# save backup if patches will be applied
if loaded_db_version < DB_VERSION:
self.library_dir = library_dir
self.save_library_backup_to_disk()
self.library_dir = None
# migrate DB step by step from one version to the next
# (migration_method, db_version, initial_db_version)
migrations = [
(self.__apply_db7_migration, 7, None), # changes: value_type, tags
(self.__apply_db8_migration, 8, None), # changes: tag_colors
(self.__apply_db9_migration, 9, None), # changes: entries
(self.__apply_db100_migration, 100, None), # changes: tag_parents
(self.__apply_db101_migration, 101, None), # changes: versions
(self.__apply_db102_migration, 102, None), # changes: tag_parents
(self.__apply_db103_migration, 103, None), # changes: tags
(self.__apply_db104_migration, 104, None), # changes: deletes preferences
(self.__apply_db200_migration, 200, None), # changes: field tables
(self.__apply_db201_migration, 201, 200), # changes: field tables
(self.__apply_db202_migration, 202, None), # changes: tag_parents
(self.__apply_db300_migration, 300, None), # changes: deletes folders
]
for migration, v, iv in migrations:
if loaded_db_version < v and (iv is None or initial_db_version < iv):
logger.info(f"[Library][Migration][{v}] Starting DB Migration")
with Session(self.engine) as session:
# any error causes transaction to rollback
migration(session, library_dir)
loaded_db_version = v
self.set_version(session, DB_VERSION_CURRENT_KEY, v)
session.commit()
logger.info(f"[Library][Migration][{v}] Completed DB Migration")
assert loaded_db_version >= DB_VERSION, (
"Ran all migrations, but the DB is still not on the newest version"
)
logger.info(f"[Library] Library migrated to DB version {DB_VERSION}")
migrations.run()
except MigrationError as e:
return LibraryStatus(success=False, message=e.args[0])
# everything is fine, set the library path
self.library_dir = library_dir
return LibraryStatus(success=True, library_path=library_dir)
def __apply_db7_migration(self, session: Session, _library_dir: Path):
"""Migrate DB from DB_VERSION 6 to 7."""
logger.info("[Library][Migration][7] Applying patches to DB_VERSION: 6 library...")
# Repair tags that may have a disambiguation_id pointing towards a deleted tag.
# TODO: combine into single sql statement
all_tag_ids = session.scalars(text("SELECT DISTINCT id FROM tags")).all()
disam_stmt = (
update(Tag)
.where(Tag.disambiguation_id.not_in(all_tag_ids))
.values(disambiguation_id=None)
)
session.execute(disam_stmt)
session.flush()
def __apply_db8_migration(self, session: Session, library_dir: Path):
"""Migrate DB from DB_VERSION 7 to 8."""
# Add the missing color_border column to the TagColorGroups table.
session.execute(
text("ALTER TABLE tag_colors ADD COLUMN color_border BOOLEAN DEFAULT FALSE NOT NULL")
)
session.flush()
logger.info("[Library][Migration][8] Added color_border column to tag_colors table")
# collect new default tag colors
tag_colors: list[TagColorGroup] = [
color
for color in default_color_groups.shades()
if color.slug in ["burgundy", "dark-teal", "dark_lavender"]
]
# Add any new default colors introduced in DB_VERSION 8
for color in tag_colors:
session.add(color)
session.flush()
logger.info(
"[Library][Migration][8] Migrated tag colors to DB_VERSION 8+",
color_name=tag_colors,
)
# Update Neon colors to use the the color_border property
for color in default_color_groups.neon():
neon_stmt = (
update(TagColorGroup)
.where(
and_(
TagColorGroup.namespace == color.namespace,
TagColorGroup.slug == color.slug,
)
)
.values(
slug=color.slug,
namespace=color.namespace,
name=color.name,
primary=color.primary,
secondary=color.secondary,
color_border=color.color_border,
)
)
session.execute(neon_stmt)
session.flush()
def __apply_db9_migration(self, session: Session, library_dir: Path):
"""Migrate DB from DB_VERSION 8 to 9."""
# Apply database schema changes
add_filename_column = text(
"ALTER TABLE entries ADD COLUMN filename TEXT NOT NULL DEFAULT ''"
)
session.execute(add_filename_column)
session.flush()
logger.info("[Library][Migration][9] Added filename column to entries table")
# Populate the new filename column.
for entry in self.__all_entries(session):
entry.filename = entry.path.name
session.merge(entry)
session.flush()
logger.info("[Library][Migration][9] Populated filename column in entries table")
def __apply_db100_migration(self, session: Session, library_dir: Path):
"""Migrate DB to DB_VERSION 100."""
# Repair parent-child tag relationships that are the wrong way around.
stmt = update(TagParent).values(
parent_id=TagParent.child_id,
child_id=TagParent.parent_id,
)
session.execute(stmt)
session.flush()
logger.info("[Library][Migration][100] Refactored TagParent table")
def __apply_db101_migration(self, session: Session, library_dir: Path):
"""Migrate DB to DB_VERSION 101."""
# Ensure version rows are present
session.add(Version(key=DB_VERSION_INITIAL_KEY, value=100))
session.flush()
def __apply_db102_migration(self, session: Session, library_dir: Path):
"""Migrate DB to DB_VERSION 102."""
# delete TagParents with a dangling parent reference
stmt = delete(TagParent).where(TagParent.parent_id.not_in(select(Tag.id).distinct()))
session.execute(stmt)
session.flush()
logger.info("[Library][Migration][102] Verified TagParent table data")
def __apply_db103_migration(self, session: Session, library_dir: Path):
"""Migrate DB from DB_VERSION 102 to 103."""
# add the new hidden column for tags
session.execute(text("ALTER TABLE tags ADD COLUMN is_hidden BOOLEAN NOT NULL DEFAULT 0"))
session.flush()
logger.info("[Library][Migration][103] Added is_hidden column to tags table")
# mark the "Archived" tag as hidden
session.query(Tag).filter(Tag.id == TAG_ARCHIVED).update({"is_hidden": True})
session.flush()
logger.info("[Library][Migration][103] Updated archived tag to be hidden")
def __apply_db104_migration(self, session: Session, library_dir: Path):
"""Migrate DB from DB_VERSION 103 to 104."""
# Convert file extension list to ts_ignore file, if a .ts_ignore file does not exist
self.__migrate_sql_to_ts_ignore(session, library_dir)
session.execute(text("DROP TABLE preferences"))
session.flush()
def __migrate_sql_to_ts_ignore(self, session: Session, library_dir: Path):
# Do not continue if existing '.ts_ignore' file is found
ts_ignore = library_dir / TS_FOLDER_NAME / IGNORE_NAME
if Path(ts_ignore).exists():
return
# Load legacy extension data
extensions: list[str] = ujson.loads(
unwrap(
session.scalar(text("SELECT value FROM preferences WHERE key = 'EXTENSION_LIST'"))
)
)
is_exclude_list: bool = unwrap(
session.scalar(text("SELECT value FROM preferences WHERE key = 'IS_EXCLUDE_LIST'"))
)
with open(ts_ignore, "w") as f:
f.write(migrate_ext_list(extensions, is_exclude_list))
def __apply_db200_migration(self, session: Session, library_dir: Path):
"""Migrate DB to DB_VERSION 200."""
# Drop unused 'boolean_fields' and 'value_type' tables
logger.info("[Library][Migration][200] Dropping boolean_fields and value_type tables...")
session.execute(text("DROP TABLE boolean_fields"))
session.execute(text("DROP TABLE value_type"))
# Add 'name' column to text_fields and datetime_fields tables
logger.info("[Library][Migration][200] Adding name columns to field tables...")
stmt = text('ALTER TABLE text_fields ADD COLUMN name VARCHAR DEFAULT ""')
session.execute(stmt)
stmt = text('ALTER TABLE datetime_fields ADD COLUMN name VARCHAR DEFAULT ""')
session.execute(stmt)
# Drop unnecessary 'position' columns
logger.info("[Library][Migration][200] Dropping position columns to field tables...")
session.execute(text("ALTER TABLE datetime_fields DROP COLUMN position"))
session.execute(text("ALTER TABLE text_fields DROP COLUMN position"))
# Add 'is_multiline' column to text_fields table
logger.info("[Library][Migration][200] Adding is_multiline column to text_fields...")
stmt = text("ALTER TABLE text_fields ADD COLUMN is_multiline BOOLEAN NOT NULL DEFAULT 0")
session.execute(stmt)
session.flush()
# Move values from old `type_key` columns into new `name` columns
logger.info("[Library][Migration][200] Moving values from type_key columns to name...")
session.execute(text("UPDATE text_fields SET name = type_key"))
session.execute(text("UPDATE datetime_fields SET name = type_key"))
session.flush()
# Change `name` values to title case
logger.info("[Library][Migration][200] Normalizing TextField names...")
for text_field in session.execute(select(TextField)).scalars():
# NOTE: The only exception to the "Title Case" conversion is the "URL" field.
text_field.name = text_field.name.title().replace("Url", "URL").replace("_", " ")
logger.info("[Library][Migration][200] Normalizing DatetimeField names...")
for datetime_field in session.execute(select(DatetimeField)).scalars():
datetime_field.name = datetime_field.name.title().replace("_", " ")
session.flush()
# Add correct `is_multiline` values to text_fields table
logger.info("[Library][Migration][200] Updating is_multiline for legacy TEXT_BOXes...")
text_boxes = [
x.get("name") for x in LEGACY_FIELD_MAP.values() if x.get("is_multiline") is True
]
update_stmt = (
update(TextField).where(TextField.name.in_(text_boxes)).values(is_multiline=True)
)
session.execute(update_stmt)
session.flush()
# Repair legacy "Description" fields to use is_multiline = True
logger.info("[Library][Migration][200] Repairing legacy Description fields...")
desc_stmt = (
update(TextField)
.where(TextField.name == "Description" and TextField.is_multiline == False) # noqa: E712
.values(is_multiline=True)
)
session.execute(desc_stmt)
# Repair legacy "Comments" fields to use is_multiline = True
logger.info("[Library][Migration][200] Repairing legacy Comment fields...")
comm_stmt = (
update(TextField)
.where(TextField.name == "Comments" and TextField.is_multiline == False) # noqa: E712
.values(is_multiline=True)
)
session.execute(comm_stmt)
# Add default field templates
logger.info("[Library][Migration][200] Adding default field templates...")
for template in get_default_field_templates():
session.add(template)
session.flush()
# DB indices for improved performance
session.execute(
text("CREATE INDEX IF NOT EXISTS idx_tags_name_shorthand ON tags (name, shorthand)")
)
session.execute(
text("CREATE INDEX IF NOT EXISTS idx_tag_parents_child_id ON tag_parents (child_id)")
)
session.execute(
text("CREATE INDEX IF NOT EXISTS idx_tag_entries_entry_id ON tag_entries (entry_id)")
)
def __apply_db201_migration(self, session: Session, library_dir: Path):
"""Migrate DB to DB_VERSION 201."""
create_text_fields_table = text("""
CREATE TABLE text_fields_new (
id INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT,
name VARCHAR NOT NULL,
entry_id INTEGER NOT NULL,
value VARCHAR,
is_multiline BOOLEAN NOT NULL,
FOREIGN KEY(entry_id) REFERENCES entries (id)
)
""")
create_datetime_fields_table = text("""
CREATE TABLE datetime_fields_new (
id INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT,
name VARCHAR NOT NULL,
entry_id INTEGER NOT NULL,
value VARCHAR,
FOREIGN KEY(entry_id) REFERENCES entries (id)
)
""")
logger.info("[Library][Migration][201] Dropping type_key from text_fields table...")
session.execute(create_text_fields_table)
session.flush()
session.execute(
text("""
INSERT INTO text_fields_new (id, name, entry_id, value, is_multiline)
SELECT id, name, entry_id, value, is_multiline
FROM text_fields
""")
)
session.execute(text("DROP TABLE text_fields"))
session.execute(text("ALTER TABLE text_fields_new RENAME TO text_fields"))
logger.info("[Library][Migration][201] Dropping type_key from datetime_fields table...")
session.execute(create_datetime_fields_table)
session.flush()
session.execute(
text("""
INSERT INTO datetime_fields_new (id, name, entry_id, value)
SELECT id, name, entry_id, value
FROM datetime_fields
""")
)
session.execute(text("DROP TABLE datetime_fields"))
session.execute(text("ALTER TABLE datetime_fields_new RENAME TO datetime_fields"))
session.flush()
def __apply_db202_migration(self, session: Session, library_dir: Path):
"""Migrate DB to DB_VERSION 202."""
stmt = delete(TagParent).where(TagParent.child_id.not_in(select(Tag.id).distinct()))
session.execute(stmt)
session.flush()
logger.info("[Library][Migration][202] Verified TagParent table data")
def __apply_db300_migration(self, session: Session, library_dir: Path):
## remove folder_id column from entries table
# create new table in the desired scheme (without folder_id column)
session.execute(
text("""
CREATE TABLE entries_new (
id INTEGER NOT NULL,
path VARCHAR NOT NULL,
suffix VARCHAR NOT NULL,
date_created DATETIME,
date_modified DATETIME,
date_added DATETIME,
filename TEXT NOT NULL DEFAULT '',
PRIMARY KEY (id),
UNIQUE (path)
)
""")
)
session.flush()
# transfer data to new table
session.execute(
text("""
INSERT INTO entries_new (id, path, suffix, date_created, date_modified, date_added,
filename)
SELECT id, path, suffix, date_created, date_modified, date_added, filename
FROM entries
""")
)
# delete old table
session.execute(text("DROP TABLE entries"))
# rename new table to old table
session.execute(text("ALTER TABLE entries_new RENAME TO entries"))
session.flush()
## drop table "folders"
session.execute(text("DROP TABLE folders"))
session.flush()
@property
def field_templates(self) -> Sequence[BaseFieldTemplate]:
with Session(self.engine) as session:
@@ -1061,7 +669,8 @@ class Library:
with Session(self.engine) as session:
return unwrap(session.scalar(select(func.count(Entry.id))))
def __all_entries(self, session: Session, with_joins: bool = False) -> Iterator[Entry]:
@staticmethod
def _all_entries(session: Session, with_joins: bool = False) -> Iterator[Entry]:
"""Load entries without joins."""
stmt = select(Entry)
if with_joins:
@@ -1090,7 +699,7 @@ class Library:
def all_entries(self, with_joins: bool = False) -> Iterator[Entry]:
"""Load entries without joins."""
with Session(self.engine) as session:
return self.__all_entries(session, with_joins)
return Library._all_entries(session, with_joins)
@property
def tags(self) -> list[Tag]:
@@ -1838,16 +1447,17 @@ class Library:
session.rollback()
return None
def save_library_backup_to_disk(self) -> Path:
assert isinstance(self.library_dir, Path)
makedirs(str(self.library_dir / TS_FOLDER_NAME / BACKUP_FOLDER_NAME), exist_ok=True)
@staticmethod
def save_library_backup_to_disk(library_dir: Path) -> Path:
assert isinstance(library_dir, Path)
makedirs(str(library_dir / TS_FOLDER_NAME / BACKUP_FOLDER_NAME), exist_ok=True)
filename = f"ts_library_backup_{datetime.now(UTC).strftime('%Y_%m_%d_%H%M%S')}.sqlite"
target_path = self.library_dir / TS_FOLDER_NAME / BACKUP_FOLDER_NAME / filename
target_path = library_dir / TS_FOLDER_NAME / BACKUP_FOLDER_NAME / filename
shutil.copy2(
self.library_dir / TS_FOLDER_NAME / SQL_FILENAME,
library_dir / TS_FOLDER_NAME / SQL_FILENAME,
target_path,
)
@@ -2139,8 +1749,12 @@ class Library:
Args:
key(str): The key for the name of the version type to set.
"""
with Session(self.engine) as session:
engine = sqlalchemy.inspect(self.engine)
return Library._get_version(self.engine, key)
@staticmethod
def _get_version(engine, key: str) -> int:
with Session(engine) as session:
engine = sqlalchemy.inspect(engine)
try:
# "Version" table added in DB_VERSION 101
if engine and engine.has_table("versions"):
@@ -2160,17 +1774,6 @@ class Library:
except Exception:
return 0
def set_version(self, session: Session, key: str, value: int) -> None:
"""Set a version value to the DB.
Args:
session(Session): The SQLAlchemy DB Session to use.
key(str): The key for the name of the version type to set.
value(int): The version value to set.
"""
# Insert if key has no value yet, otherwise update the value
session.merge(Version(key=key, value=value))
def mirror_entry_fields(self, entries: list[Entry]) -> None:
"""Mirror fields among multiple Entry items."""
all_fields: set[BaseField] = set()
@@ -0,0 +1,551 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: MIT
from collections.abc import Callable
from pathlib import Path
from typing import override
import structlog
import ujson
from sqlalchemy import Engine, and_, delete, select, text, update
from sqlalchemy.orm import Session
from tagstudio.core.constants import IGNORE_NAME, TAG_ARCHIVED, TS_FOLDER_NAME
from tagstudio.core.library.alchemy import default_color_groups
from tagstudio.core.library.alchemy.constants import (
DB_VERSION,
DB_VERSION_CURRENT_KEY,
DB_VERSION_INITIAL_KEY,
DEFAULT_FIELD_TEMPLATES,
)
from tagstudio.core.library.alchemy.fields import LEGACY_FIELD_MAP, DatetimeField, TextField
from tagstudio.core.library.alchemy.joins import TagParent
from tagstudio.core.library.alchemy.models import Tag, TagColorGroup, Version
from tagstudio.core.library.ignore import migrate_ext_list
from tagstudio.core.utils.types import unwrap
from tagstudio.qt.translations import Translations
logger = structlog.get_logger(__name__)
class MigrationError(Exception):
pass
class DBMigration:
version: int = None # pyright: ignore[reportAssignmentType]
initial_version: int | None = None
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log: Callable[[str], str]):
raise NotImplementedError
class DBMigrations:
def __init__(self, library_dir: Path, engine: Engine) -> None:
from tagstudio.core.library.alchemy.library import Library
self.library_dir = library_dir
self.engine = engine
# Don't check DB version when creating new library
self.loaded_db_version = Library._get_version(engine, DB_VERSION_CURRENT_KEY)
self.initial_db_version = Library._get_version(engine, DB_VERSION_INITIAL_KEY)
# ======================== Library Database Version Checking =======================
# DB_VERSION 6 is the first supported SQLite DB version.
# If the DB_VERSION is >= 100, that means it's a compound major + minor version.
# - Dividing by 100 and flooring gives the major (breaking changes) version.
# - If a DB has major version higher than the current program, don't load it.
# - If only the minor version is higher, it's still allowed to load.
if self.loaded_db_version < 6 or (
self.loaded_db_version >= 100 and self.loaded_db_version // 100 > DB_VERSION // 100
):
mismatch_text = Translations["status.library_version_mismatch"]
found_text = Translations["status.library_version_found"]
expected_text = Translations["status.library_version_expected"]
raise MigrationError(
f"{mismatch_text}\n"
f"{found_text} v{self.loaded_db_version}, "
f"{expected_text} v{DB_VERSION}"
)
logger.info(
f"[Library][Migration] Starting with library DB version: {self.loaded_db_version}"
)
@property
def required(self) -> bool:
return self.loaded_db_version < DB_VERSION
def run(self):
# migrate DB step by step from one version to the next
# (migration_method, db_version, initial_db_version)
migrations: list[type[DBMigration]] = [
MigrationTo7, # changes: value_type, tags
MigrationTo8, # changes: tag_colors
MigrationTo9, # changes: entries
MigrationTo100, # changes: tag_parents
MigrationTo101, # changes: versions
MigrationTo102, # changes: tag_parents
MigrationTo103, # changes: tags
MigrationTo104, # changes: deletes preferences
MigrationTo200, # changes: field tables
MigrationTo201, # changes: field tables
MigrationTo202, # changes: tag_parents
MigrationTo300, # changes: deletes folders
]
with Session(self.engine) as session:
for migration in migrations:
if self.loaded_db_version < migration.version and (
migration.initial_version is None
or self.initial_db_version < migration.initial_version
):
logger.info(f"[Library][Migration][{migration.version}] Starting DB Migration")
# any error causes transaction to rollback
migration.run(
session,
self.library_dir,
lambda msg, v=migration.version: f"[Library][Migration][{v}] {msg}",
)
self.loaded_db_version = migration.version
try:
self.__set_version(session, DB_VERSION_CURRENT_KEY, migration.version)
except Exception as e:
logger.info(
f"[Library][Migration][{migration.version}] "
"Couldn't update version, continuing without commit",
error=e,
)
session.flush()
else:
session.commit()
logger.info(f"[Library][Migration][{migration.version}] Completed DB Migration")
assert self.loaded_db_version >= DB_VERSION, (
"Ran all migrations, but the DB is still not on the newest version"
)
logger.info(f"[Library][Migration] Library migrated to DB version {DB_VERSION}")
def __set_version(self, session: Session, key: str, value: int) -> None:
"""Set a version value to the DB.
Args:
session(Session): The SQLAlchemy DB Session to use.
key(str): The key for the name of the version type to set.
value(int): The version value to set.
"""
# Insert if key has no value yet, otherwise update the value
session.merge(Version(key=key, value=value))
class MigrationTo7(DBMigration):
version = 7
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB from DB_VERSION 6 to 7."""
logger.info(fmt_log("Applying patches to DB_VERSION: 6 library..."))
# Repair tags that may have a disambiguation_id pointing towards a deleted tag.
# TODO: combine into single sql statement
all_tag_ids = session.scalars(text("SELECT DISTINCT id FROM tags")).all()
disam_stmt = (
update(Tag)
.where(Tag.disambiguation_id.not_in(all_tag_ids))
.values(disambiguation_id=None)
)
session.execute(disam_stmt)
session.flush()
class MigrationTo8(DBMigration):
version = 8
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB from DB_VERSION 7 to 8."""
# Add the missing color_border column to the TagColorGroups table.
session.execute(
text("ALTER TABLE tag_colors ADD COLUMN color_border BOOLEAN DEFAULT FALSE NOT NULL")
)
session.flush()
logger.info(fmt_log("Added color_border column to tag_colors table"))
# collect new default tag colors
tag_colors: list[TagColorGroup] = [
color
for color in default_color_groups.shades()
if color.slug in ["burgundy", "dark-teal", "dark_lavender"]
]
# Add any new default colors introduced in DB_VERSION 8
for color in tag_colors:
session.add(color)
session.flush()
logger.info(
fmt_log("Migrated tag colors to DB_VERSION 8+"),
color_name=tag_colors,
)
# Update Neon colors to use the the color_border property
for color in default_color_groups.neon():
neon_stmt = (
update(TagColorGroup)
.where(
and_(
TagColorGroup.namespace == color.namespace,
TagColorGroup.slug == color.slug,
)
)
.values(
slug=color.slug,
namespace=color.namespace,
name=color.name,
primary=color.primary,
secondary=color.secondary,
color_border=color.color_border,
)
)
session.execute(neon_stmt)
session.flush()
class MigrationTo9(DBMigration):
version = 9
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB from DB_VERSION 8 to 9."""
# Apply database schema changes
add_filename_column = text(
"ALTER TABLE entries ADD COLUMN filename TEXT NOT NULL DEFAULT ''"
)
session.execute(add_filename_column)
session.flush()
logger.info(fmt_log("Added filename column to entries table"))
# Populate the new filename column.
from tagstudio.core.library.alchemy.library import Library
for entry in Library._all_entries(session):
entry.filename = entry.path.name
session.merge(entry)
session.flush()
logger.info(fmt_log("Populated filename column in entries table"))
class MigrationTo100(DBMigration):
version = 100
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB to DB_VERSION 100."""
# Repair parent-child tag relationships that are the wrong way around.
stmt = update(TagParent).values(
parent_id=TagParent.child_id,
child_id=TagParent.parent_id,
)
session.execute(stmt)
session.flush()
logger.info(fmt_log("Refactored TagParent table"))
class MigrationTo101(DBMigration):
version = 101
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB to DB_VERSION 101."""
# Create versions table
session.execute(
text("""
CREATE TABLE versions (
"key" VARCHAR NOT NULL PRIMARY KEY,
value INTEGER NOT NULL
)
""")
)
session.flush()
# Ensure version rows are present
session.add(Version(key=DB_VERSION_INITIAL_KEY, value=100))
session.flush()
logger.info(fmt_log("Created versions table"))
class MigrationTo102(DBMigration):
version = 102
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB to DB_VERSION 102."""
# delete TagParents with a dangling parent reference
stmt = delete(TagParent).where(TagParent.parent_id.not_in(select(Tag.id).distinct()))
session.execute(stmt)
session.flush()
logger.info(fmt_log("Verified TagParent table data"))
class MigrationTo103(DBMigration):
version = 103
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB from DB_VERSION 102 to 103."""
# add the new hidden column for tags
session.execute(text("ALTER TABLE tags ADD COLUMN is_hidden BOOLEAN NOT NULL DEFAULT 0"))
session.flush()
logger.info(fmt_log("Added is_hidden column to tags table"))
# mark the "Archived" tag as hidden
session.query(Tag).filter(Tag.id == TAG_ARCHIVED).update({"is_hidden": True})
session.flush()
logger.info(fmt_log("Updated archived tag to be hidden"))
class MigrationTo104(DBMigration):
version = 104
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB from DB_VERSION 103 to 104."""
# Convert file extension list to ts_ignore file, if a .ts_ignore file does not exist
cls.__migrate_sql_to_ts_ignore(session, library_dir)
session.execute(text("DROP TABLE preferences"))
session.flush()
@classmethod
def __migrate_sql_to_ts_ignore(cls, session: Session, library_dir: Path):
# Do not continue if existing '.ts_ignore' file is found
ts_ignore = library_dir / TS_FOLDER_NAME / IGNORE_NAME
if Path(ts_ignore).exists():
return
# Load legacy extension data
extensions: list[str] = ujson.loads(
unwrap(
session.scalar(text("SELECT value FROM preferences WHERE key = 'EXTENSION_LIST'"))
)
)
is_exclude_list: bool = unwrap(
session.scalar(text("SELECT value FROM preferences WHERE key = 'IS_EXCLUDE_LIST'"))
)
with open(ts_ignore, "w") as f:
f.write(migrate_ext_list(extensions, is_exclude_list))
class MigrationTo200(DBMigration):
version = 200
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB to DB_VERSION 200."""
# Drop unused 'boolean_fields' and 'value_type' tables
logger.info(fmt_log("Dropping boolean_fields and value_type tables..."))
session.execute(text("DROP TABLE boolean_fields"))
session.execute(text("DROP TABLE value_type"))
# Add 'name' column to text_fields and datetime_fields tables
logger.info(fmt_log("Adding name columns to field tables..."))
stmt = text('ALTER TABLE text_fields ADD COLUMN name VARCHAR DEFAULT ""')
session.execute(stmt)
stmt = text('ALTER TABLE datetime_fields ADD COLUMN name VARCHAR DEFAULT ""')
session.execute(stmt)
# Drop unnecessary 'position' columns
logger.info(fmt_log("Dropping position columns to field tables..."))
session.execute(text("ALTER TABLE datetime_fields DROP COLUMN position"))
session.execute(text("ALTER TABLE text_fields DROP COLUMN position"))
# Add 'is_multiline' column to text_fields table
logger.info(fmt_log("Adding is_multiline column to text_fields..."))
stmt = text("ALTER TABLE text_fields ADD COLUMN is_multiline BOOLEAN NOT NULL DEFAULT 0")
session.execute(stmt)
session.flush()
# Move values from old `type_key` columns into new `name` columns
logger.info(fmt_log("Moving values from type_key columns to name..."))
session.execute(text("UPDATE text_fields SET name = type_key"))
session.execute(text("UPDATE datetime_fields SET name = type_key"))
session.flush()
# Change `name` values to title case
logger.info(fmt_log("Normalizing TextField names..."))
for text_field in session.execute(select(TextField)).scalars():
# NOTE: The only exception to the "Title Case" conversion is the "URL" field.
text_field.name = text_field.name.title().replace("Url", "URL").replace("_", " ")
logger.info(fmt_log("Normalizing DatetimeField names..."))
for datetime_field in session.execute(select(DatetimeField)).scalars():
datetime_field.name = datetime_field.name.title().replace("_", " ")
session.flush()
# Add correct `is_multiline` values to text_fields table
logger.info(fmt_log("Updating is_multiline for legacy TEXT_BOXes..."))
text_boxes = [
x.get("name") for x in LEGACY_FIELD_MAP.values() if x.get("is_multiline") is True
]
update_stmt = (
update(TextField).where(TextField.name.in_(text_boxes)).values(is_multiline=True)
)
session.execute(update_stmt)
session.flush()
# Repair legacy "Description" fields to use is_multiline = True
logger.info(fmt_log("Repairing legacy Description fields..."))
desc_stmt = (
update(TextField)
.where(TextField.name == "Description" and TextField.is_multiline == False) # noqa: E712
.values(is_multiline=True)
)
session.execute(desc_stmt)
# Repair legacy "Comments" fields to use is_multiline = True
logger.info(fmt_log("Repairing legacy Comment fields..."))
comm_stmt = (
update(TextField)
.where(TextField.name == "Comments" and TextField.is_multiline == False) # noqa: E712
.values(is_multiline=True)
)
session.execute(comm_stmt)
# Add default field templates
logger.info(fmt_log("Adding default field templates..."))
for template in DEFAULT_FIELD_TEMPLATES:
session.add(template)
session.flush()
# DB indices for improved performance
session.execute(
text("CREATE INDEX IF NOT EXISTS idx_tags_name_shorthand ON tags (name, shorthand)")
)
session.execute(
text("CREATE INDEX IF NOT EXISTS idx_tag_parents_child_id ON tag_parents (child_id)")
)
session.execute(
text("CREATE INDEX IF NOT EXISTS idx_tag_entries_entry_id ON tag_entries (entry_id)")
)
class MigrationTo201(DBMigration):
version = 201
initial_version = 200
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB to DB_VERSION 201."""
create_text_fields_table = text("""
CREATE TABLE text_fields_new (
id INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT,
name VARCHAR NOT NULL,
entry_id INTEGER NOT NULL,
value VARCHAR,
is_multiline BOOLEAN NOT NULL,
FOREIGN KEY(entry_id) REFERENCES entries (id)
)
""")
create_datetime_fields_table = text("""
CREATE TABLE datetime_fields_new (
id INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT,
name VARCHAR NOT NULL,
entry_id INTEGER NOT NULL,
value VARCHAR,
FOREIGN KEY(entry_id) REFERENCES entries (id)
)
""")
logger.info(fmt_log("Dropping type_key from text_fields table..."))
session.execute(create_text_fields_table)
session.flush()
session.execute(
text("""
INSERT INTO text_fields_new (id, name, entry_id, value, is_multiline)
SELECT id, name, entry_id, value, is_multiline
FROM text_fields
""")
)
session.execute(text("DROP TABLE text_fields"))
session.execute(text("ALTER TABLE text_fields_new RENAME TO text_fields"))
logger.info(fmt_log("Dropping type_key from datetime_fields table..."))
session.execute(create_datetime_fields_table)
session.flush()
session.execute(
text("""
INSERT INTO datetime_fields_new (id, name, entry_id, value)
SELECT id, name, entry_id, value
FROM datetime_fields
""")
)
session.execute(text("DROP TABLE datetime_fields"))
session.execute(text("ALTER TABLE datetime_fields_new RENAME TO datetime_fields"))
session.flush()
class MigrationTo202(DBMigration):
version = 202
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
"""Migrate DB to DB_VERSION 202."""
stmt = delete(TagParent).where(TagParent.child_id.not_in(select(Tag.id).distinct()))
session.execute(stmt)
session.flush()
logger.info(fmt_log("Verified TagParent table data"))
class MigrationTo300(DBMigration):
version = 300
@override
@classmethod
def run(cls, session: Session, library_dir: Path, fmt_log):
## remove folder_id column from entries table
# create new table in the desired scheme (without folder_id column)
session.execute(
text("""
CREATE TABLE entries_new (
id INTEGER NOT NULL,
path VARCHAR NOT NULL,
suffix VARCHAR NOT NULL,
date_created DATETIME,
date_modified DATETIME,
date_added DATETIME,
filename TEXT NOT NULL DEFAULT '',
PRIMARY KEY (id),
UNIQUE (path)
)
""")
)
session.flush()
# transfer data to new table
session.execute(
text("""
INSERT INTO entries_new (id, path, suffix, date_created, date_modified, date_added,
filename)
SELECT id, path, suffix, date_created, date_modified, date_added, filename
FROM entries
""")
)
# delete old table
session.execute(text("DROP TABLE entries"))
# rename new table to old table
session.execute(text("ALTER TABLE entries_new RENAME TO entries"))
session.flush()
## drop table "folders"
session.execute(text("DROP TABLE folders"))
session.flush()
-1
View File
@@ -305,7 +305,6 @@ class MediaCategories:
".nef",
".nrw",
".orf",
".r3d",
".raf",
".raw",
".rw2",
+6 -2
View File
@@ -6,7 +6,7 @@ from pathlib import Path
import ffmpeg
from tagstudio.renderers.vendored.probe import probe
from tagstudio.qt.previews.vendored.probe import probe
def is_readable_video(filepath: Path | str):
@@ -23,7 +23,11 @@ def is_readable_video(filepath: Path | str):
return False
for stream in result["streams"]:
# DRM check
if stream.get("codec_tag_string") in ["drma", "drms", "drmi"]:
if stream.get("codec_tag_string") in [
"drma",
"drms",
"drmi",
]:
return False
except ffmpeg.Error:
return False
File diff suppressed because it is too large Load Diff
@@ -1,7 +1,11 @@
# SPDX-FileCopyrightText: (c) 2017 The Blender Foundation
#!/usr/bin/env python3
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
# <pep8 compliant>
## This file is a modified script that gets the thumbnail data stored in a blend file
@@ -1,4 +1,4 @@
# SPDX-FileCopyrightText: (c) 2022 Karl Kroening (kkroening)
# SPDX-FileCopyrightText: (c) 2022 Karl Kroening (kkroening)
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
# Vendored from ffmpeg-python and ffmpeg-python PR#790 by amamic1803
@@ -1,4 +1,4 @@
# SPDX-FileCopyrightText: (c) 2022 James Robert (jiaaro), http://jiaaro.com
# SPDX-FileCopyrightText: (c) 2022 James Robert (jiaaro)
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
# Vendored from pydub
@@ -43,7 +43,7 @@ from pydub.utils import (
)
from tagstudio.core.utils.silent_subprocess import silent_popen
from tagstudio.renderers.vendored.pydub.utils import _mediainfo_json
from tagstudio.qt.previews.vendored.pydub.utils import _mediainfo_json
basestring = str
xrange = range
@@ -256,7 +256,7 @@ class _AudioSegment:
self.sample_width = 4
self.frame_width = self.channels * self.sample_width
super().__init__(*args, **kwargs)
super(_AudioSegment, self).__init__(*args, **kwargs)
@property
def raw_data(self):
@@ -1,4 +1,4 @@
# SPDX-FileCopyrightText: (c) 2011 James Robert (jiaaro), http://jiaaro.com
# SPDX-FileCopyrightText: (c) 2011 James Robert, http://jiaaro.com
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
# Vendored from pydub
+1 -1
View File
@@ -837,7 +837,7 @@ class QtDriver(DriverMixin, QObject):
logger.info("Backing Up Library...")
self.main_window.status_bar.showMessage(Translations["status.library_backup_in_progress"])
start_time = time.time()
target_path = self.lib.save_library_backup_to_disk()
target_path = Library.save_library_backup_to_disk(unwrap(self.lib.library_dir))
end_time = time.time()
self.main_window.status_bar.showMessage(
Translations.format(
-158
View File
@@ -1,158 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import tarfile
import zipfile
from io import BytesIO
from pathlib import Path
from typing import Literal
import py7zr
import py7zr.io
import rarfile
import structlog
from PIL import Image
from tagstudio.core.media_types import MediaCategories
from tagstudio.core.utils.types import unwrap
from tagstudio.renderers.raster_image import image_from_bytes
logger = structlog.get_logger(__name__)
type Archive = zipfile.ZipFile | rarfile.RarFile | SevenZipFile | TarFile
class SevenZipFile(py7zr.SevenZipFile):
"""Wrapper around py7zr.SevenZipFile to mimic zipfile.ZipFile's API."""
def __init__(self, filepath: Path, mode: Literal["r"]) -> None:
super().__init__(filepath, mode)
def read(self, name: str) -> bytes:
# SevenZipFile must be reset after every extraction
# See https://py7zr.readthedocs.io/en/stable/api.html#py7zr.SevenZipFile.extract
self.reset()
factory = py7zr.io.BytesIOFactory(limit=10485760) # 10 MiB
self.extract(targets=[name], factory=factory)
return factory.get(name).read()
class TarFile:
"""Wrapper around tarfile.TarFile to mimic zipfile.ZipFile's API."""
def __init__(self, filepath: Path, mode: Literal["r"]) -> None:
self.tar: tarfile.TarFile
self.filepath = filepath
self.mode: Literal["r"] = mode
def namelist(self) -> list[str]:
return self.tar.getnames()
def read(self, name: str) -> bytes:
return unwrap(self.tar.extractfile(name)).read()
def __enter__(self) -> "TarFile":
self.tar = tarfile.open(name=self.filepath, mode=self.mode).__enter__()
return self
def __exit__(self, *args) -> None: # pyright: ignore[reportUnknownParameterType, reportMissingParameterType]
self.tar.__exit__(*args)
def open_archive(filepath: Path, ext: str = "") -> Archive:
"""Open an archive with its corresponding archiver.
Args:
filepath (Path): The path to the archive.
ext (str): The file extension.
Returns:
Archive: The opened archive.
"""
archiver: type[Archive] = zipfile.ZipFile
if ext in {".7z", ".cb7", ".s7z"}:
archiver = SevenZipFile
elif ext in {".cbr", ".rar"}:
archiver = rarfile.RarFile
elif ext in {".cbt", ".tar", ".tgz"}:
archiver = TarFile
return archiver(filepath, "r")
def first_image_in_archive(archive: Archive) -> Image.Image | None:
"""Find and extract the first renderable image in the archive.
Args:
archive (Archive): The current archive.
Returns:
Image: The first renderable image in the archive.
"""
for file_name in archive.namelist(): # pyright: ignore[reportUnknownVariableType]
ext = Path(file_name).suffix
if MediaCategories.IMAGE_RASTER_TYPES.contains(ext):
image_data = archive.read(file_name) # pyright: ignore[reportUnknownVariableType]
return image_from_bytes(BytesIO(image_data))
return None
def archive_thumb(
filepath: Path, ext: str = "", image_names: list[Path] | list[str] | None = None
) -> Image.Image | None:
"""Extract an embedded preview image from an archive.
Args:
filepath (Path): The path to the archive.
ext (str): The file extension. Used to help determine more specific archive type.
image_names: (list[Path] | list[str] | None): List of embedded image names to search for.
Returns:
Image: The first image found in the archive.
"""
try:
with open_archive(filepath, ext) as archive:
# If no list of image names to search for was provided, default to the first image.
if not image_names:
return first_image_in_archive(archive)
for image_name in image_names:
if image_name in archive.namelist():
file_data = archive.read(str(image_name)) # pyright: ignore[reportUnknownVariableType]
return image_from_bytes(BytesIO(file_data))
except Exception as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return None
def apple_embedded_thumb(filepath: Path) -> Image.Image | None:
"""Extract and render an apple embedded thumbnail (iWork, Apple Creative Studio)."""
image_names: list[str] = [
"preview.jpg",
"QuickLook/Preview.heic",
"QuickLook/Thumbnail.jpg",
"QuickLook/Thumbnail.heic",
"QuickLook/Thumbnail.webp",
"QuickLook/Icon.webp",
]
return archive_thumb(filepath, image_names=image_names)
def krita_thumb(filepath: Path) -> Image.Image | None:
"""Extract and render a thumbnail for a Krita file."""
image_names = ["preview.png"]
return archive_thumb(filepath, image_names=image_names)
def open_doc_thumb(filepath: Path) -> Image.Image | None:
"""Extract and render a thumbnail for an OpenDocument file."""
image_names = ["Thumbnails/thumbnail.png"]
return archive_thumb(filepath, image_names=image_names)
def powerpoint_thumb(filepath: Path) -> Image.Image | None:
"""Extract and render a thumbnail for a Microsoft PowerPoint file."""
image_names = ["docProps/thumbnail.jpeg"]
return archive_thumb(filepath, image_names=image_names)
-149
View File
@@ -1,149 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import math
from io import BytesIO
from pathlib import Path
from warnings import catch_warnings
import numpy as np
import structlog
from mutagen import flac, id3, mp4
from mutagen._util import MutagenError
from PIL import Image, ImageDraw
from tagstudio.renderers.vendored.pydub.audio_segment import (
_AudioSegment as AudioSegment, # pyright: ignore[reportPrivateUsage]
)
logger = structlog.get_logger(__name__)
def audio_album_thumb(filepath: Path, ext: str) -> Image.Image | None:
"""Return an album cover thumb from an audio file if a cover is present.
Args:
filepath (Path): The path of the file.
ext (str): The file extension (with leading ".").
"""
image: Image.Image | None = None
try:
if not filepath.is_file():
raise FileNotFoundError
artwork = None
if ext in [".mp3"]:
id3_tags: id3.ID3 = id3.ID3(filepath)
id3_covers: list = id3_tags.getall("APIC") # pyright: ignore[reportUnknownVariableType]
if id3_covers:
artwork = Image.open(BytesIO(id3_covers[0].data))
elif ext in [".flac"]:
flac_tags: flac.FLAC = flac.FLAC(filepath)
flac_covers: list = flac_tags.pictures # pyright: ignore[reportUnknownVariableType]
if flac_covers:
artwork = Image.open(BytesIO(flac_covers[0].data))
elif ext in [".mp4", ".m4a", ".aac"]:
mp4_tags: mp4.MP4 = mp4.MP4(filepath)
mp4_covers: list | None = mp4_tags.get("covr") # pyright: ignore[reportUnknownVariableType]
if mp4_covers:
artwork = Image.open(BytesIO(mp4_covers[0]))
if artwork:
image = artwork
except (
FileNotFoundError,
id3.ID3NoHeaderError, # pyright: ignore[reportPrivateImportUsage]
mp4.MP4MetadataError,
mp4.MP4StreamInfoError,
MutagenError,
) as e:
logger.error("Couldn't read album artwork", path=filepath, error=type(e).__name__)
return image
def audio_waveform_thumb(
filepath: Path, ext: str, size: int, pixel_ratio: float
) -> Image.Image | None:
"""Render a waveform image from an audio file.
Args:
filepath (Path): The path of the file.
ext (str): The file extension (with leading ".").
size (tuple[int,int]): The size of the thumbnail.
pixel_ratio (float): The screen pixel ratio.
"""
# BASE_SCALE used for drawing on a larger image and resampling down
# to provide an antialiased effect.
base_scale: int = 2
samples_per_bar: int = 3
size_scaled: int = size * base_scale
allow_small_min: bool = False
im: Image.Image | None = None
try:
bar_count: int = min(math.floor((size // pixel_ratio) / 5), 64)
audio = AudioSegment.from_file(filepath, ext[1:]) # pyright: ignore[reportUnknownVariableType]
data = np.frombuffer(buffer=audio._data, dtype=np.int16)
data_indices = np.linspace(1, len(data), num=bar_count * samples_per_bar)
bar_margin: float = ((size_scaled / (bar_count * 3)) * base_scale) / 2
line_width: float = ((size_scaled - bar_margin) / (bar_count * 3)) * base_scale
bar_height: float = (size_scaled) - (size_scaled // bar_margin)
count: int = 0
maximum_item: int = 0
max_array: list[int] = []
highest_line: int = 0
for i in range(-1, len(data_indices)):
d = data[math.ceil(data_indices[i]) - 1]
if count < samples_per_bar:
count = count + 1
with catch_warnings(record=True):
if abs(d) > maximum_item:
maximum_item = int(abs(d))
else:
max_array.append(maximum_item)
if maximum_item > highest_line:
highest_line = maximum_item
maximum_item = 0
count = 1
line_ratio = max(highest_line / bar_height, 1)
im = Image.new("RGB", (size_scaled, size_scaled), color="#000000")
draw = ImageDraw.Draw(im)
current_x = bar_margin
for item in max_array:
item_height = item / line_ratio
# If small minimums are not allowed, raise all values
# smaller than the line width to the same value.
if not allow_small_min:
item_height = max(item_height, line_width)
current_y = (bar_height - item_height + (size_scaled // bar_margin)) // 2
draw.rounded_rectangle(
(
current_x,
current_y,
(current_x + line_width),
(current_y + item_height),
),
radius=100 * base_scale,
fill=("#FF0000"),
outline=("#FFFF00"),
width=max(math.ceil(line_width / 6), base_scale),
)
current_x = current_x + line_width + bar_margin
im.resize((size, size), Image.Resampling.BILINEAR)
except Exception as e:
logger.error("Couldn't render waveform", path=filepath.name, error=type(e).__name__)
return im
-43
View File
@@ -1,43 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
from pathlib import Path
import structlog
from PIL import Image
from PySide6.QtCore import Qt
from PySide6.QtGui import QGuiApplication
from tagstudio.renderers.vendored.blender_renderer import (
blend_thumb, # pyright: ignore[reportUnknownVariableType]
)
logger = structlog.get_logger(__name__)
def blender_thumb(filepath: Path) -> Image.Image | None:
"""Get an emended thumbnail from a Blender file, if a thumbnail is present.
Args:
filepath (Path): The path of the file.
"""
bg_color: str = (
"#1e1e1e"
if QGuiApplication.styleHints().colorScheme() is Qt.ColorScheme.Dark
else "#FFFFFF"
)
im: Image.Image | None = None
try:
if (blend_image := blend_thumb(str(filepath))) is not None:
bg = Image.new("RGB", blend_image.size, color=bg_color)
bg.paste(blend_image, mask=blend_image.getchannel(3))
im = bg
else:
logger.info(
f"[ThumbRenderer][BLENDER][INFO] {filepath.name} "
"Doesn't have an embedded thumbnail."
)
except Exception as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
-80
View File
@@ -1,80 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import base64
import sqlite3
import struct
import xml.etree.ElementTree as ET
from io import BytesIO
from pathlib import Path
import structlog
from PIL import Image
logger = structlog.get_logger(__name__)
def clip_studio_thumb(filepath: Path) -> Image.Image | None:
"""Extract the thumbnail from the SQLite database embedded in a .clip file.
Args:
filepath (Path): The path of the .clip file.
Returns:
Image: The embedded thumbnail, if extractable.
"""
im: Image.Image | None = None
try:
with open(filepath, "rb") as f:
blob = f.read()
sqlite_index = blob.find(b"SQLite format 3")
if sqlite_index == -1:
return im
with sqlite3.connect(":memory:") as conn:
conn.deserialize(blob[sqlite_index:])
thumbnail = conn.execute("SELECT ImageData FROM CanvasPreview").fetchone()
if thumbnail:
im = Image.open(BytesIO(thumbnail[0]))
conn.close()
except Exception as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
def pdn_thumb(filepath: Path) -> Image.Image | None:
"""Extract the base64-encoded thumbnail from a .pdn file header.
Args:
filepath (Path): The path of the .pdn file.
Returns:
Image: the decoded PNG thumbnail or None by default.
"""
im: Image.Image | None = None
with open(filepath, "rb") as f:
try:
# First 4 bytes are the magic number
if f.read(4) != b"PDN3":
return im
# Header length is a little-endian 24-bit int
header_size = struct.unpack("<i", f.read(3) + b"\x00")[0]
thumb_element = ET.fromstring(f.read(header_size)).find("./*thumb")
if thumb_element is None:
return im
encoded_png = thumb_element.get("png")
if encoded_png:
decoded_png = base64.b64decode(encoded_png)
im = Image.open(BytesIO(decoded_png))
if im.mode == "RGBA":
new_bg = Image.new("RGB", im.size, color="#1e1e1e")
new_bg.paste(im, mask=im.getchannel(3))
im = new_bg
except Exception as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
-73
View File
@@ -1,73 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import xml.etree.ElementTree as ET
from io import BytesIO
from pathlib import Path
from xml.etree.ElementTree import Element
import structlog
from PIL import Image
from tagstudio.core.media_types import MediaCategories
from tagstudio.core.utils.types import unwrap
from tagstudio.renderers.archive import Archive, first_image_in_archive, open_archive
from tagstudio.renderers.raster_image import image_from_bytes
logger = structlog.get_logger(__name__)
def epub_thumb(filepath: Path, ext: str) -> Image.Image | None:
"""Extracts the cover specified by ComicInfo.xml or first image found in the ePub file.
Args:
filepath (Path): The path to the ePub file.
ext (str): The file extension.
Returns:
Image: The cover specified in ComicInfo.xml,
the first image found in the ePub file, or None by default.
"""
im: Image.Image | None = None
try:
with open_archive(filepath, ext) as archive:
if "ComicInfo.xml" in archive.namelist():
comic_info = ET.fromstring(archive.read("ComicInfo.xml"))
im = _cover_from_comic_info(archive, comic_info, "FrontCover")
if not im:
im = _cover_from_comic_info(archive, comic_info, "InnerCover")
if not im:
im = first_image_in_archive(archive)
except Exception as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
def _cover_from_comic_info(
archive: Archive, comic_info: Element, cover_type: str
) -> Image.Image | None:
"""Extract the cover specified in ComicInfo.xml.
Args:
archive (Archive): The current ePub file.
comic_info (Element): The parsed ComicInfo.xml.
cover_type (str): The type of cover to load.
Returns:
Image: The cover specified in ComicInfo.xml.
"""
im: Image.Image | None = None
cover = comic_info.find(f"./*Page[@Type='{cover_type}']")
if cover is not None:
pages = [f for f in archive.namelist() if f != "ComicInfo.xml"] # pyright: ignore[reportUnknownVariableType]
page_name = pages[int(unwrap(cover.get("Image")))] # pyright: ignore[reportUnknownVariableType]
ext = Path(page_name).suffix
if MediaCategories.IMAGE_RASTER_TYPES.contains(ext):
image_data = archive.read(page_name) # pyright: ignore[reportUnknownVariableType]
im = image_from_bytes(BytesIO(image_data))
return im
-108
View File
@@ -1,108 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import math
from pathlib import Path
from typing import cast
import numpy as np
import structlog
from PIL import Image, ImageDraw, ImageFont
from tagstudio.core.constants import FONT_SAMPLE_SIZES, FONT_SAMPLE_TEXT
from tagstudio.qt.helpers.color_overlay import auto_theme_overlay
from tagstudio.qt.helpers.text_wrapper import wrap_full_text
logger = structlog.get_logger(__name__)
def font_small_thumb(filepath: Path, size: int) -> Image.Image | None:
"""Render a small font preview ("Aa") thumbnail from a font file.
Args:
filepath (Path): The path of the file.
size (tuple[int,int]): The size of the thumbnail.
"""
im: Image.Image | None = None
try:
bg = Image.new("RGB", (size, size), color="#000000")
raw = Image.new("RGB", (size * 3, size * 3), color="#000000")
draw = ImageDraw.Draw(raw)
font = ImageFont.truetype(filepath, size=size)
# NOTE: While a stroke effect is desired, the text
# method only allows for outer strokes, which looks
# a bit weird when rendering fonts.
draw.text(
(size // 8, size // 8),
"Aa",
font=font,
fill="#FF0000",
# stroke_width=math.ceil(size / 96),
# stroke_fill="#FFFF00",
)
# NOTE: Change to getchannel(1) if using an outline.
data = np.asarray(raw.getchannel(0))
m, n = data.shape[:2]
col: np.ndarray = cast(np.ndarray, data.any(0))
row: np.ndarray = cast(np.ndarray, data.any(1))
cropped_data = np.asarray(raw)[
row.argmax() : m - row[::-1].argmax(),
col.argmax() : n - col[::-1].argmax(),
]
cropped_im: Image.Image = Image.fromarray(cropped_data, "RGB")
margin: int = math.ceil(size // 16)
orig_x, orig_y = cropped_im.size
new_x, new_y = (size, size)
if orig_x > orig_y:
new_x = size
new_y = math.ceil(size * (orig_y / orig_x))
elif orig_y > orig_x:
new_y = size
new_x = math.ceil(size * (orig_x / orig_y))
cropped_im = cropped_im.resize(
size=(new_x - (margin * 2), new_y - (margin * 2)),
resample=Image.Resampling.BILINEAR,
)
bg.paste(
cropped_im,
box=(margin, margin + ((size - new_y) // 2)),
)
im = bg
except OSError as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
def font_full_preview(filepath: Path, size: int) -> Image.Image | None:
"""Render a large font preview ("Alphabet") thumbnail from a font file.
Args:
filepath (Path): The path of the file.
size (tuple[int,int]): The size of the thumbnail.
"""
# Scale the sample font sizes to the preview image
# resolution,assuming the sizes are tuned for 256px.
im: Image.Image | None = None
try:
scaled_sizes: list[int] = [math.floor(x * (size / 256)) for x in FONT_SAMPLE_SIZES]
bg = Image.new("RGBA", (size, size), color="#00000000")
draw = ImageDraw.Draw(bg)
lines_of_padding = 2
y_offset = 0.0
for font_size in scaled_sizes:
font = ImageFont.truetype(filepath, size=font_size)
text_wrapped: str = wrap_full_text(FONT_SAMPLE_TEXT, font=font, width=size, draw=draw)
draw.multiline_text((0, y_offset), text_wrapped, font=font)
y_offset += (len(text_wrapped.split("\n")) + lines_of_padding) * draw.textbbox(
(0, 0), "A", font=font
)[-1]
im = auto_theme_overlay(bg, use_alpha=False)
except OSError as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
-58
View File
@@ -1,58 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import os
import struct
import xml.etree.ElementTree as ET
import zlib
from pathlib import Path
import structlog
from PIL import Image
from tagstudio.core.utils.types import unwrap
logger = structlog.get_logger(__name__)
def medibang_paint_thumb(filepath: Path) -> Image.Image | None:
"""Extract the thumbnail from a .mdp file.
Args:
filepath (Path): The path of the .mdp file.
Returns:
Image: The embedded thumbnail.
"""
im: Image.Image | None = None
try:
with open(filepath, "rb") as f:
magic = struct.unpack("<7sx", f.read(8))[0]
if magic != b"mdipack":
return im
bin_header = struct.unpack("<LLL", f.read(12))
xml_header = ET.fromstring(f.read(bin_header[1]))
mdibin_count = len(xml_header.findall("./*Layer")) + 1
for _ in range(mdibin_count):
pac_header = struct.unpack("<3sxLLLL48s64s", f.read(132))
if not pac_header[6].startswith(b"thumb"):
f.seek(pac_header[3], os.SEEK_CUR)
continue
thumb_element = unwrap(xml_header.find("Thumb"))
dimensions = (
int(unwrap(thumb_element.get("width"))),
int(unwrap(thumb_element.get("height"))),
)
thumb_blob = f.read(pac_header[3])
if pac_header[2] == 1:
thumb_blob = zlib.decompress(thumb_blob, bufsize=pac_header[4])
im = Image.frombytes("RGBA", dimensions, thumb_blob, "raw", "BGRA")
break
except Exception as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
-66
View File
@@ -1,66 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
# NOTE: This file contains Qt imports because Qt is used as the vector renderer in this case.
# This is NOT considered part of the Qt frontend, but is technically tangled with the Qt import.
from io import BytesIO
from pathlib import Path
import structlog
from PIL import Image
from PySide6.QtCore import QBuffer, QFile, QFileDevice, QIODeviceBase, QSizeF
from PySide6.QtGui import QImage
from PySide6.QtPdf import QPdfDocument, QPdfDocumentRenderOptions
from tagstudio.qt.helpers.image_effects import replace_transparent_pixels
logger = structlog.get_logger(__name__)
def pdf_thumb(filepath: Path, size: int, ext: str) -> Image.Image | None:
"""Render a thumbnail for a PDF or Adobe Illustrator file.
filepath (Path): The path of the file.
size (int): The size of the icon.
ext (str): The file extension.
"""
im: Image.Image | None = None
file: QFile = QFile(filepath)
success: bool = file.open(QIODeviceBase.OpenModeFlag.ReadOnly, QFileDevice.Permission.ReadUser)
if not success:
logger.error("Couldn't render thumbnail", filepath=filepath)
return im
document: QPdfDocument = QPdfDocument()
document.load(file)
file.close()
# Transform page_size in points to pixels with proper aspect ratio
page_size: QSizeF = document.pagePointSize(0)
ratio_hw: float = page_size.height() / page_size.width()
if ratio_hw >= 1:
page_size *= size / page_size.height()
else:
page_size *= size / page_size.width()
# Enlarge image for anti-aliasing
scale_factor = 2.5 if ext in {".pdf"} else 1
page_size *= scale_factor
# Render image with no anti-aliasing for speed
render_options: QPdfDocumentRenderOptions = QPdfDocumentRenderOptions()
render_options.setRenderFlags(
QPdfDocumentRenderOptions.RenderFlag.TextAliased
| QPdfDocumentRenderOptions.RenderFlag.ImageAliased
| QPdfDocumentRenderOptions.RenderFlag.PathAliased
)
# Convert QImage to PIL Image
q_image: QImage = document.render(0, page_size.toSize(), render_options)
buffer: QBuffer = QBuffer()
buffer.open(QBuffer.OpenModeFlag.ReadWrite)
try:
q_image.save(buffer, "PNG") # pyright: ignore
im = Image.open(BytesIO(buffer.buffer().data()))
finally:
buffer.close()
# Replace transparent pixels with white (otherwise Background defaults to transparent)
return replace_transparent_pixels(im)
-126
View File
@@ -1,126 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import os
from io import BytesIO
from pathlib import Path
import cv2
import numpy as np
import rawpy
import structlog
from PIL import Image, ImageOps, UnidentifiedImageError
from PIL.Image import DecompressionBombError
from pillow_heif import register_heif_opener # pyright: ignore[reportUnknownVariableType]
from rawpy import (
LibRawFileUnsupportedError, # pyright: ignore[reportPrivateImportUsage]
LibRawIOError, # pyright: ignore[reportPrivateImportUsage]
)
from tagstudio.core.utils.types import unwrap
logger = structlog.get_logger(__name__)
try:
import pillow_jxl # noqa: F401 # pyright: ignore
except ImportError as e:
logger.error('[ThumbRenderer] Could not import the "pillow_jxl" module', error=e)
register_heif_opener()
os.environ["OPENCV_IO_ENABLE_OPENEXR"] = "1"
def raster_image_thumb(filepath: Path) -> Image.Image | None:
"""Render a thumbnail for a standard image type.
Args:
filepath (Path): The path of the file.
"""
im: Image.Image | None = None
try:
with filepath.open("rb") as file:
im = image_from_bytes(BytesIO(file.read()))
except (
FileNotFoundError,
UnidentifiedImageError,
DecompressionBombError,
NotImplementedError,
) as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
def exr_image_thumb(filepath: Path) -> Image.Image | None:
"""Render a thumbnail for a EXR image type.
Args:
filepath (Path): The path of the file.
"""
im: Image.Image | None = None
try:
# Load the EXR data to an array and rotate the color space from BGRA -> RGBA
raw_array = cv2.imread(str(filepath), cv2.IMREAD_UNCHANGED)
assert raw_array is not None
raw_array[..., :3] = raw_array[..., 2::-1]
# Correct the gamma of the raw array
gamma = 2.2
array_gamma = np.power(np.clip(raw_array, 0, 1), 1 / gamma)
array = (array_gamma * 255).astype(np.uint8)
im = Image.fromarray(array, mode="RGBA")
# Paste solid background
if im.mode == "RGBA":
new_bg = Image.new("RGB", im.size, color="#1e1e1e")
new_bg.paste(im, mask=im.getchannel(3))
im = new_bg
except Exception as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
def raw_image_thumb(filepath: Path) -> Image.Image | None:
"""Render a thumbnail for a RAW image type.
Args:
filepath (Path): The path of the file.
"""
im: Image.Image | None = None
try:
with rawpy.imread(str(filepath)) as raw:
rgb = raw.postprocess(use_camera_wb=True)
im = Image.frombytes(
"RGB",
(rgb.shape[1], rgb.shape[0]),
rgb,
decoder_name="raw",
)
except (
DecompressionBombError,
LibRawIOError,
LibRawFileUnsupportedError,
) as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
def image_from_bytes(image_data: BytesIO) -> Image.Image:
"""Load a raster image and add a background if it's transparent.
Args:
image_data (BytesIO): The binary image data.
Returns:
Image.Image: The loaded raster image, with a background if needed.
"""
im: Image.Image = Image.open(image_data)
if im.mode != "RGB" and im.mode != "RGBA":
im = im.convert(mode="RGBA")
if im.mode == "RGBA":
new_bg = Image.new("RGB", im.size, color="#1e1e1e")
new_bg.paste(im, mask=im.getchannel(3))
im = new_bg
return unwrap(ImageOps.exif_transpose(im))
-30
View File
@@ -1,30 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
from pathlib import Path
import srctools
import structlog
from PIL import Image
logger = structlog.get_logger(__name__)
def vtf_thumb(filepath: Path) -> Image.Image | None:
"""Extract and render a thumbnail for VTF (Valve Texture Format) images.
Uses the srctools library for reading VTF files.
Args:
filepath (Path): The path of the file.
"""
im: Image.Image | None = None
try:
with open(filepath, "rb") as f:
vtf = srctools.VTF.read(f)
im = vtf.get(frame=0).to_PIL()
except (ValueError, FileNotFoundError) as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
-55
View File
@@ -1,55 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
from pathlib import Path
import cv2
import structlog
from PIL import Image, ImageDraw, UnidentifiedImageError
from PIL.Image import DecompressionBombError
from PySide6.QtCore import Qt
from PySide6.QtGui import QGuiApplication
from tagstudio.core.utils.encoding import detect_char_encoding
logger = structlog.get_logger(__name__)
def text_thumb(filepath: Path) -> Image.Image | None:
"""Render a thumbnail for a plaintext file.
Args:
filepath (Path): The path of the file.
"""
im: Image.Image | None = None
bg_color: str = (
"#1e1e1e"
if QGuiApplication.styleHints().colorScheme() is Qt.ColorScheme.Dark
else "#FFFFFF"
)
fg_color: str = (
"#FFFFFF"
if QGuiApplication.styleHints().colorScheme() is Qt.ColorScheme.Dark
else "#111111"
)
try:
encoding = detect_char_encoding(filepath)
with open(filepath, encoding=encoding) as text_file:
text = text_file.read(256)
bg = Image.new("RGB", (256, 256), color=bg_color)
draw = ImageDraw.Draw(bg)
draw.text((16, 16), text, fill=fg_color)
im = bg
except (
UnidentifiedImageError,
cv2.error,
DecompressionBombError,
UnicodeDecodeError,
OSError,
FileNotFoundError,
) as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im
-56
View File
@@ -1,56 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
# NOTE: This file contains Qt imports because Qt is used as the vector renderer in this case.
# This is NOT considered part of the Qt frontend, but is technically tangled with the Qt import.
from io import BytesIO
from pathlib import Path
import structlog
from PIL import (
Image,
UnidentifiedImageError,
)
from PySide6.QtCore import QBuffer, Qt
from PySide6.QtGui import QImage, QPainter
from PySide6.QtSvg import QSvgRenderer
logger = structlog.get_logger(__name__)
def vector_image_thumb(filepath: Path, size: int) -> Image.Image:
"""Render a thumbnail for a vector image, such as SVG.
Args:
filepath (Path): The path of the file.
size (tuple[int,int]): The size of the thumbnail.
"""
im: Image.Image | None = None
# Create an image to draw the svg to and a painter to do the drawing
q_image: QImage = QImage(size, size, QImage.Format.Format_ARGB32)
q_image.fill("#1e1e1e")
# Create an svg renderer, then render to the painter
svg: QSvgRenderer = QSvgRenderer(str(filepath))
if not svg.isValid():
raise UnidentifiedImageError
painter: QPainter = QPainter(q_image)
svg.setAspectRatioMode(Qt.AspectRatioMode.KeepAspectRatio)
svg.render(painter)
painter.end()
# Write the image to a buffer as png
buffer: QBuffer = QBuffer()
buffer.open(QBuffer.OpenModeFlag.ReadWrite)
q_image.save(buffer, "PNG") # pyright: ignore[reportCallIssue, reportArgumentType]
# Load the image from the buffer
im = Image.new("RGB", (size, size), color="#1e1e1e")
im.paste(Image.open(BytesIO(buffer.data().data())))
im = im.convert(mode="RGB")
buffer.close()
return im
-55
View File
@@ -1,55 +0,0 @@
# SPDX-FileCopyrightText: (c) TagStudio Contributors
# SPDX-License-Identifier: GPL-3.0-only
import math
from pathlib import Path
import cv2
import structlog
from cv2.typing import MatLike
from PIL import Image, UnidentifiedImageError
from PIL.Image import DecompressionBombError
from tagstudio.qt.helpers.file_tester import is_readable_video
logger = structlog.get_logger(__name__)
def video_thumb(filepath: Path) -> Image.Image | None:
"""Render a thumbnail for a video file.
Args:
filepath (Path): The path of the file.
"""
im: Image.Image | None = None
frame: MatLike | None = None
try:
if is_readable_video(filepath):
video = cv2.VideoCapture(str(filepath), cv2.CAP_FFMPEG)
# TODO: Move this check to is_readable_video()
if video.get(cv2.CAP_PROP_FRAME_COUNT) <= 0:
raise cv2.error("File is invalid or has 0 frames")
video.set(
cv2.CAP_PROP_POS_FRAMES,
(video.get(cv2.CAP_PROP_FRAME_COUNT) // 2),
)
# NOTE: Depending on the video format, compression, and
# frame count, seeking halfway does not work and the thumb
# must be pulled from the earliest available frame.
max_frame_seek: int = 10
for i in range(
0,
min(max_frame_seek, math.floor(video.get(cv2.CAP_PROP_FRAME_COUNT))),
):
success, frame = video.read()
if not success:
video.set(cv2.CAP_PROP_POS_FRAMES, i)
else:
break
if frame is not None:
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
im = Image.fromarray(frame)
except (UnidentifiedImageError, cv2.error, DecompressionBombError, OSError) as e:
logger.error("Couldn't render thumbnail", filepath=filepath, error=type(e).__name__)
return im