Persist library profiles and download history with weight capacity checks
This commit is contained in:
+48
-4
@@ -7,6 +7,7 @@ import re
|
||||
import shutil
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
|
||||
@@ -38,12 +39,47 @@ def repo_id(value):
|
||||
class Catalog:
|
||||
def __init__(self,root):
|
||||
self.root=Path(root);self.lock=threading.RLock();self.job=None;self.cancel=threading.Event()
|
||||
self.cache={};self.cache_lock=threading.Lock();self.history=[]
|
||||
path=self.root/'downloads.json'
|
||||
if path.exists():
|
||||
self.history=json.loads(path.read_text())
|
||||
for job in self.history:
|
||||
if job['state']=='downloading':job.update(state='interrupted',error='Deck wurde neu gestartet. Datei erneut auswählen und herunterladen.')
|
||||
else:
|
||||
for entry in self.root.glob('*/entry.json'):
|
||||
try:
|
||||
x=json.loads(entry.read_text())
|
||||
self.history.append(dict(id='import-'+entry.parent.name,repo=x['repo'],file=x['file'],kind=x['kind'],state='complete',bytes=x['size'],total=x['size'],created_at=x.get('downloaded_at'),error=None))
|
||||
except (OSError,ValueError,KeyError):pass
|
||||
if self.history:self._save_history()
|
||||
def _save_history(self):
|
||||
self.root.mkdir(parents=True,exist_ok=True,mode=0o700)
|
||||
temp=self.root/'downloads.tmp';temp.write_text(json.dumps(self.history));temp.replace(self.root/'downloads.json')
|
||||
def dismiss(self,job_id):
|
||||
with self.lock:
|
||||
job=next((x for x in self.history if x['id']==job_id),None)
|
||||
if not job:raise ValueError('Download nicht gefunden.')
|
||||
if job['state']=='downloading':raise ValueError('Laufenden Download zuerst abbrechen.')
|
||||
job['dismissed']=True
|
||||
if self.job and self.job.get('id')==job_id:self.job=None
|
||||
self._save_history()
|
||||
return {'dismissed':True,'model_deleted':False}
|
||||
def entry(self,ident):
|
||||
if not isinstance(ident,str) or not re.fullmatch('[a-f0-9]{64}',ident):raise ValueError('Ungültige Modelldatei-ID.')
|
||||
entry=next((x for x in self.status()['entries'] if x['id']==ident),None)
|
||||
if not entry:raise ValueError('Modelldatei nicht in der Bibliothek.')
|
||||
path=self.root/ident/('model'+PurePosixPath(entry['file']).suffix)
|
||||
if not path.is_file() or path.stat().st_size!=entry['size']:raise ValueError('Modelldatei fehlt oder ist unvollständig.')
|
||||
return entry
|
||||
def search(self,q,kind):
|
||||
if kind not in KINDS or not isinstance(q,str) or len(q)>120:raise ValueError('Ungültige Suche.')
|
||||
rows=metadata('/api/models?'+urllib.parse.urlencode(dict(search=q,filter=KINDS[kind],limit=20,sort='downloads',direction=-1)))
|
||||
return {'models':[dict(repo=x['id'],downloads=x.get('downloads'),gated=x.get('gated',False)) for x in rows]}
|
||||
def files(self,repo):
|
||||
repo=repo_id(repo)
|
||||
with self.cache_lock:
|
||||
cached=self.cache.get(repo)
|
||||
if cached and time.monotonic()-cached[0]<300:return cached[1]
|
||||
data=metadata('/api/models/'+repo+'?blobs=true')
|
||||
revision=data.get('sha','')
|
||||
if not re.fullmatch('[a-f0-9]{40}',revision):raise ValueError('Keine feste Repository-Version verfügbar.')
|
||||
@@ -54,20 +90,25 @@ class Catalog:
|
||||
size=f.get('size',f.get('lfs',{}).get('size'))
|
||||
if not isinstance(size,int) or size<1:continue
|
||||
files.append(dict(name=name,size=size,sha256=f.get('lfs',{}).get('sha256')))
|
||||
return dict(repo=repo,revision=revision,files=files,gated=data.get('gated',False),license=(data.get('cardData') or {}).get('license'),url='https://huggingface.co/'+repo)
|
||||
result=dict(repo=repo,revision=revision,files=files,gated=data.get('gated',False),license=(data.get('cardData') or {}).get('license'),url='https://huggingface.co/'+repo)
|
||||
with self.cache_lock:
|
||||
if len(self.cache)>=64:self.cache.pop(next(iter(self.cache)))
|
||||
self.cache[repo]=(time.monotonic(),result)
|
||||
return result
|
||||
def status(self):
|
||||
with self.lock:
|
||||
entries=[]
|
||||
if self.root.exists():
|
||||
for p in self.root.glob('*/entry.json'):
|
||||
try:
|
||||
item=json.loads(p.read_text());entries.append(item)
|
||||
item=json.loads(p.read_text());item['id']=p.parent.name;entries.append(item)
|
||||
except (OSError,ValueError):pass
|
||||
return dict(entries=entries,job=dict(self.job) if self.job else None,free_bytes=shutil.disk_usage(self.root if self.root.exists() else self.root.parent).free)
|
||||
return dict(entries=entries,downloads=[dict(x) for x in reversed(self.history) if not x.get("dismissed")],job=dict(self.job) if self.job else None,free_bytes=shutil.disk_usage(self.root if self.root.exists() else self.root.parent).free)
|
||||
def start(self,repo,filename,revision,kind):
|
||||
if kind not in KINDS:raise ValueError('Ungültiger Bereich.')
|
||||
with self.lock:
|
||||
if self.job and self.job['state']=='downloading':raise ValueError('Ein Download läuft bereits.')
|
||||
with self.cache_lock:self.cache.pop(repo,None)
|
||||
data=self.files(repo)
|
||||
if data['gated']:raise ValueError('Zugangsbeschränkte Modelle werden noch nicht unterstützt.')
|
||||
if revision!=data['revision']:raise ValueError('Repository wurde geändert. Dateiliste neu laden.')
|
||||
@@ -80,7 +121,8 @@ class Catalog:
|
||||
if (target/'entry.json').exists():raise ValueError('Datei bereits in der Bibliothek.')
|
||||
target.mkdir(exist_ok=True,mode=0o700)
|
||||
self.cancel.clear()
|
||||
self.job=dict(id=ident,state='downloading',repo=repo,file=filename,bytes=0,total=item['size'],error=None)
|
||||
self.job=dict(id=uuid.uuid4().hex,entry_id=ident,state='downloading',repo=repo,file=filename,kind=kind,revision=revision,created_at=time.time(),bytes=0,total=item['size'],error=None)
|
||||
self.history.append(self.job);self._save_history()
|
||||
threading.Thread(target=self._download,args=(data,item,target,kind),daemon=True).start()
|
||||
return dict(self.job)
|
||||
def stop(self):
|
||||
@@ -111,3 +153,5 @@ class Catalog:
|
||||
self.job['state']='cancelled' if isinstance(exc,InterruptedError) else 'failed'
|
||||
self.job['error']=str(exc) if isinstance(exc,ValueError) else ('Abgebrochen.' if isinstance(exc,InterruptedError) else 'Download fehlgeschlagen; Verbindung oder Anbieter prüfen.')
|
||||
partial.unlink(missing_ok=True)
|
||||
finally:
|
||||
with self.lock:self._save_history()
|
||||
|
||||
Reference in New Issue
Block a user