71 lines
4.9 KiB
Python
71 lines
4.9 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_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_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())
|