Files
Athena-Deck/test_catalog.py
T

102 lines
7.0 KiB
Python

import hashlib
import io
import tempfile
import time
import unittest
from unittest.mock import patch
from catalog import Catalog, safe_url, repo_id, is_derived_model
class CatalogTests(unittest.TestCase):
def test_reject_untrusted_urls_and_repos(self):
for value in ['http://huggingface.co/a','https://127.0.0.1/x','https://huggingface.co.evil.org/a','https://user@huggingface.co/a']:
with self.assertRaises(ValueError):safe_url(value)
for value in ['../foo','a/b/c','https://x','a/..']:
with self.assertRaises(ValueError):repo_id(value)
def test_newest_search_is_sorted_upstream_and_returns_creation_date(self):
from urllib.parse import urlsplit,parse_qs
with tempfile.TemporaryDirectory() as d,patch('catalog.metadata',return_value=[{'id':'owner/new','createdAt':'2026-09-28T00:00:00Z'}]) as fetch:
result=Catalog(d).search('qwen','image','newest')
query=parse_qs(urlsplit(fetch.call_args.args[0]).query)
self.assertEqual(query['sort'],['createdAt']);self.assertEqual(query['direction'],['-1'])
self.assertEqual(result['models'][0]['created_at'],'2026-09-28T00:00:00Z')
with self.assertRaises(ValueError):Catalog(d).search('','image','invalid')
def test_base_filter_checks_lineage_and_filters_before_limit(self):
base={'id':'maker/base','cardData':{},'tags':['text-to-video']}
for metadata in ({'cardData':{'base_model':'maker/base'}},{'cardData':{'base_model':['maker/base']}},{'tags':['base_model:quantized:maker/base']},{'tags':['lora']},{'tags':['gguf']}):
self.assertTrue(is_derived_model(dict(base,**metadata)))
self.assertFalse(is_derived_model(base));self.assertFalse(is_derived_model(dict(base,id='maker/adapter-in-name',cardData=None)))
with tempfile.TemporaryDirectory() as d,patch('catalog.metadata',return_value=[dict(base,id='derived/'+str(i),tags=['lora']) for i in range(20)]+[base]) as fetch:
value=Catalog(d).search('','video','newest',True)
self.assertEqual([x['repo'] for x in value['models']],['maker/base']);self.assertEqual(value['scanned'],21)
self.assertIn('limit=100',fetch.call_args.args[0]);self.assertIn('expand=cardData',fetch.call_args.args[0])
with self.assertRaises(ValueError):Catalog(d).search('','video','downloads','true')
def test_direct_repository_bypasses_discovery_filters_explicitly(self):
item={'id':'jpetrina/example','tags':['gguf']}
with tempfile.TemporaryDirectory() as d,patch('catalog.metadata',return_value=item):
for q in ('jpetrina/example','https://huggingface.co/jpetrina/example'):
result=Catalog(d).search(q,'chat','downloads',True)
self.assertEqual(result['models'][0]['repo'],item['id']);self.assertIn('nicht angewendet',result['notice'])
def test_gguf_filename_fallback_finds_conversational_without_pipeline_tag(self):
item={'id':'maker/model-GGUF','tags':['gguf','conversational']}
with tempfile.TemporaryDirectory() as d,patch('catalog.metadata',side_effect=[[],[item]]) as fetch:
result=Catalog(d).search('model.gguf','chat')
self.assertEqual(result['models'][0]['repo'],item['id']);self.assertIn('search=model&',fetch.call_args_list[0].args[0]);self.assertIn('filter=gguf',fetch.call_args.args[0])
def test_projector_search_includes_vision_repositories(self):
with tempfile.TemporaryDirectory() as d,patch('catalog.metadata',return_value=[{'id':'maker/vision','pipeline_tag':'image-text-to-text'}]) as fetch:
result=Catalog(d).search('vision','chat',purpose='projector')
self.assertEqual(result['models'][0]['repo'],'maker/vision')
self.assertIn('filter=gguf',fetch.call_args.args[0])
with self.assertRaises(ValueError):Catalog(d).search('vision','chat',purpose='invalid')
def test_download_integrity_and_library(self):
for valid in (True,False):
with self.subTest(valid=valid),tempfile.TemporaryDirectory() as directory:
c=Catalog(directory);payload=b'test-model-file'
data=dict(repo='owner/model',revision='a'*40,gated=False,files=[dict(name='weights.gguf',size=len(payload),sha256=hashlib.sha256(payload if valid else b'wrong').hexdigest())])
with patch.object(c,'files',return_value=data),patch('catalog.remote',return_value=io.BytesIO(payload)),patch('catalog.shutil.disk_usage',return_value=type('Disk',(),{'free':100*1024**3})()):
c.start('owner/model','weights.gguf','a'*40,'chat')
for _ in range(100):
if c.status()['job']['state']!='downloading':break
time.sleep(.01)
self.assertEqual(c.status()['job']['state'],'complete' if valid else 'failed')
self.assertEqual(len(c.status()['entries']),1 if valid else 0)
def test_changed_revision_and_disk_reserve(self):
with tempfile.TemporaryDirectory() as directory:
c=Catalog(directory);data=dict(repo='a/b',revision='b'*40,gated=False,files=[dict(name='m.gguf',size=20,sha256=None)])
with patch.object(c,'files',return_value=data):
with self.assertRaises(ValueError):c.start('a/b','m.gguf','a'*40,'chat')
with patch('catalog.shutil.disk_usage',return_value=type('Disk',(),{'free':1024})()):
with self.assertRaises(ValueError):c.start('a/b','m.gguf','b'*40,'chat')
def test_queue_progress_duplicate_and_removal(self):
import threading
release=threading.Event();progress=threading.Event()
class Stream(io.BytesIO):
def read(self,n=-1):
if self.tell()==0:
time.sleep(.3);return super().read(1)
progress.set();release.wait(3);return super().read(n)
with tempfile.TemporaryDirectory() as d:
c=Catalog(d);data=dict(repo='a/b',revision='a'*40,gated=False,files=[dict(name=n,size=2,sha256=None) for n in ['a.gguf','b.gguf','c.gguf']])
with patch.object(c,'files',return_value=data),patch('catalog.remote',side_effect=[Stream(b'aa'),io.BytesIO(b'bb')]),patch('catalog.shutil.disk_usage',return_value=type('Disk',(),{'free':100*1024**3})()):
first=c.start('a/b','a.gguf','a'*40,'chat');self.assertTrue(progress.wait(2))
status=c.status()['job'];self.assertGreater(status['bytes_per_second'],0);self.assertGreater(status['eta_seconds'],0)
second=c.start('a/b','b.gguf','a'*40,'chat');self.assertEqual(second['state'],'queued')
with self.assertRaises(ValueError):c.start('a/b','b.gguf','a'*40,'chat')
third=c.start('a/b','c.gguf','a'*40,'chat');c.dismiss(third['id'])
release.set()
for _ in range(200):
if len(c.status()['entries'])==2 and c.status()['job']['state']=='complete':break
time.sleep(.01)
self.assertEqual(len(c.status()['entries']),2);self.assertEqual(c.status()['job']['id'],second['id']);self.assertFalse(c.pending)
def test_shutdown_does_not_start_queued_file(self):
with tempfile.TemporaryDirectory() as d:
c=Catalog(d);c.pending=[('sentinel',)];c.stop(shutdown=True)
with patch('catalog.threading.Thread') as thread:c._next();thread.assert_not_called()
def test_cancellation(self):
with tempfile.TemporaryDirectory() as directory:
c=Catalog(directory);c.job=dict(state='downloading');c.stop()
from pathlib import Path
target=Path(directory)
with patch('catalog.remote',return_value=io.BytesIO(b'x')):
c._download(dict(repo='a/b',revision='a'*40),dict(name='m.gguf',size=1,sha256=None),target,'chat')
self.assertEqual(c.job['state'],'cancelled');self.assertFalse((target/'download.part').exists())