mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
feat(autocomplete): aliases, extra-networks, translations, baseline fixes
- tag aliases (optional 4-tuple) and translations (companion file) - lora and wildcard triggers; embeddings fold into tag-mode results - new GET /sdapi/v1/wildcards endpoint - LRU cap, IME safety, multi-textarea safety, modifier-key pass-through - manifest-retry and config-bridge hardening - tags-fetch: alias/translation harvesters with authenticated rate limits
This commit is contained in:
+309
-26
@@ -123,18 +123,38 @@ def fetch_with_retry(session: requests.Session, url: str, params: dict | None =
|
||||
raise
|
||||
|
||||
|
||||
# -- Auth --
|
||||
|
||||
def parse_login_key(key: str | None) -> tuple[str, str] | None:
|
||||
"""Parse a `login:api_key` string. Returns (login, api_key) or None if unset/invalid."""
|
||||
if not key or ":" not in key:
|
||||
return None
|
||||
login, api_key = key.split(":", 1)
|
||||
login = login.strip()
|
||||
api_key = api_key.strip()
|
||||
if not login or not api_key:
|
||||
return None
|
||||
return login, api_key
|
||||
|
||||
|
||||
# -- Fetchers --
|
||||
|
||||
def fetch_danbooru(min_count: int = 10, partial_path: str = "", **_kwargs) -> list:
|
||||
"""Fetch tags from Danbooru API, paginated."""
|
||||
def fetch_danbooru(min_count: int = 10, partial_path: str = "", api_key: str | None = None, **_kwargs) -> list:
|
||||
"""Fetch tags from Danbooru API, paginated.
|
||||
With `api_key` set (login:api_key) the rate limit rises from anon (~1 rps) to authenticated (~10 rps).
|
||||
"""
|
||||
auth = parse_login_key(api_key)
|
||||
sleep_sec = 0.1 if auth else 1.0
|
||||
start_page, tags = load_partial(partial_path)
|
||||
page = (start_page + 1) if start_page is not None else 1
|
||||
session = requests.Session()
|
||||
session.headers["User-Agent"] = USER_AGENT
|
||||
while True:
|
||||
url = f"https://danbooru.donmai.us/tags.json?limit=1000&page={page}&search[order]=count"
|
||||
params: dict[str, str] = {"limit": "1000", "page": str(page), "search[order]": "count"}
|
||||
if auth:
|
||||
params["login"], params["api_key"] = auth
|
||||
try:
|
||||
resp = fetch_with_retry(session, url)
|
||||
resp = fetch_with_retry(session, "https://danbooru.donmai.us/tags.json", params=params)
|
||||
except requests.RequestException as e:
|
||||
print(f" Error on page {page}: {e}", file=sys.stderr)
|
||||
break
|
||||
@@ -153,21 +173,27 @@ def fetch_danbooru(min_count: int = 10, partial_path: str = "", **_kwargs) -> li
|
||||
if page % SAVE_INTERVAL == 0:
|
||||
save_partial(partial_path, page, tags)
|
||||
page += 1
|
||||
time.sleep(0.5)
|
||||
time.sleep(sleep_sec)
|
||||
tags.sort(key=lambda t: t[2], reverse=True)
|
||||
return tags
|
||||
|
||||
|
||||
def fetch_e621(min_count: int = 10, partial_path: str = "", **_kwargs) -> list:
|
||||
"""Fetch tags from e621 API, paginated."""
|
||||
def fetch_e621(min_count: int = 10, partial_path: str = "", api_key: str | None = None, **_kwargs) -> list:
|
||||
"""Fetch tags from e621 API, paginated.
|
||||
With `api_key` set (login:api_key) the rate limit rises from anon (~1 rps) to authenticated.
|
||||
"""
|
||||
auth = parse_login_key(api_key)
|
||||
sleep_sec = 0.25 if auth else 1.0
|
||||
start_page, tags = load_partial(partial_path)
|
||||
page = (start_page + 1) if start_page is not None else 1
|
||||
session = requests.Session()
|
||||
session.headers["User-Agent"] = USER_AGENT
|
||||
while True:
|
||||
url = f"https://e621.net/tags.json?limit=320&page={page}&search[order]=count"
|
||||
params: dict[str, str] = {"limit": "320", "page": str(page), "search[order]": "count"}
|
||||
if auth:
|
||||
params["login"], params["api_key"] = auth
|
||||
try:
|
||||
resp = fetch_with_retry(session, url)
|
||||
resp = fetch_with_retry(session, "https://e621.net/tags.json", params=params)
|
||||
except requests.RequestException as e:
|
||||
print(f" Error on page {page}: {e}", file=sys.stderr)
|
||||
break
|
||||
@@ -186,11 +212,63 @@ def fetch_e621(min_count: int = 10, partial_path: str = "", **_kwargs) -> list:
|
||||
if page % SAVE_INTERVAL == 0:
|
||||
save_partial(partial_path, page, tags)
|
||||
page += 1
|
||||
time.sleep(1.0)
|
||||
time.sleep(sleep_sec)
|
||||
tags.sort(key=lambda t: t[2], reverse=True)
|
||||
return tags
|
||||
|
||||
|
||||
def fetch_danbooru_style_aliases(base_url: str, api_key: str | None, partial_path: str, sleep_sec: float) -> dict[str, list[str]]:
|
||||
"""Fetch active tag aliases from a Danbooru-style /tag_aliases.json endpoint.
|
||||
Works for both danbooru.donmai.us and e621.net, which share the same schema.
|
||||
Returns {consequent_name: [antecedent_names]}.
|
||||
"""
|
||||
partial = {}
|
||||
start_page, collected = load_partial(partial_path)
|
||||
if collected and isinstance(collected, dict):
|
||||
partial = collected
|
||||
page = (start_page + 1) if start_page is not None else 1
|
||||
auth = parse_login_key(api_key)
|
||||
session = requests.Session()
|
||||
session.headers["User-Agent"] = USER_AGENT
|
||||
while True:
|
||||
params: dict[str, str] = {"limit": "1000", "page": str(page), "search[status]": "active"}
|
||||
if auth:
|
||||
params["login"], params["api_key"] = auth
|
||||
try:
|
||||
resp = fetch_with_retry(session, base_url, params=params)
|
||||
except requests.RequestException as e:
|
||||
print(f" Error on alias page {page}: {e}", file=sys.stderr)
|
||||
break
|
||||
data = resp.json()
|
||||
if not data:
|
||||
break
|
||||
for row in data:
|
||||
ant = row.get("antecedent_name")
|
||||
con = row.get("consequent_name")
|
||||
if not ant or not con:
|
||||
continue
|
||||
partial.setdefault(con, []).append(ant)
|
||||
print(f" Alias page {page}: {len(data)} rows (total consequents: {len(partial)})", file=sys.stderr)
|
||||
if page % SAVE_INTERVAL == 0:
|
||||
# .partial for alias harvests stores the dict directly under `tags` for reuse of the loader.
|
||||
save_partial(partial_path, page, partial) # type: ignore[arg-type]
|
||||
page += 1
|
||||
time.sleep(sleep_sec)
|
||||
return partial
|
||||
|
||||
|
||||
def fetch_danbooru_aliases(api_key: str | None = None, partial_path: str = "", **_kwargs) -> dict[str, list[str]]:
|
||||
"""Harvest active tag aliases from Danbooru. Authenticated runs go ~10x faster."""
|
||||
sleep_sec = 0.1 if parse_login_key(api_key) else 1.0
|
||||
return fetch_danbooru_style_aliases("https://danbooru.donmai.us/tag_aliases.json", api_key, partial_path, sleep_sec)
|
||||
|
||||
|
||||
def fetch_e621_aliases(api_key: str | None = None, partial_path: str = "", **_kwargs) -> dict[str, list[str]]:
|
||||
"""Harvest active tag aliases from e621 (same schema as danbooru)."""
|
||||
sleep_sec = 0.25 if parse_login_key(api_key) else 1.0
|
||||
return fetch_danbooru_style_aliases("https://e621.net/tag_aliases.json", api_key, partial_path, sleep_sec)
|
||||
|
||||
|
||||
def fetch_gelbooru(base_url: str, min_count: int = 10, api_key: str | None = None, rate_limit: float = 0.5, partial_path: str = "") -> list:
|
||||
"""Fetch tags from a Gelbooru-compatible API (rule34, gelbooru, etc.).
|
||||
|
||||
@@ -259,16 +337,117 @@ def fetch_rule34(min_count: int = 10, api_key: str | None = None, partial_path:
|
||||
return fetch_gelbooru("https://api.rule34.xxx/index.php", min_count=min_count, api_key=api_key, partial_path=partial_path)
|
||||
|
||||
|
||||
def fetch_sankaku(min_count: int = 10, partial_path: str = "", **_kwargs) -> list:
|
||||
def fetch_danbooru_translations(api_key: str | None = None, partial_path: str = "", **_kwargs) -> dict[str, str]:
|
||||
"""Harvest foreign-name to canonical mappings from Danbooru's wiki.
|
||||
Each wiki page has `title` (canonical tag) and `other_names[]` (alternate names, often JA/KR/ZH).
|
||||
Returns {other_name_lower: title_lower}. Full harvest; authenticated runs go ~10x faster.
|
||||
"""
|
||||
auth = parse_login_key(api_key)
|
||||
sleep_sec = 0.1 if auth else 1.0
|
||||
partial: dict[str, str] = {}
|
||||
start_page, collected = load_partial(partial_path)
|
||||
if collected and isinstance(collected, dict):
|
||||
partial = collected
|
||||
page = (start_page + 1) if start_page is not None else 1
|
||||
session = requests.Session()
|
||||
session.headers["User-Agent"] = USER_AGENT
|
||||
while True:
|
||||
params: dict[str, str] = {"limit": "1000", "page": str(page), "only": "title,other_names"}
|
||||
if auth:
|
||||
params["login"], params["api_key"] = auth
|
||||
try:
|
||||
resp = fetch_with_retry(session, "https://danbooru.donmai.us/wiki_pages.json", params=params)
|
||||
except requests.RequestException as e:
|
||||
print(f" Error on wiki page {page}: {e}", file=sys.stderr)
|
||||
break
|
||||
data = resp.json()
|
||||
if not data:
|
||||
break
|
||||
added = 0
|
||||
for wiki in data:
|
||||
title = (wiki.get("title") or "").strip().lower()
|
||||
if not title:
|
||||
continue
|
||||
for other in wiki.get("other_names") or []:
|
||||
other_norm = (other or "").strip().lower()
|
||||
if not other_norm or other_norm == title:
|
||||
continue
|
||||
# First-wins: if two wiki pages claim the same foreign term, keep the first seen.
|
||||
partial.setdefault(other_norm, title)
|
||||
added += 1
|
||||
print(f" Wiki page {page}: {len(data)} entries ({added} names, total: {len(partial)})", file=sys.stderr)
|
||||
if page % SAVE_INTERVAL == 0:
|
||||
save_partial(partial_path, page, partial) # type: ignore[arg-type]
|
||||
page += 1
|
||||
time.sleep(sleep_sec)
|
||||
return partial
|
||||
|
||||
|
||||
def fetch_rule34_aliases(partial_path: str = "", **_kwargs) -> dict[str, list[str]]:
|
||||
"""Harvest rule34.xxx aliases by scraping the public /index.php?page=alias&s=list listing.
|
||||
No API key required (alias list is public). 50 rows per page; `pid` advances by 50; empty page = done.
|
||||
"""
|
||||
import re
|
||||
from html import unescape
|
||||
partial: dict[str, list[str]] = {}
|
||||
start_pid, collected = load_partial(partial_path)
|
||||
if collected and isinstance(collected, dict):
|
||||
partial = collected
|
||||
pid = (start_pid + 50) if start_pid is not None else 0
|
||||
session = requests.Session()
|
||||
# A browser-like UA is needed; the default python-requests UA gets 403.
|
||||
session.headers["User-Agent"] = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 Chrome/138.0.0.0 Safari/537.36"
|
||||
# Matches the 4 <td> cells of an alias row: checkbox, alias-link, canonical-link, reason.
|
||||
# Each tag link's text is `tagname` (possibly with ` (count)` suffix), so we extract from the tags= URL param instead.
|
||||
row_re = re.compile(r'<tr class="(?:even|odd)"[^>]*>.*?</tr>', re.DOTALL)
|
||||
tag_re = re.compile(r'tags=([^"&]+)')
|
||||
while True:
|
||||
try:
|
||||
resp = fetch_with_retry(session, "https://rule34.xxx/index.php",
|
||||
params={"page": "alias", "s": "list", "pid": str(pid)})
|
||||
except requests.RequestException as e:
|
||||
print(f" Error on alias pid={pid}: {e}", file=sys.stderr)
|
||||
break
|
||||
rows = row_re.findall(resp.text)
|
||||
if not rows:
|
||||
break
|
||||
added = 0
|
||||
for row in rows:
|
||||
matches = tag_re.findall(row)
|
||||
if len(matches) < 2:
|
||||
continue
|
||||
ant = unescape(matches[0]).strip()
|
||||
con = unescape(matches[1]).strip()
|
||||
if not ant or not con:
|
||||
continue
|
||||
partial.setdefault(con, []).append(ant)
|
||||
added += 1
|
||||
print(f" Alias pid={pid}: {len(rows)} rows ({added} parsed, total consequents: {len(partial)})", file=sys.stderr)
|
||||
if len(rows) < 50:
|
||||
break # last page
|
||||
if (pid // 50) % SAVE_INTERVAL == 0:
|
||||
save_partial(partial_path, pid, partial) # type: ignore[arg-type]
|
||||
pid += 50
|
||||
time.sleep(1.0)
|
||||
return partial
|
||||
|
||||
|
||||
def fetch_sankaku(min_count: int = 10, partial_path: str = "", translations_out: dict | None = None, **_kwargs) -> list:
|
||||
"""Fetch tags from Sankaku Complex (chan.sankakucomplex.com).
|
||||
|
||||
Uses the public JSON API at sankakuapi.com which supports order=count,
|
||||
so we can stop early when counts drop below min_count.
|
||||
Tag names come as English with spaces - converted to lowercase.
|
||||
|
||||
If `translations_out` is provided, it is populated with {foreign_term_lower: canonical_lower}
|
||||
harvested from each tag's name_ja + translations[] fields in the same pass, avoiding a second
|
||||
full-API walk.
|
||||
"""
|
||||
start_page, tags = load_partial(partial_path)
|
||||
page = (start_page + 1) if start_page is not None else 1
|
||||
page_size = 200
|
||||
# Languages to include in translations. EN is already the canonical tag, so it's excluded.
|
||||
translation_langs = {"ja", "ko", "zh", "de", "it", "pt", "ru", "fr", "es"}
|
||||
session = requests.Session()
|
||||
session.headers["User-Agent"] = USER_AGENT
|
||||
while True:
|
||||
@@ -290,6 +469,18 @@ def fetch_sankaku(min_count: int = 10, partial_path: str = "", **_kwargs) -> lis
|
||||
continue
|
||||
name = name.strip().lower()
|
||||
tags.append([name, tag.get("type", 0), count])
|
||||
if translations_out is None:
|
||||
continue
|
||||
# name_ja is present on most tags; translations[] covers other languages when available.
|
||||
name_ja = (tag.get("name_ja") or "").strip().lower()
|
||||
if name_ja and name_ja != name:
|
||||
translations_out.setdefault(name_ja, name)
|
||||
for entry in tag.get("translations") or []:
|
||||
lang = (entry.get("lang") or "").strip().lower()
|
||||
term = (entry.get("translation") or "").strip().lower()
|
||||
if not term or lang not in translation_langs or term == name:
|
||||
continue
|
||||
translations_out.setdefault(term, name)
|
||||
below_threshold = all(tag.get("post_count", 0) < min_count for tag in data)
|
||||
print(f" Page {page}: {len(data)} tags (total: {len(tags)})", file=sys.stderr)
|
||||
if below_threshold:
|
||||
@@ -344,27 +535,55 @@ def fetch_idol(min_count: int = 10, partial_path: str = "", **_kwargs) -> list:
|
||||
|
||||
|
||||
SOURCES = {
|
||||
"danbooru": {"fetch": fetch_danbooru, "type_map": DANBOORU_TYPE_MAP},
|
||||
"e621": {"fetch": fetch_e621, "type_map": E621_TYPE_MAP},
|
||||
"rule34": {"fetch": fetch_rule34, "type_map": RULE34_TYPE_MAP},
|
||||
"sankaku": {"fetch": fetch_sankaku, "type_map": SANKAKU_TYPE_MAP},
|
||||
"danbooru": {
|
||||
"fetch": fetch_danbooru, "type_map": DANBOORU_TYPE_MAP,
|
||||
"fetch_aliases": fetch_danbooru_aliases,
|
||||
"fetch_translations": fetch_danbooru_translations,
|
||||
},
|
||||
"e621": {
|
||||
"fetch": fetch_e621, "type_map": E621_TYPE_MAP,
|
||||
"fetch_aliases": fetch_e621_aliases,
|
||||
},
|
||||
"rule34": {
|
||||
"fetch": fetch_rule34, "type_map": RULE34_TYPE_MAP,
|
||||
"fetch_aliases": fetch_rule34_aliases,
|
||||
},
|
||||
"sankaku": {"fetch": fetch_sankaku, "type_map": SANKAKU_TYPE_MAP}, # translations harvested inline
|
||||
"idol": {"fetch": fetch_idol, "type_map": IDOL_TYPE_MAP},
|
||||
}
|
||||
|
||||
|
||||
def write_dict(name: str, tags: list, type_map: dict, output_path: str, separator: str = "_"):
|
||||
def write_dict(name: str, tags: list, type_map: dict, output_path: str,
|
||||
separator: str = "_", aliases: dict[str, list[str]] | None = None):
|
||||
"""Write dict JSON file atomically.
|
||||
|
||||
type_map remaps source-native category IDs to unified IDs.
|
||||
separator controls the word separator in tag names:
|
||||
"_" (default) → "high_resolution" (booru convention, anime/illustration models)
|
||||
" " → "high resolution" (natural language, SDXL/Flux-style models)
|
||||
aliases, if provided, is {canonical_name: [alternative_names]}. A 4-tuple is emitted
|
||||
only for tags that have non-empty aliases, keeping the on-disk size of alias-free
|
||||
tags unchanged for backward compatibility with clients reading 3-tuples.
|
||||
"""
|
||||
aliases = aliases or {}
|
||||
normalized = []
|
||||
matched_aliases = 0
|
||||
for t in tags:
|
||||
tag_name = t[0].replace(" ", "_") if separator == "_" else t[0].replace("_", " ")
|
||||
source_name = t[0]
|
||||
canonical_key = source_name.strip().lower()
|
||||
tag_name = source_name.replace(" ", "_") if separator == "_" else source_name.replace("_", " ")
|
||||
category = type_map.get(t[1], 0)
|
||||
normalized.append([tag_name, category, t[2]])
|
||||
tag_aliases = aliases.get(canonical_key) or aliases.get(tag_name.lower())
|
||||
if tag_aliases:
|
||||
# Normalize alias word-separator to match the tag's separator choice.
|
||||
if separator == "_":
|
||||
tag_aliases = [a.replace(" ", "_") for a in tag_aliases]
|
||||
else:
|
||||
tag_aliases = [a.replace("_", " ") for a in tag_aliases]
|
||||
normalized.append([tag_name, category, t[2], tag_aliases])
|
||||
matched_aliases += 1
|
||||
else:
|
||||
normalized.append([tag_name, category, t[2]])
|
||||
data = {
|
||||
"name": name,
|
||||
"version": date.today().isoformat(),
|
||||
@@ -377,23 +596,77 @@ def write_dict(name: str, tags: list, type_map: dict, output_path: str, separato
|
||||
json.dump(data, f, ensure_ascii=False, separators=(",", ":"))
|
||||
os.replace(tmp_path, output_path)
|
||||
size_mb = os.path.getsize(output_path) / (1024 * 1024)
|
||||
print(f" Written: {output_path} ({len(tags)} tags, {size_mb:.1f} MB)", file=sys.stderr)
|
||||
alias_note = f", {matched_aliases} with aliases" if matched_aliases else ""
|
||||
print(f" Written: {output_path} ({len(tags)} tags{alias_note}, {size_mb:.1f} MB)", file=sys.stderr)
|
||||
|
||||
|
||||
def fetch_source(name: str, output: str, min_count: int, api_key: str | None = None, separator: str = "_"):
|
||||
"""Fetch and write a single source."""
|
||||
def write_translations(mapping: dict[str, str], output_path: str):
|
||||
"""Write the translations companion file atomically.
|
||||
`mapping` is {foreign_term: canonical_tag_name}, both lowercased and underscore-normalized.
|
||||
"""
|
||||
if not mapping:
|
||||
return
|
||||
os.makedirs(os.path.dirname(output_path) or '.', exist_ok=True)
|
||||
tmp_path = output_path + ".tmp"
|
||||
with open(tmp_path, "w", encoding="utf-8") as f:
|
||||
json.dump(mapping, f, ensure_ascii=False, separators=(",", ":"), sort_keys=True)
|
||||
os.replace(tmp_path, output_path)
|
||||
size_kb = os.path.getsize(output_path) / 1024
|
||||
print(f" Translations: {output_path} ({len(mapping)} terms, {size_kb:.1f} KB)", file=sys.stderr)
|
||||
|
||||
|
||||
def fetch_source(name: str, output: str, min_count: int, keys: dict[str, str] | None = None, separator: str = "_"):
|
||||
"""Fetch and write a single source. `keys` maps source name to `login:token` (or rule34's `uid:token`)."""
|
||||
if name not in SOURCES:
|
||||
print(f"Unknown source: {name}. Available: {', '.join(SOURCES.keys())}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
source = SOURCES[name]
|
||||
keys = keys or {}
|
||||
api_key = keys.get(name)
|
||||
partial_path = output + ".partial"
|
||||
print(f"Fetching {name} (min_count={min_count})...", file=sys.stderr)
|
||||
tags = source["fetch"](min_count=min_count, api_key=api_key, partial_path=partial_path)
|
||||
|
||||
# Sankaku collects translations inline during the tag walk to avoid a second full API pass.
|
||||
translations_inline: dict[str, str] | None = {} if name == "sankaku" else None
|
||||
tags = source["fetch"](
|
||||
min_count=min_count, api_key=api_key, partial_path=partial_path,
|
||||
translations_out=translations_inline,
|
||||
) if translations_inline is not None else source["fetch"](
|
||||
min_count=min_count, api_key=api_key, partial_path=partial_path,
|
||||
)
|
||||
if not tags:
|
||||
print(f" No tags fetched for {name}", file=sys.stderr)
|
||||
return
|
||||
write_dict(name, tags, source["type_map"], output, separator=separator)
|
||||
|
||||
# Optional alias harvest via a separate .partial file (independent resume).
|
||||
aliases: dict[str, list[str]] = {}
|
||||
if source.get("fetch_aliases"):
|
||||
alias_partial = output + ".aliases.partial"
|
||||
print(f"Fetching {name} aliases...", file=sys.stderr)
|
||||
try:
|
||||
aliases = source["fetch_aliases"](api_key=api_key, partial_path=alias_partial)
|
||||
clear_partial(alias_partial)
|
||||
except Exception as e:
|
||||
print(f" Alias fetch failed for {name}: {e}", file=sys.stderr)
|
||||
aliases = {}
|
||||
|
||||
write_dict(name, tags, source["type_map"], output, separator=separator, aliases=aliases)
|
||||
clear_partial(partial_path)
|
||||
|
||||
# Translations: either inline (sankaku) or via a dedicated fetcher (danbooru wiki).
|
||||
translations: dict[str, str] = translations_inline or {}
|
||||
if source.get("fetch_translations"):
|
||||
tr_partial = output + ".translations.partial"
|
||||
print(f"Fetching {name} translations...", file=sys.stderr)
|
||||
try:
|
||||
translations = source["fetch_translations"](api_key=api_key, partial_path=tr_partial) or {}
|
||||
clear_partial(tr_partial)
|
||||
except Exception as e:
|
||||
print(f" Translation fetch failed for {name}: {e}", file=sys.stderr)
|
||||
if translations:
|
||||
tr_output = output.replace(".json", ".translations.json") if output.endswith(".json") else f"{output}.translations.json"
|
||||
write_translations(translations, tr_output)
|
||||
|
||||
from importlib.util import spec_from_file_location, module_from_spec
|
||||
spec = spec_from_file_location("tags_manifest", os.path.join(os.path.dirname(__file__), "tags-manifest.py"))
|
||||
mod = module_from_spec(spec)
|
||||
@@ -408,19 +681,29 @@ def main():
|
||||
parser.add_argument("--output", "-o", help="Output file path (for single source)")
|
||||
parser.add_argument("--output-dir", "-d", help="Output directory (for 'all')")
|
||||
parser.add_argument("--min-count", "-m", type=int, default=10, help="Minimum post count to include (default: 10)")
|
||||
parser.add_argument("--key", "-k", help="API key as USER_ID:API_KEY (required for rule34)")
|
||||
parser.add_argument("--key", "-k", help="Rule34 API key as USER_ID:API_KEY")
|
||||
parser.add_argument("--danbooru-key", help="Danbooru login:api_key (raises rate limit from 1 to 10 rps)")
|
||||
parser.add_argument("--e621-key", help="e621 login:api_key (raises rate limit)")
|
||||
parser.add_argument("--spaces", action="store_true", help="Use spaces instead of underscores in tag names (for natural language models like SDXL/Flux)")
|
||||
args = parser.parse_args()
|
||||
separator = " " if args.spaces else "_"
|
||||
# Per-source key map. Rule34 keeps the original `--key` (uid:token); danbooru/e621 get dedicated flags.
|
||||
keys: dict[str, str] = {}
|
||||
if args.key:
|
||||
keys["rule34"] = args.key
|
||||
if args.danbooru_key:
|
||||
keys["danbooru"] = args.danbooru_key
|
||||
if args.e621_key:
|
||||
keys["e621"] = args.e621_key
|
||||
|
||||
if args.source == "all":
|
||||
output_dir = args.output_dir or "."
|
||||
for name in SOURCES:
|
||||
output = os.path.join(output_dir, f"{name}.json")
|
||||
fetch_source(name, output, args.min_count, api_key=args.key, separator=separator)
|
||||
fetch_source(name, output, args.min_count, keys=keys, separator=separator)
|
||||
else:
|
||||
output = args.output or f"{args.source}.json"
|
||||
fetch_source(args.source, output, args.min_count, api_key=args.key, separator=separator)
|
||||
fetch_source(args.source, output, args.min_count, keys=keys, separator=separator)
|
||||
|
||||
print("Done.", file=sys.stderr)
|
||||
|
||||
|
||||
+13
-2
@@ -30,18 +30,25 @@ DESCRIPTIONS = {
|
||||
|
||||
|
||||
def build_entry(filepath: str) -> dict:
|
||||
"""Build a manifest entry from a tag JSON file."""
|
||||
"""Build a manifest entry from a tag JSON file.
|
||||
If a `<name>.translations.json` companion sits next to the file, the entry gets
|
||||
`translations: true` so the client-side downloader knows to pull the companion too.
|
||||
"""
|
||||
with open(filepath, encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
name = data.get("name") or os.path.splitext(os.path.basename(filepath))[0]
|
||||
size_mb = round(os.path.getsize(filepath) / (1024 * 1024), 1)
|
||||
return {
|
||||
entry = {
|
||||
"name": name,
|
||||
"description": DESCRIPTIONS.get(name, ""),
|
||||
"version": data.get("version", ""),
|
||||
"tag_count": len(data.get("tags", [])),
|
||||
"size_mb": size_mb,
|
||||
}
|
||||
translations_path = os.path.join(os.path.dirname(filepath), f"{name}.translations.json")
|
||||
if os.path.isfile(translations_path):
|
||||
entry["translations"] = True
|
||||
return entry
|
||||
|
||||
|
||||
def update_manifest(directory: str) -> bool:
|
||||
@@ -87,6 +94,10 @@ def main():
|
||||
basename = os.path.splitext(os.path.basename(filepath))[0]
|
||||
if basename in exclude or basename == "manifest":
|
||||
continue
|
||||
# Skip translation companion files. They're pulled in automatically via `translations: true`
|
||||
# flags on their parent dict entries; standalone entries would be malformed.
|
||||
if basename.endswith(".translations"):
|
||||
continue
|
||||
if not os.path.isfile(filepath):
|
||||
print(f" Skipping {filepath}: not found", file=sys.stderr)
|
||||
continue
|
||||
|
||||
+244
-44
@@ -42,6 +42,14 @@ const CATEGORY_NAMES = {
|
||||
13: 'color',
|
||||
};
|
||||
|
||||
// Glyph + color per result kind. Renders in place of the category dot for non-tag results.
|
||||
const KIND_GLYPHS = {
|
||||
tag: { glyph: '●', color: null }, // color pulled from tag category
|
||||
lora: { glyph: '◆', color: '#8a66ff' },
|
||||
embed: { glyph: '▲', color: '#1abc9c' },
|
||||
wildcard: { glyph: '★', color: '#f1c40f' },
|
||||
};
|
||||
|
||||
let active = false;
|
||||
|
||||
// -- Utilities (ported from Enso) --
|
||||
@@ -103,34 +111,98 @@ function caretViewportY(textarea) {
|
||||
class TagIndex {
|
||||
constructor(data) {
|
||||
this.categories = data.categories || {};
|
||||
// Build sorted array of {name, category, count} from raw [name, catId, count] tuples
|
||||
this.tags = data.tags.map(([name, category, count]) => ({
|
||||
// Tuples are [name, catId, count] or [name, catId, count, aliases]. Default `aliases = []`
|
||||
// keeps legacy 3-tuple dictionaries working unchanged.
|
||||
this.tags = data.tags.map(([name, category, count, aliases = []]) => ({
|
||||
name: name.toLowerCase(),
|
||||
display: name,
|
||||
category,
|
||||
count,
|
||||
aliases,
|
||||
}));
|
||||
this.tags.sort((a, b) => a.name.localeCompare(b.name));
|
||||
// Alias index parallel to this.tags. Each entry has .name so lowerBound works on both.
|
||||
this.aliasEntries = [];
|
||||
for (const tag of this.tags) {
|
||||
if (!tag.aliases || tag.aliases.length === 0) continue;
|
||||
for (const alias of tag.aliases) {
|
||||
this.aliasEntries.push({ name: alias.toLowerCase(), display: alias, tag });
|
||||
}
|
||||
}
|
||||
this.aliasEntries.sort((a, b) => a.name.localeCompare(b.name));
|
||||
// Optional translations companion: foreign_term -> canonical_tag_name.
|
||||
// tagByName is keyed on canonical lowercased name for O(1) resolution from a translation hit.
|
||||
this.translations = new Map();
|
||||
this.tagByName = new Map(this.tags.map((t) => [t.name, t]));
|
||||
if (data.translations && typeof data.translations === 'object') {
|
||||
for (const [foreign, canonical] of Object.entries(data.translations)) {
|
||||
if (typeof foreign !== 'string' || typeof canonical !== 'string') continue;
|
||||
this.translations.set(foreign.toLowerCase(), { canonical: canonical.toLowerCase(), foreign });
|
||||
}
|
||||
}
|
||||
// Sorted translation keys for prefix+substring scan via lowerBound.
|
||||
this.translationEntries = [...this.translations.entries()]
|
||||
.map(([foreignLower, { canonical, foreign }]) => ({ name: foreignLower, foreign, canonical }))
|
||||
.sort((a, b) => a.name.localeCompare(b.name));
|
||||
}
|
||||
|
||||
/** Prefix search with binary search. Returns matches sorted by count descending. */
|
||||
/** Prefix search with binary search across canonical names and aliases. Returns matches sorted by count descending. */
|
||||
search(prefix, limit = 20) {
|
||||
const query = prefix.toLowerCase().replace(/ /g, '_');
|
||||
if (!query) return [];
|
||||
const start = lowerBound(this.tags, query);
|
||||
// Canonical prefix matches
|
||||
const matches = [];
|
||||
const start = lowerBound(this.tags, query);
|
||||
for (let i = start; i < this.tags.length && matches.length < limit * 5; i++) {
|
||||
if (!this.tags[i].name.startsWith(query)) break;
|
||||
matches.push(this.tags[i]);
|
||||
}
|
||||
// Substring fallback for 4+ chars if prefix found nothing
|
||||
// Alias prefix matches. Annotate so render can show "canonical (alias)".
|
||||
const aliasStart = lowerBound(this.aliasEntries, query);
|
||||
for (let i = aliasStart; i < this.aliasEntries.length && matches.length < limit * 10; i++) {
|
||||
const entry = this.aliasEntries[i];
|
||||
if (!entry.name.startsWith(query)) break;
|
||||
matches.push({ ...entry.tag, matchedVia: 'alias', matchedAlias: entry.display });
|
||||
}
|
||||
// Substring fallback (canonical + aliases) for 4+ char queries when prefix matching returned nothing.
|
||||
if (matches.length === 0 && query.length >= 4) {
|
||||
for (let i = 0; i < this.tags.length && matches.length < limit * 5; i++) {
|
||||
if (this.tags[i].name.includes(query)) matches.push(this.tags[i]);
|
||||
}
|
||||
for (let i = 0; i < this.aliasEntries.length && matches.length < limit * 10; i++) {
|
||||
const entry = this.aliasEntries[i];
|
||||
if (entry.name.includes(query)) matches.push({ ...entry.tag, matchedVia: 'alias', matchedAlias: entry.display });
|
||||
}
|
||||
}
|
||||
matches.sort((a, b) => b.count - a.count);
|
||||
return matches.slice(0, limit);
|
||||
// Translation lookup. Prefix scan over foreign terms, resolving to canonical tags when present.
|
||||
if (this.translationEntries.length > 0) {
|
||||
const tStart = lowerBound(this.translationEntries, query);
|
||||
for (let i = tStart; i < this.translationEntries.length && matches.length < limit * 10; i++) {
|
||||
const entry = this.translationEntries[i];
|
||||
if (!entry.name.startsWith(query)) break;
|
||||
const canonicalTag = this.tagByName.get(entry.canonical);
|
||||
if (canonicalTag) matches.push({ ...canonicalTag, matchedVia: 'translation', matchedTerm: entry.foreign });
|
||||
}
|
||||
// Substring fallback over translation keys (CJK/short foreign terms benefit from 2-char threshold)
|
||||
if (query.length >= 2) {
|
||||
for (let i = 0; i < this.translationEntries.length && matches.length < limit * 10; i++) {
|
||||
const entry = this.translationEntries[i];
|
||||
if (entry.name.includes(query) && !entry.name.startsWith(query)) {
|
||||
const canonicalTag = this.tagByName.get(entry.canonical);
|
||||
if (canonicalTag) matches.push({ ...canonicalTag, matchedVia: 'translation', matchedTerm: entry.foreign });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Dedupe by canonical name; prefer canonical (no matchedVia) over alias/translation matches.
|
||||
const seen = new Map();
|
||||
for (const tag of matches) {
|
||||
const existing = seen.get(tag.name);
|
||||
if (!existing || (existing.matchedVia && !tag.matchedVia)) seen.set(tag.name, tag);
|
||||
}
|
||||
const result = [...seen.values()];
|
||||
result.sort((a, b) => b.count - a.count);
|
||||
return result.slice(0, limit);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -194,47 +266,100 @@ const engine = {
|
||||
|
||||
// -- Textarea integration --
|
||||
|
||||
/** Extract the current word being typed at the cursor position. */
|
||||
/**
|
||||
* Extract the current completion context at the cursor position.
|
||||
*
|
||||
* Returns { word, start, end, mode } where:
|
||||
* mode === 'tag': ordinary tag completion
|
||||
* mode === 'lora': inside an unclosed `<lora:...` span; `start` points at the `<`
|
||||
* mode === 'wildcard': inside an unclosed `__...` span; `start` points at the first `_`
|
||||
*
|
||||
* `start..end` is the replacement range the appropriate insert function should overwrite.
|
||||
* Embeddings are served under `mode === 'tag'` and merged into tag-mode results by the engine.
|
||||
*/
|
||||
function getCurrentWord(textarea) {
|
||||
const { value, selectionStart } = textarea;
|
||||
if (selectionStart !== textarea.selectionEnd) return null; // has selection
|
||||
// Scan backward from cursor to find word start
|
||||
let start = selectionStart;
|
||||
while (start > 0) {
|
||||
const ch = value[start - 1];
|
||||
// Scan backward from cursor to the nearest hard separator
|
||||
let wordStart = selectionStart;
|
||||
while (wordStart > 0) {
|
||||
const ch = value[wordStart - 1];
|
||||
if (ch === ',' || ch === '\n') break;
|
||||
start--;
|
||||
wordStart--;
|
||||
}
|
||||
// Skip leading whitespace
|
||||
while (start < selectionStart && value[start] === ' ') start++;
|
||||
const word = value.slice(start, selectionStart);
|
||||
if (!word) return null;
|
||||
// Skip if inside angle brackets (LoRA/embedding syntax)
|
||||
// Skip leading whitespace between the separator and the typed word
|
||||
while (wordStart < selectionStart && value[wordStart] === ' ') wordStart++;
|
||||
const segment = value.slice(wordStart, selectionStart);
|
||||
// LoRA / extra-network trigger: unclosed `<` with `kind:` prefix
|
||||
const before = value.slice(0, selectionStart);
|
||||
const lastOpen = before.lastIndexOf('<');
|
||||
const lastClose = before.lastIndexOf('>');
|
||||
if (lastOpen > lastClose) return null;
|
||||
// Skip if inside wildcard syntax
|
||||
const wcBefore = before.slice(start);
|
||||
if (wcBefore.startsWith('__') && !wcBefore.endsWith('__')) return null;
|
||||
return { word, start, end: selectionStart };
|
||||
if (lastOpen > lastClose && lastOpen >= wordStart) {
|
||||
const inside = before.slice(lastOpen + 1); // e.g. "lora:foo" or "lora:" or "lor"
|
||||
const colon = inside.indexOf(':');
|
||||
// Require `<lora:`; before the colon the kind is ambiguous (could be lora/embed/hypernet).
|
||||
if (colon >= 0 && inside.slice(0, colon).toLowerCase() === 'lora') {
|
||||
return { word: inside.slice(colon + 1), start: lastOpen, end: selectionStart, mode: 'lora' };
|
||||
}
|
||||
// Inside `<...` but not yet a recognized kind, suppress completion.
|
||||
return null;
|
||||
}
|
||||
// Wildcard trigger: unclosed `__` that doesn't close within the current word
|
||||
if (segment.startsWith('__') && !segment.slice(2).includes('__')) {
|
||||
return { word: segment.slice(2), start: wordStart, end: selectionStart, mode: 'wildcard' };
|
||||
}
|
||||
// Ordinary tag
|
||||
if (!segment) return null;
|
||||
return { word: segment, start: wordStart, end: selectionStart, mode: 'tag' };
|
||||
}
|
||||
|
||||
/** Escape bare parens so tag names like `fate_(series)` aren't parsed as attention syntax. */
|
||||
function escapeParensForPrompt(name) {
|
||||
return name.replace(/([()])/g, '\\$1');
|
||||
}
|
||||
|
||||
/**
|
||||
* Insert an extra-network reference at the current trigger position.
|
||||
* kind === 'lora': inserts `<lora:name:1.0>` over the range including the leading `<`
|
||||
* kind === 'wildcard': inserts `__name__` over the range including the leading `__`
|
||||
* Embeddings use insertTag directly so they go through comma-separator and paren-escape logic.
|
||||
*/
|
||||
function insertExtraNetwork(textarea, item, kind) {
|
||||
const info = getCurrentWord(textarea);
|
||||
if (!info || info.mode !== kind) return;
|
||||
const { value } = textarea;
|
||||
const before = value.slice(0, info.start);
|
||||
const after = value.slice(info.end);
|
||||
let insertion;
|
||||
if (kind === 'lora') {
|
||||
insertion = `<lora:${item.display ?? item.name}:1.0>`;
|
||||
} else if (kind === 'wildcard') {
|
||||
insertion = `__${item.display ?? item.name}__`;
|
||||
} else {
|
||||
return;
|
||||
}
|
||||
textarea.value = before + insertion + after;
|
||||
const cursorPos = before.length + insertion.length;
|
||||
textarea.selectionStart = cursorPos;
|
||||
textarea.selectionEnd = cursorPos;
|
||||
if (typeof updateInput === 'function') updateInput(textarea);
|
||||
}
|
||||
|
||||
/** Insert a tag at the current word position, replacing the typed prefix. */
|
||||
function insertTag(textarea, tagName) {
|
||||
const info = getCurrentWord(textarea);
|
||||
if (!info) return;
|
||||
if (!info || info.mode !== 'tag') return;
|
||||
const { value } = textarea;
|
||||
const before = value.slice(0, info.start);
|
||||
const after = value.slice(info.end);
|
||||
// Build insertion: tag + separator
|
||||
// Build insertion: tag + separator. Parens in tag names are escaped so the prompt parser doesn't read them as attention syntax.
|
||||
const useComma = window.opts?.autocomplete_append_comma ?? true;
|
||||
const sep = useComma ? ',' : '';
|
||||
const needsSepBefore = before.length > 0 && before.trimEnd().length > 0 && !before.trimEnd().endsWith(',');
|
||||
const prefix = needsSepBefore ? `${sep} ` : '';
|
||||
let suffix = `${sep} `;
|
||||
if (after.length > 0 && after.trimStart().startsWith(',')) suffix = ' ';
|
||||
const insertion = `${prefix}${tagName}${suffix}`;
|
||||
const insertion = `${prefix}${escapeParensForPrompt(tagName)}${suffix}`;
|
||||
textarea.value = before.trimEnd() + (before.trimEnd().length > 0 ? ' ' : '') + insertion + after.trimStart();
|
||||
// Position cursor after the inserted tag + separator
|
||||
const cursorPos = before.trimEnd().length + (before.trimEnd().length > 0 ? 1 : 0) + insertion.length;
|
||||
@@ -280,10 +405,9 @@ const dropdown = {
|
||||
|
||||
show(results, textarea, query) {
|
||||
if (results.length === 0) { this.hide(); return; }
|
||||
if (this.textarea !== textarea) {
|
||||
if (this.textarea) this.resizeObserver.unobserve(this.textarea);
|
||||
this.resizeObserver.observe(textarea);
|
||||
}
|
||||
// Switching textareas: clear prior state so a stale render can't leak across.
|
||||
if (this.textarea && this.textarea !== textarea) this.hide();
|
||||
if (this.textarea !== textarea) this.resizeObserver.observe(textarea);
|
||||
this.results = results;
|
||||
this.textarea = textarea;
|
||||
this.query = query || '';
|
||||
@@ -312,24 +436,50 @@ const dropdown = {
|
||||
if (i === this.selectedIndex) li.classList.add('selected');
|
||||
const dot = document.createElement('span');
|
||||
dot.className = 'autocomplete-category';
|
||||
dot.style.color = engine.categoryColors[tag.category] || '#888';
|
||||
dot.textContent = '\u25CF';
|
||||
dot.title = engine.categoryNames[tag.category] || '';
|
||||
const kind = tag.kind || 'tag';
|
||||
const kindStyle = KIND_GLYPHS[kind] || KIND_GLYPHS.tag;
|
||||
dot.style.color = kindStyle.color || engine.categoryColors[tag.category] || '#888';
|
||||
dot.textContent = kindStyle.glyph;
|
||||
dot.title = kind === 'tag' ? (engine.categoryNames[tag.category] || '') : kind;
|
||||
const name = document.createElement('span');
|
||||
name.className = 'autocomplete-tag';
|
||||
const tagText = replaceUnderscores ? tag.display.replace(/_/g, ' ') : tag.display;
|
||||
const matchPos = tag.name.indexOf(queryNorm);
|
||||
if (matchPos >= 0 && queryNorm.length > 0) {
|
||||
const canonicalMatch = tag.name.indexOf(queryNorm);
|
||||
if (canonicalMatch >= 0 && queryNorm.length > 0) {
|
||||
const mark = document.createElement('mark');
|
||||
mark.textContent = tagText.slice(matchPos, matchPos + queryNorm.length);
|
||||
mark.textContent = tagText.slice(canonicalMatch, canonicalMatch + queryNorm.length);
|
||||
name.append(
|
||||
document.createTextNode(tagText.slice(0, matchPos)),
|
||||
document.createTextNode(tagText.slice(0, canonicalMatch)),
|
||||
mark,
|
||||
document.createTextNode(tagText.slice(matchPos + queryNorm.length)),
|
||||
document.createTextNode(tagText.slice(canonicalMatch + queryNorm.length)),
|
||||
);
|
||||
} else {
|
||||
name.textContent = tagText;
|
||||
}
|
||||
// Alias/translation-matched rows append " (foreign)" with the query fragment highlighted.
|
||||
let annotationTerm = null;
|
||||
if (tag.matchedVia === 'alias') annotationTerm = tag.matchedAlias;
|
||||
else if (tag.matchedVia === 'translation') annotationTerm = tag.matchedTerm;
|
||||
if (annotationTerm) {
|
||||
const annotationDisplay = replaceUnderscores ? annotationTerm.replace(/_/g, ' ') : annotationTerm;
|
||||
const annotationLower = annotationTerm.toLowerCase();
|
||||
const annotationMatch = annotationLower.indexOf(queryNorm);
|
||||
const prefix = tag.matchedVia === 'translation' ? ' \u{1F310} ' : ' (';
|
||||
const suffix = tag.matchedVia === 'translation' ? '' : ')';
|
||||
name.appendChild(document.createTextNode(prefix));
|
||||
if (annotationMatch >= 0 && queryNorm.length > 0) {
|
||||
const mark = document.createElement('mark');
|
||||
mark.textContent = annotationDisplay.slice(annotationMatch, annotationMatch + queryNorm.length);
|
||||
name.append(
|
||||
document.createTextNode(annotationDisplay.slice(0, annotationMatch)),
|
||||
mark,
|
||||
document.createTextNode(annotationDisplay.slice(annotationMatch + queryNorm.length)),
|
||||
);
|
||||
} else {
|
||||
name.appendChild(document.createTextNode(annotationDisplay));
|
||||
}
|
||||
if (suffix) name.appendChild(document.createTextNode(suffix));
|
||||
}
|
||||
const count = document.createElement('span');
|
||||
count.className = 'autocomplete-count';
|
||||
count.textContent = tag.count > 0 ? formatCount(tag.count) : '';
|
||||
@@ -386,8 +536,15 @@ const dropdown = {
|
||||
}
|
||||
return;
|
||||
}
|
||||
const tag = this.results[this.selectedIndex];
|
||||
if (this.textarea) insertTag(this.textarea, tag.display);
|
||||
const result = this.results[this.selectedIndex];
|
||||
if (this.textarea) {
|
||||
if (result.kind === 'lora' || result.kind === 'wildcard') {
|
||||
insertExtraNetwork(this.textarea, result, result.kind);
|
||||
} else {
|
||||
// 'embed' kind and untagged tag results both go through insertTag (comma-aware, paren-escaped).
|
||||
insertTag(this.textarea, result.display ?? result.name);
|
||||
}
|
||||
}
|
||||
this.hide();
|
||||
},
|
||||
};
|
||||
@@ -398,33 +555,58 @@ let debounceTimer = null;
|
||||
|
||||
function onInput(textarea) {
|
||||
if (!active) return;
|
||||
// IME candidate window open: value isn't committed, and Enter would race with tag accept.
|
||||
if (textarea.dataset.imeActive === '1') return;
|
||||
const minChars = window.opts?.autocomplete_min_chars ?? 3;
|
||||
const info = getCurrentWord(textarea);
|
||||
if (!info || info.word.length < minChars) {
|
||||
if (!info) {
|
||||
dropdown.hide();
|
||||
return;
|
||||
}
|
||||
// Extra-network triggers have a zero threshold so `<lora:` alone surfaces results.
|
||||
const threshold = info.mode === 'tag' ? minChars : 0;
|
||||
if (info.word.length < threshold) {
|
||||
dropdown.hide();
|
||||
return;
|
||||
}
|
||||
clearTimeout(debounceTimer);
|
||||
debounceTimer = setTimeout(() => {
|
||||
const results = engine.searchAll(info.word);
|
||||
let results;
|
||||
if (info.mode === 'lora') {
|
||||
results = window.autocompleteXn ? window.autocompleteXn.searchLoras(info.word) : [];
|
||||
} else if (info.mode === 'wildcard') {
|
||||
results = window.autocompleteXn ? window.autocompleteXn.searchWildcards(info.word) : [];
|
||||
} else {
|
||||
const tagResults = engine.searchAll(info.word);
|
||||
const embedResults = window.autocompleteXn ? window.autocompleteXn.searchEmbeddings(info.word) : [];
|
||||
// Embeddings fold into tag-mode results (a1111 tagcomplete parity).
|
||||
results = [...embedResults, ...tagResults];
|
||||
}
|
||||
dropdown.show(results, textarea, info.word);
|
||||
}, 150);
|
||||
}
|
||||
|
||||
function onKeyDown(e) {
|
||||
if (!dropdown.visible) return;
|
||||
if (e.isComposing) return; // IME candidate selection, let the browser commit the candidate
|
||||
// Modifier + nav/accept keys belong to other handlers (editAttention.js on Ctrl+Arrow,
|
||||
// generate hotkey on Ctrl+Enter). Let them through even with the dropdown open.
|
||||
const hasModifier = e.ctrlKey || e.metaKey || e.altKey;
|
||||
switch (e.key) {
|
||||
case 'ArrowDown':
|
||||
if (hasModifier) return;
|
||||
e.preventDefault();
|
||||
e.stopPropagation();
|
||||
dropdown.navigate(1);
|
||||
break;
|
||||
case 'ArrowUp':
|
||||
if (hasModifier) return;
|
||||
e.preventDefault();
|
||||
e.stopPropagation();
|
||||
dropdown.navigate(-1);
|
||||
break;
|
||||
case 'Enter':
|
||||
if (hasModifier) return;
|
||||
if (dropdown.selectedIndex >= 0) {
|
||||
e.preventDefault();
|
||||
e.stopPropagation();
|
||||
@@ -432,6 +614,7 @@ function onKeyDown(e) {
|
||||
}
|
||||
break;
|
||||
case 'Tab':
|
||||
if (hasModifier) return;
|
||||
e.preventDefault();
|
||||
e.stopPropagation();
|
||||
dropdown.accept();
|
||||
@@ -450,7 +633,14 @@ function onKeyDown(e) {
|
||||
function attachAutocomplete(textarea) {
|
||||
textarea.addEventListener('input', () => onInput(textarea));
|
||||
textarea.addEventListener('keydown', onKeyDown);
|
||||
textarea.addEventListener('compositionstart', () => { textarea.dataset.imeActive = '1'; });
|
||||
textarea.addEventListener('compositionend', () => { delete textarea.dataset.imeActive; });
|
||||
textarea.addEventListener('focusin', () => {
|
||||
if (dropdown.visible && dropdown.textarea && dropdown.textarea !== textarea) dropdown.hide();
|
||||
});
|
||||
textarea.addEventListener('focusout', () => {
|
||||
// Cancel any in-flight debounced dropdown.show; otherwise it fires against a stale textarea.
|
||||
clearTimeout(debounceTimer);
|
||||
setTimeout(() => dropdown.hide(), 200);
|
||||
});
|
||||
}
|
||||
@@ -484,13 +674,21 @@ function patchActiveButton() {
|
||||
// -- Config bridge --
|
||||
|
||||
/** Monkey-patch script config bridge textboxes to push autocomplete config changes to window.opts immediately. */
|
||||
let bridgeWarnedMissingDescriptor = false;
|
||||
function patchConfigBridge() {
|
||||
const proto = Object.getOwnPropertyDescriptor(HTMLTextAreaElement.prototype, 'value');
|
||||
if (!proto?.get || !proto?.set) {
|
||||
if (!bridgeWarnedMissingDescriptor) {
|
||||
log('autoComplete', { bridge: 'skipped', reason: 'HTMLTextAreaElement.prototype.value descriptor missing' });
|
||||
bridgeWarnedMissingDescriptor = true;
|
||||
}
|
||||
return;
|
||||
}
|
||||
const elements = gradioApp().querySelectorAll('[id$="_tag_autocomplete_config_json"]');
|
||||
for (const el of elements) {
|
||||
const textarea = el.querySelector('textarea');
|
||||
if (!textarea || textarea.acBridgePatched) continue;
|
||||
textarea.acBridgePatched = true;
|
||||
const proto = Object.getOwnPropertyDescriptor(HTMLTextAreaElement.prototype, 'value');
|
||||
Object.defineProperty(textarea, 'value', {
|
||||
set(newValue) {
|
||||
const oldValue = proto.get.call(textarea);
|
||||
@@ -500,7 +698,7 @@ function patchConfigBridge() {
|
||||
const cfg = JSON.parse(newValue);
|
||||
for (const [key, val] of Object.entries(cfg)) window.opts[key] = val;
|
||||
executeCallbacks(optionsChangedCallbacks);
|
||||
} catch { /* ignore parse errors */ }
|
||||
} catch { /* ignore parse errors; the bridge is best-effort */ }
|
||||
}
|
||||
},
|
||||
get() { return proto.get.call(textarea); },
|
||||
@@ -538,6 +736,7 @@ async function initAutocomplete() {
|
||||
document.head.appendChild(style);
|
||||
dropdown.init();
|
||||
await engine.loadEnabled();
|
||||
if (window.autocompleteXn) window.autocompleteXn.loadAll();
|
||||
// Attach to all prompt textareas; even if no dictionaries loaded yet, they may be enabled later via script UI
|
||||
let attached = 0;
|
||||
PROMPT_IDS.forEach((id) => {
|
||||
@@ -560,6 +759,7 @@ async function initAutocomplete() {
|
||||
active = newActive;
|
||||
patchActiveButton();
|
||||
}
|
||||
if (window.autocompleteXn) window.autocompleteXn.loadAll();
|
||||
}
|
||||
onOptionsChanged(optionsChangedCallback);
|
||||
// Watch for config updates from the script UI bridge
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
/*
|
||||
* Extra-networks completion for SD.Next prompt textareas.
|
||||
*
|
||||
* Companion to autocomplete.js: exposes sorted indices for LoRAs, embeddings, and wildcards,
|
||||
* each backed by an existing enumeration endpoint. Dispatch and insertion are driven from
|
||||
* autocomplete.js via the mode returned by getCurrentWord().
|
||||
*
|
||||
* This file relies on globals declared in autocomplete.js (lowerBound, log, engine).
|
||||
*/
|
||||
|
||||
/* global lowerBound */
|
||||
|
||||
// -- Indices --
|
||||
|
||||
class XnIndex {
|
||||
constructor(items) {
|
||||
// items: [{ name, display }]. Sorted in-place by lowercase name.
|
||||
this.items = items.map(({ name, display }) => ({
|
||||
name: String(name).toLowerCase(),
|
||||
display: display ?? name,
|
||||
}));
|
||||
this.items.sort((a, b) => a.name.localeCompare(b.name));
|
||||
}
|
||||
|
||||
search(prefix, limit = 20) {
|
||||
const query = String(prefix).toLowerCase();
|
||||
if (!query) return [];
|
||||
const start = lowerBound(this.items, query);
|
||||
const matches = [];
|
||||
for (let i = start; i < this.items.length && matches.length < limit; i++) {
|
||||
if (!this.items[i].name.startsWith(query)) break;
|
||||
matches.push(this.items[i]);
|
||||
}
|
||||
// Substring fallback for 3+ char queries (extra-network names are usually short)
|
||||
if (matches.length === 0 && query.length >= 3) {
|
||||
for (let i = 0; i < this.items.length && matches.length < limit; i++) {
|
||||
if (this.items[i].name.includes(query)) matches.push(this.items[i]);
|
||||
}
|
||||
}
|
||||
return matches.slice(0, limit);
|
||||
}
|
||||
}
|
||||
|
||||
// -- Engine --
|
||||
|
||||
const xnEngine = {
|
||||
lora: new XnIndex([]),
|
||||
embed: new XnIndex([]),
|
||||
wildcard: new XnIndex([]),
|
||||
|
||||
async fetchJson(path) {
|
||||
try {
|
||||
const resp = await fetch(`${window.api}${path}`, { credentials: 'include' });
|
||||
if (!resp.ok) throw new Error(`${resp.status}`);
|
||||
return await resp.json();
|
||||
} catch (e) {
|
||||
log('autoComplete', { xnFetchFailed: path, error: String(e) });
|
||||
return null;
|
||||
}
|
||||
},
|
||||
|
||||
async loadAll() {
|
||||
// LoRAs: [{name, alias, path, metadata}, ...]
|
||||
const loraData = await this.fetchJson('/loras');
|
||||
if (Array.isArray(loraData)) {
|
||||
const items = [];
|
||||
for (const lo of loraData) {
|
||||
if (lo?.name) items.push({ name: lo.name });
|
||||
if (lo?.alias && lo.alias !== lo.name) items.push({ name: lo.alias });
|
||||
}
|
||||
this.lora = new XnIndex(items);
|
||||
}
|
||||
// Embeddings: {loaded: [...], skipped: [...]}
|
||||
const embData = await this.fetchJson('/embeddings');
|
||||
if (embData && typeof embData === 'object') {
|
||||
const loaded = Array.isArray(embData.loaded) ? embData.loaded : [];
|
||||
this.embed = new XnIndex(loaded.map((name) => ({ name })));
|
||||
}
|
||||
// Wildcards: [{name}, ...]
|
||||
const wcData = await this.fetchJson('/wildcards');
|
||||
if (Array.isArray(wcData)) {
|
||||
this.wildcard = new XnIndex(wcData.filter((w) => w?.name).map((w) => ({ name: w.name })));
|
||||
}
|
||||
log('autoComplete', {
|
||||
xnLoaded: true,
|
||||
lora: this.lora.items.length,
|
||||
embed: this.embed.items.length,
|
||||
wildcard: this.wildcard.items.length,
|
||||
});
|
||||
},
|
||||
|
||||
searchLoras(prefix, limit = 20) {
|
||||
return this.lora.search(prefix, limit).map((item) => ({ ...item, kind: 'lora' }));
|
||||
},
|
||||
|
||||
searchEmbeddings(prefix, limit = 20) {
|
||||
return this.embed.search(prefix, limit).map((item) => ({ ...item, kind: 'embed' }));
|
||||
},
|
||||
|
||||
searchWildcards(prefix, limit = 20) {
|
||||
return this.wildcard.search(prefix, limit).map((item) => ({ ...item, kind: 'wildcard' }));
|
||||
},
|
||||
};
|
||||
|
||||
// Expose globally so autocomplete.js can dispatch to it.
|
||||
window.autocompleteXn = xnEngine;
|
||||
@@ -86,6 +86,7 @@ class Api:
|
||||
self.add_api_route("/sdapi/v1/detailers", endpoints.get_detailers, methods=["GET"], response_model=list[models.ItemDetailer])
|
||||
self.add_api_route("/sdapi/v1/prompt-styles", endpoints.get_prompt_styles, methods=["GET"], response_model=list[models.ItemStyle])
|
||||
self.add_api_route("/sdapi/v1/embeddings", endpoints.get_embeddings, methods=["GET"], response_model=models.ResEmbeddings)
|
||||
self.add_api_route("/sdapi/v1/wildcards", endpoints.get_wildcards, methods=["GET"], response_model=list[dict], tags=["Enumerators"])
|
||||
self.add_api_route("/sdapi/v1/sd-vae", endpoints.get_sd_vaes, methods=["GET"], response_model=list[models.ItemVae])
|
||||
self.add_api_route("/sdapi/v1/extensions", endpoints.get_extensions_list, methods=["GET"], response_model=list[models.ItemExtension])
|
||||
self.add_api_route("/sdapi/v1/extra-networks", endpoints.get_extra_networks, methods=["GET"], response_model=list[models.ItemExtraNetwork])
|
||||
|
||||
@@ -6,6 +6,7 @@ are hosted on HuggingFace and downloaded on demand.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import collections
|
||||
import json
|
||||
import os
|
||||
|
||||
@@ -16,7 +17,10 @@ from modules.logger import log
|
||||
|
||||
|
||||
autocomplete_dir: str = ""
|
||||
cache: dict[str, dict] = {}
|
||||
# LRU cap. Realistic usage enables up to ~16 dictionaries at once; the bound also protects
|
||||
# against bloat when users disable/re-enable many dicts in one session.
|
||||
CACHE_MAX_ENTRIES = 16
|
||||
cache: collections.OrderedDict[str, dict] = collections.OrderedDict()
|
||||
HF_REPO = "CalamitousFelicitousness/prompt-vocab"
|
||||
HF_BASE = f"https://huggingface.co/datasets/{HF_REPO}/resolve/main"
|
||||
MANIFEST_CACHE_SEC = 300 # re-fetch manifest every 5 minutes
|
||||
@@ -49,13 +53,27 @@ def get_cached(name: str) -> dict:
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=404, detail=f"Not found: {name} ({e})") from e
|
||||
stat = os.stat(path)
|
||||
# Translations live in an optional companion file; its mtime is folded into the cache key
|
||||
# so edits to either file invalidate a stale entry.
|
||||
translations_path = os.path.join(autocomplete_dir, f"{name}.translations.json")
|
||||
translations_mtime = os.stat(translations_path).st_mtime if os.path.isfile(translations_path) else 0.0
|
||||
entry = cache.get(name)
|
||||
if entry and entry['mtime'] == stat.st_mtime:
|
||||
if entry and entry['mtime'] == stat.st_mtime and entry.get('translations_mtime', 0.0) == translations_mtime:
|
||||
cache.move_to_end(name)
|
||||
return entry
|
||||
with open(path, encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
if translations_mtime:
|
||||
try:
|
||||
with open(translations_path, encoding='utf-8') as tf:
|
||||
translations = json.load(tf)
|
||||
if isinstance(translations, dict):
|
||||
data['translations'] = translations
|
||||
except Exception as e:
|
||||
log.warning(f'Autocomplete: failed to load translations for "{name}": {e}')
|
||||
entry = {
|
||||
'mtime': stat.st_mtime,
|
||||
'translations_mtime': translations_mtime,
|
||||
'size': stat.st_size,
|
||||
'meta': {
|
||||
'name': data.get('name', name),
|
||||
@@ -69,6 +87,9 @@ def get_cached(name: str) -> dict:
|
||||
'content': data,
|
||||
}
|
||||
cache[name] = entry
|
||||
cache.move_to_end(name)
|
||||
while len(cache) > CACHE_MAX_ENTRIES:
|
||||
cache.popitem(last=False)
|
||||
return entry
|
||||
|
||||
|
||||
@@ -117,6 +138,7 @@ async def get_content(name: str) -> ItemAutocompleteContent:
|
||||
version=content.get('version', ''),
|
||||
categories=content.get('categories', {}),
|
||||
tags=content.get('tags', []),
|
||||
translations=content.get('translations'),
|
||||
)
|
||||
|
||||
|
||||
@@ -139,6 +161,9 @@ def fetch_manifest_sync() -> list[dict]:
|
||||
manifest_cache['fetched_at'] = now
|
||||
return entries
|
||||
except Exception as e:
|
||||
# Zero the timestamp so the next call retries immediately instead of serving the
|
||||
# last-known-good payload for the full 5-minute window after a transient failure.
|
||||
manifest_cache['fetched_at'] = 0
|
||||
log.warning(f"Autocomplete: Failed to fetch manifest: {e}")
|
||||
return manifest_cache.get('data', [])
|
||||
|
||||
@@ -196,7 +221,10 @@ async def list_remote() -> list[ItemAutocompleteRemote]:
|
||||
|
||||
|
||||
def download_sync(name: str) -> str:
|
||||
"""Download a tag file from HuggingFace to the local autocomplete directory."""
|
||||
"""Download a tag file from HuggingFace to the local autocomplete directory.
|
||||
If the manifest entry declares `translations: true`, fetches the `{name}.translations.json`
|
||||
companion too. Companion failure is logged but does not fail the primary download.
|
||||
"""
|
||||
import requests
|
||||
if '/' in name or '\\' in name or '..' in name:
|
||||
raise HTTPException(status_code=400, detail="Invalid name")
|
||||
@@ -217,6 +245,21 @@ def download_sync(name: str) -> str:
|
||||
os.replace(tmp, target)
|
||||
cache.pop(name, None)
|
||||
log.info(f'Autocomplete: name="{name}" url={url} ({size / 1024 / 1024:.2f}MB) downloaded')
|
||||
# Optional companion translations file. Manifest flag controls whether to attempt the download.
|
||||
manifest_entry = next((e for e in manifest_cache.get('data', []) if e.get('name') == name), None)
|
||||
if manifest_entry and manifest_entry.get('translations'):
|
||||
tr_url = f"{HF_BASE}/{name}.translations.json"
|
||||
tr_target = os.path.join(autocomplete_dir, f"{name}.translations.json")
|
||||
try:
|
||||
tr_resp = requests.get(tr_url, timeout=60)
|
||||
tr_resp.raise_for_status()
|
||||
tr_tmp = tr_target + ".tmp"
|
||||
with open(tr_tmp, 'wb') as f:
|
||||
f.write(tr_resp.content)
|
||||
os.replace(tr_tmp, tr_target)
|
||||
log.info(f'Autocomplete: name="{name}" translations downloaded')
|
||||
except Exception as e:
|
||||
log.warning(f'Autocomplete: failed to fetch translations for "{name}": {e}')
|
||||
return target
|
||||
|
||||
|
||||
|
||||
@@ -86,6 +86,11 @@ def get_embeddings():
|
||||
return models.ResEmbeddings(loaded=[], skipped=[])
|
||||
return models.ResEmbeddings(loaded=list(db.word_embeddings.keys()), skipped=list(db.skipped_embeddings.keys()))
|
||||
|
||||
def get_wildcards():
|
||||
"""List wildcard basenames (relative path with `.txt` stripped) from the configured wildcards directory."""
|
||||
from modules import ui_extra_networks_wildcards
|
||||
return [{"name": n} for n in ui_extra_networks_wildcards.list_wildcard_names()]
|
||||
|
||||
def get_extra_networks(page: str | None = None, name: str | None = None, filename: str | None = None, title: str | None = None, fullname: str | None = None, hash: str | None = None): # pylint: disable=redefined-builtin
|
||||
"""List extra networks (LoRA, checkpoints, embeddings, etc.) with optional filtering by page, name, filename, title, fullname, or hash."""
|
||||
res = []
|
||||
|
||||
@@ -521,7 +521,8 @@ class ItemAutocompleteContent(BaseModel):
|
||||
name: str = Field(title="Name", description="Autocomplete file identifier")
|
||||
version: str = Field(default="", title="Version", description="Version string")
|
||||
categories: dict = Field(default_factory=dict, title="Categories", description="Category definitions with name and color")
|
||||
tags: list = Field(default_factory=list, title="Tags", description="Tag entries as [name, category_id, post_count] tuples")
|
||||
tags: list = Field(default_factory=list, title="Tags", description="Tag entries as [name, category_id, post_count, aliases?] tuples")
|
||||
translations: Optional[dict[str, str]] = Field(default=None, title="Translations", description="Optional foreign_term -> canonical_tag_name map")
|
||||
|
||||
class ItemAutocompleteRemote(BaseModel):
|
||||
name: str = Field(title="Name", description="Autocomplete file identifier")
|
||||
|
||||
@@ -7,6 +7,20 @@ from modules.logger import log
|
||||
wildcards_list = []
|
||||
|
||||
|
||||
def list_wildcard_names() -> list[str]:
|
||||
"""Enumerate wildcard basenames (relative path, `.txt` stripped). Shared with the autocomplete API."""
|
||||
wildcards_dir = shared.opts.wildcards_dir
|
||||
if not wildcards_dir or not os.path.isdir(wildcards_dir):
|
||||
return []
|
||||
files = files_cache.list_files(wildcards_dir, ext_filter=[".txt"], recursive=True)
|
||||
names = []
|
||||
for filename in files:
|
||||
relname = os.path.relpath(filename, wildcards_dir)
|
||||
names.append(os.path.splitext(relname)[0])
|
||||
names.sort()
|
||||
return names
|
||||
|
||||
|
||||
class ExtraNetworksPageWildcards(ui_extra_networks.ExtraNetworksPage):
|
||||
def __init__(self):
|
||||
super().__init__('Wildcards')
|
||||
|
||||
Reference in New Issue
Block a user