@@ -19,7 +19,7 @@ use crate::{
1919 Histogram , automerge, calculate_color_histogram, calculate_sticker_file_hash,
2020 fetch_sticker_file,
2121 } ,
22- util:: { Emoji , Required , decode_sticker_set_id, is_wrong_file_id_error} ,
22+ util:: { Emoji , FloatIteratorExt , Required , decode_sticker_set_id, is_wrong_file_id_error} ,
2323} ;
2424
2525#[ derive( Clone ) ]
@@ -420,32 +420,7 @@ impl ImportService {
420420
421421 // todo: tag animated?
422422 for sticker in stickers_not_in_database_yet {
423- let result = self
424- . fetch_sticker_and_save_to_db ( sticker. clone ( ) , set. name . clone ( ) )
425- . await ;
426-
427- match result {
428- Err ( InternalError :: Teloxide ( teloxide:: RequestError :: Api ( api_err) ) )
429- if is_wrong_file_id_error ( & api_err) =>
430- {
431- tracing:: warn!( "invalid file_id for a sticker, continuing" ) ;
432- }
433- Err ( InternalError :: Database ( DatabaseError :: TryingToInsertRemovedSticker ) ) => {
434- tracing:: info!( "trying to insert removed sticker" )
435- }
436- Err ( other) => {
437- if other. is_timeout_error ( ) {
438- tracing:: warn!( "sticker fetch timed out, continuing" ) ;
439- } else {
440- return Err ( other) ;
441- }
442- }
443- Ok ( ( ) ) => { }
444- }
445-
446- self . analyze_sticker ( sticker. file . unique_id . clone ( ) ) . await ?;
447- self . possibly_auto_ban_set ( & set. name , set. stickers . len ( ) )
448- . await ?;
423+ self . import_new_sticker ( sticker, set. clone ( ) ) . await ?;
449424 }
450425 for sticker in saved_stickers. clone ( ) {
451426 let Some ( s) = set. stickers . iter ( ) . find ( |s| s. file . unique_id == sticker. id ) else {
@@ -457,6 +432,16 @@ impl ImportService {
457432 . update_sticker ( sticker. id , s. file . id . clone ( ) )
458433 . await ?;
459434 }
435+ let files = self . database . get_sticker_files_by_ids ( & [ sticker. sticker_file_id ] ) . await ?;
436+ assert ! ( files. len( ) <= 1 ) ;
437+ if let Some ( file) = files. first ( ) {
438+ if file. thumbnail_file_id . is_none ( ) && s. thumb . is_some ( ) {
439+ tracing:: info!( set_id = %sticker. sticker_set_id, "found previously non-imported thumbnail" ) ;
440+ self . import_new_sticker ( s. clone ( ) , set. clone ( ) ) . await ?;
441+ }
442+ } else {
443+ self . import_new_sticker ( s. clone ( ) , set. clone ( ) ) . await ?;
444+ }
460445 }
461446
462447 let deleted_stickers = saved_stickers
@@ -493,6 +478,40 @@ impl ImportService {
493478 Ok ( ( ) )
494479 }
495480
481+
482+ async fn import_new_sticker ( & self , sticker : teloxide:: types:: Sticker , set : teloxide:: types:: StickerSet ) -> Result < ( ) , InternalError > {
483+ let result = self
484+ . fetch_sticker_and_save_to_db ( sticker. clone ( ) , set. name . clone ( ) )
485+ . await ;
486+
487+ match result {
488+ Err ( InternalError :: Teloxide ( teloxide:: RequestError :: Api ( api_err) ) )
489+ if is_wrong_file_id_error ( & api_err) =>
490+ {
491+ tracing:: warn!( "invalid file_id for a sticker, continuing" ) ;
492+ }
493+ Err ( InternalError :: Database ( DatabaseError :: TryingToInsertRemovedSticker ) ) => {
494+ tracing:: info!( "trying to insert removed sticker" )
495+ }
496+ Err ( other) => {
497+ if other. is_timeout_error ( ) {
498+ tracing:: warn!( "sticker fetch timed out, continuing" ) ;
499+ } else {
500+ return Err ( other) ;
501+ }
502+ }
503+ Ok ( ( ) ) => { }
504+ }
505+
506+ if let Err ( err) = self . analyze_sticker ( sticker. file . unique_id . clone ( ) ) . await {
507+ tracing:: error!( error = %err, "error during analyze" ) ;
508+ }
509+ if let Err ( err) = self . possibly_auto_ban_set ( & set. name , set. stickers . len ( ) ) . await {
510+ tracing:: error!( error = %err, "error during auto ban" ) ;
511+ }
512+ Ok ( ( ) )
513+ }
514+
496515 #[ tracing:: instrument( skip( self ) , err( Debug ) ) ]
497516 async fn get_clip_embedding (
498517 & self ,
@@ -599,8 +618,17 @@ impl ImportService {
599618 }
600619 let matches = self
601620 . vector_db
602- . find_banned_stickers_given_vector ( clip_vector. clone ( ) , 5 , Some ( 0.6 ) )
621+ . find_banned_stickers_given_vector ( clip_vector. clone ( ) , 5 , Some ( 0.7 ) )
603622 . await ?;
623+ let Some ( worst_score) = matches. iter ( ) . map ( |m| m. score ) . fmax ( ) else {
624+ return Ok ( false ) ;
625+ } ;
626+ let best_regular_matches = self
627+ . vector_db
628+ . find_stickers_given_vector ( clip_vector. clone ( ) , 100 , 0 , Some ( worst_score) )
629+ . await ?;
630+ let best_regular_matches_len = best_regular_matches. len ( ) ;
631+
604632 let mut should_ban = false ;
605633 for m in matches {
606634 if let Some ( banned_sticker) = self . database . get_banned_sticker ( & m. file_hash ) . await ? {
@@ -610,13 +638,11 @@ impl ImportService {
610638
611639 // if the sticker under consideration for a ban matches a banned sticker closer than
612640 // the best match from a different set, also ban
613- let best_regular_matches = self
614- . vector_db
615- . find_stickers_given_vector ( clip_vector. clone ( ) , 100 , 0 , Some ( m. score ) )
616- . await ?; // TODO: call outside loop
617- let best_regular_matches_len = best_regular_matches. len ( ) ;
618- for m in best_regular_matches {
641+ for m in & best_regular_matches {
619642 let score_non_banned_sticker = m. score ;
643+ if score_non_banned_sticker < score_banned_sticker {
644+ continue ;
645+ }
620646 if let Some ( matched_sticker) = self
621647 . database
622648 . get_some_sticker_by_file_id ( & m. file_hash )
0 commit comments