Enable public model catalogue and isolated file downloads; define native target
This commit is contained in:
@@ -0,0 +1,41 @@
|
||||
import hashlib
|
||||
import io
|
||||
import tempfile
|
||||
import time
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
from catalog import Catalog, safe_url, repo_id
|
||||
|
||||
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_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())
|
||||
Reference in New Issue
Block a user