summaryrefslogtreecommitdiff
path: root/server/szurubooru/func/posts.py
diff options
context:
space:
mode:
authorHunternif <hunternif@gmail.com>2022-10-02 03:25:10 +0100
committerHunternif <hunternif@gmail.com>2022-10-02 04:18:27 +0100
commita82fa8940158784285872434659b5cb7fb9b817a (patch)
treed0ffd2795ef2bca181801f04aa83aef9f8031527 /server/szurubooru/func/posts.py
parentbc44d8055bf3b65365bf5329a8a4ad4ac4177cfc (diff)
client,server: reverse search post by signature
Diffstat (limited to 'server/szurubooru/func/posts.py')
-rw-r--r--server/szurubooru/func/posts.py18
1 files changed, 13 insertions, 5 deletions
diff --git a/server/szurubooru/func/posts.py b/server/szurubooru/func/posts.py
index 107fc5d..fac7cd5 100644
--- a/server/szurubooru/func/posts.py
+++ b/server/szurubooru/func/posts.py
@@ -21,6 +21,7 @@ from szurubooru.func import (
users,
util,
)
+from szurubooru.func.image_hash import NpMatrix
logger = logging.getLogger(__name__)
@@ -954,8 +955,15 @@ def search_by_image_exact(image_content: bytes) -> Optional[model.Post]:
def search_by_image(image_content: bytes) -> List[Tuple[float, model.Post]]:
query_signature = image_hash.generate_signature(image_content)
- query_words = image_hash.generate_words(query_signature)
+ return search_by_signature(query_signature)
+
+def search_by_signature(
+ signature: NpMatrix,
+ limit: int = 100,
+ distance_cutoff: float = image_hash.DISTANCE_CUTOFF,
+) -> List[Tuple[float, model.Post]]:
+ query_words = image_hash.generate_words(signature)
"""
The unnest function is used here to expand one row containing the 'words'
array into multiple rows each containing a singular word.
@@ -969,10 +977,10 @@ def search_by_image(image_content: bytes) -> List[Tuple[float, model.Post]]:
FROM post_signature AS s, unnest(s.words, :q) AS a(word, query)
WHERE a.word = a.query
GROUP BY s.post_id
- ORDER BY score DESC LIMIT 100;
+ ORDER BY score DESC LIMIT :limit;
"""
- candidates = db.session.execute(dbquery, {"q": query_words})
+ candidates = db.session.execute(dbquery, {"q": query_words, "limit": limit})
data = tuple(
zip(
*[
@@ -983,13 +991,13 @@ def search_by_image(image_content: bytes) -> List[Tuple[float, model.Post]]:
)
if data:
candidate_post_ids, sigarray = data
- distances = image_hash.normalized_distance(sigarray, query_signature)
+ distances = image_hash.normalized_distance(sigarray, signature)
return [
(distance, try_get_post_by_id(candidate_post_id))
for candidate_post_id, distance in zip(
candidate_post_ids, distances
)
- if distance < image_hash.DISTANCE_CUTOFF
+ if distance < distance_cutoff
]
else:
return []