.search() works for FTS4, yields dicts

Closes #198, refs #197
This commit is contained in:
Simon Willison 2020-11-06 10:23:16 -08:00
commit a5b85ebdf8
4 changed files with 164 additions and 70 deletions

View file

@ -11,6 +11,7 @@ import json
import os import os
import pathlib import pathlib
import re import re
from sqlite_fts4 import rank_bm25
import sys import sys
import textwrap import textwrap
import uuid import uuid
@ -189,6 +190,9 @@ class Database:
else: else:
register(fn) register(fn)
def register_fts4_bm25(self):
self.register_function(rank_bm25, deterministic=True)
def execute(self, sql, parameters=None): def execute(self, sql, parameters=None):
if self._tracer: if self._tracer:
self._tracer(sql, parameters) self._tracer(sql, parameters)
@ -1330,52 +1334,51 @@ class Table(Queryable):
assert fts_table, "Full-text search is not configured for table '{}'".format( assert fts_table, "Full-text search is not configured for table '{}'".format(
self.name self.name
) )
if self.db[fts_table].virtual_table_using == "FTS5": virtual_table_using = self.db[fts_table].virtual_table_using
sql = textwrap.dedent( sql = textwrap.dedent(
""" """
with {original} as ( with {original} as (
select
rowid,
{columns}
from [{dbtable}]
)
select select
{original}.*, rowid,
[{fts}].rank as {rank} {columns}
from from [{dbtable}]
[{original}] )
join [{fts}] on [{original}].rowid = [{fts}].rowid select
where {original}.*,
[{fts}] match :query {rank_implementation} as {rank}
order by from
{order} [{original}]
{limit} join [{fts_table}] on [{original}].rowid = [{fts_table}].rowid
""" where
).strip() [{fts_table}] match :query
order by
{order}
{limit}
"""
).strip()
if virtual_table_using == "FTS5":
rank_implementation = "[{}].rank".format(fts_table)
else: else:
if order == rank or order is None: self.db.register_fts4_bm25()
order = "rowid" rank_implementation = "-rank_bm25(matchinfo([{}], 'pcnalx'))".format(
sql = textwrap.dedent( fts_table
"""
select * from "{dbtable}" where rowid in (
select rowid from [{fts}]
where [{fts}] match :query
) )
order by {order}
"""
).strip()
return sql.format( return sql.format(
dbtable=self.name, dbtable=self.name,
original=original, original=original,
columns=columns_sql, columns=columns_sql,
rank=rank, rank=rank,
fts=fts_table, rank_implementation=rank_implementation,
fts_table=fts_table,
order=order or "{} desc".format(rank), order=order or "{} desc".format(rank),
limit="limit {}".format(limit) if limit else "", limit="limit {}".format(limit) if limit else "",
).strip() ).strip()
def search(self, q, order=None): def search(self, q, order=None):
return self.db.execute(self.search_sql(order=order), {"query": q}).fetchall() cursor = self.db.execute(self.search_sql(order=order), {"query": q})
columns = [c[0] for c in cursor.description]
for row in cursor:
yield dict(zip(columns, row))
def value_or_default(self, key, value): def value_or_default(self, key, value):
return self._defaults[key] if value is DEFAULT else value return self._defaults[key] if value is DEFAULT else value

View file

@ -510,14 +510,14 @@ def test_rebuild_fts(db_path, tables):
["c1", "c2", "c3"], fts_version="FTS5", create_triggers=True ["c1", "c2", "c3"], fts_version="FTS5", create_triggers=True
) )
# Search should work # Search should work
assert db["fts4_table"].search("verb1") assert list(db["fts4_table"].search("verb1"))
assert db["fts5_table"].search("verb1") assert list(db["fts5_table"].search("verb1"))
# Deleting _fts_segments to break FTS4 # Deleting _fts_segments to break FTS4
with db.conn: with db.conn:
db["fts4_table_fts_segments"].delete_where() db["fts4_table_fts_segments"].delete_where()
# Now this should error: # Now this should error:
with pytest.raises(sqlite3.DatabaseError): with pytest.raises(sqlite3.DatabaseError):
db["fts4_table"].search("verb1") list(db["fts4_table"].search("verb1"))
# Replicate docsize error from this issue for FTS5 # Replicate docsize error from this issue for FTS5
# https://github.com/simonw/sqlite-utils/issues/149 # https://github.com/simonw/sqlite-utils/issues/149
assert db["fts5_table_fts_docsize"].count == 10000 assert db["fts5_table_fts_docsize"].count == 10000
@ -530,10 +530,10 @@ def test_rebuild_fts(db_path, tables):
assert 0 == result.exit_code assert 0 == result.exit_code
fixed_tables = tables or ["fts4_table", "fts5_table"] fixed_tables = tables or ["fts4_table", "fts5_table"]
if "fts4_table" in fixed_tables: if "fts4_table" in fixed_tables:
assert db["fts4_table"].search("verb1") assert list(db["fts4_table"].search("verb1"))
else: else:
with pytest.raises(sqlite3.DatabaseError): with pytest.raises(sqlite3.DatabaseError):
db["fts4_table"].search("verb1") list(db["fts4_table"].search("verb1"))
if "fts5_table" in fixed_tables: if "fts5_table" in fixed_tables:
assert db["fts5_table_fts_docsize"].count == 10000 assert db["fts5_table_fts_docsize"].count == 10000
else: else:

View file

@ -29,9 +29,25 @@ def test_enable_fts(fresh_db):
"searchable_fts_docsize", "searchable_fts_docsize",
"searchable_fts_stat", "searchable_fts_stat",
] == fresh_db.table_names() ] == fresh_db.table_names()
assert [("tanuki are running tricksters", "Japan", "foo")] == table.search("tanuki") assert [
assert [("racoons are biting trash pandas", "USA", "bar")] == table.search("usa") {
assert [] == table.search("bar") "rowid": 1,
"text": "tanuki are running tricksters",
"country": "Japan",
"not_searchable": "foo",
"rank": 0.0,
}
] == list(table.search("tanuki"))
assert [
{
"rowid": 2,
"text": "racoons are biting trash pandas",
"country": "USA",
"not_searchable": "bar",
"rank": 0.0,
}
] == list(table.search("usa"))
assert [] == list(table.search("bar"))
def test_enable_fts_escape_table_names(fresh_db): def test_enable_fts_escape_table_names(fresh_db):
@ -49,9 +65,25 @@ def test_enable_fts_escape_table_names(fresh_db):
"http://example.com_fts_docsize", "http://example.com_fts_docsize",
"http://example.com_fts_stat", "http://example.com_fts_stat",
] == fresh_db.table_names() ] == fresh_db.table_names()
assert [("tanuki are running tricksters", "Japan", "foo")] == table.search("tanuki") assert [
assert [("racoons are biting trash pandas", "USA", "bar")] == table.search("usa") {
assert [] == table.search("bar") "rowid": 1,
"text": "tanuki are running tricksters",
"country": "Japan",
"not_searchable": "foo",
"rank": 0.0,
}
] == list(table.search("tanuki"))
assert [
{
"rowid": 2,
"text": "racoons are biting trash pandas",
"country": "USA",
"not_searchable": "bar",
"rank": 0.0,
}
] == list(table.search("usa"))
assert [] == list(table.search("bar"))
def test_enable_fts_table_names_containing_spaces(fresh_db): def test_enable_fts_table_names_containing_spaces(fresh_db):
@ -72,12 +104,21 @@ def test_populate_fts(fresh_db):
table = fresh_db["populatable"] table = fresh_db["populatable"]
table.insert(search_records[0]) table.insert(search_records[0])
table.enable_fts(["text", "country"], fts_version="FTS4") table.enable_fts(["text", "country"], fts_version="FTS4")
assert [] == table.search("trash pandas") assert [] == list(table.search("trash pandas"))
table.insert(search_records[1]) table.insert(search_records[1])
assert [] == table.search("trash pandas") assert [] == list(table.search("trash pandas"))
# Now run populate_fts to make this record available # Now run populate_fts to make this record available
table.populate_fts(["text", "country"]) table.populate_fts(["text", "country"])
assert [("racoons are biting trash pandas", "USA", "bar")] == table.search("usa") rows = list(table.search("usa"))
assert [
{
"rowid": 2,
"text": "racoons are biting trash pandas",
"country": "USA",
"not_searchable": "bar",
"rank": 0.5108256237659907,
}
] == rows
def test_populate_fts_escape_table_names(fresh_db): def test_populate_fts_escape_table_names(fresh_db):
@ -85,12 +126,20 @@ def test_populate_fts_escape_table_names(fresh_db):
table = fresh_db["http://example.com"] table = fresh_db["http://example.com"]
table.insert(search_records[0]) table.insert(search_records[0])
table.enable_fts(["text", "country"], fts_version="FTS4") table.enable_fts(["text", "country"], fts_version="FTS4")
assert [] == table.search("trash pandas") assert [] == list(table.search("trash pandas"))
table.insert(search_records[1]) table.insert(search_records[1])
assert [] == table.search("trash pandas") assert [] == list(table.search("trash pandas"))
# Now run populate_fts to make this record available # Now run populate_fts to make this record available
table.populate_fts(["text", "country"]) table.populate_fts(["text", "country"])
assert [("racoons are biting trash pandas", "USA", "bar")] == table.search("usa") assert [
{
"rowid": 2,
"text": "racoons are biting trash pandas",
"country": "USA",
"not_searchable": "bar",
"rank": 0.5108256237659907,
}
] == list(table.search("usa"))
def test_fts_tokenize(fresh_db): def test_fts_tokenize(fresh_db):
@ -103,7 +152,7 @@ def test_fts_tokenize(fresh_db):
["text", "country"], ["text", "country"],
fts_version="FTS{}".format(fts_version), fts_version="FTS{}".format(fts_version),
) )
assert [] == table.search("bite") assert [] == list(table.search("bite"))
# Test WITH stemming # Test WITH stemming
table.disable_fts() table.disable_fts()
table.enable_fts( table.enable_fts(
@ -111,9 +160,14 @@ def test_fts_tokenize(fresh_db):
fts_version="FTS{}".format(fts_version), fts_version="FTS{}".format(fts_version),
tokenize="porter", tokenize="porter",
) )
assert [("racoons are biting trash pandas", "USA", "bar")] == table.search( rows = list(table.search("bite", order="rowid"))
"bite", order="rowid" assert len(rows) == 1
) assert {
"rowid": 2,
"text": "racoons are biting trash pandas",
"country": "USA",
"not_searchable": "bar",
}.items() <= rows[0].items()
def test_optimize_fts(fresh_db): def test_optimize_fts(fresh_db):
@ -132,15 +186,34 @@ def test_optimize_fts(fresh_db):
fresh_db[table_name].optimize() fresh_db[table_name].optimize()
def test_enable_fts_w_triggers(fresh_db): def test_enable_fts_with_triggers(fresh_db):
table = fresh_db["searchable"] table = fresh_db["searchable"]
table.insert(search_records[0]) table.insert(search_records[0])
table.enable_fts(["text", "country"], fts_version="FTS4", create_triggers=True) table.enable_fts(["text", "country"], fts_version="FTS4", create_triggers=True)
assert [("tanuki are running tricksters", "Japan", "foo")] == table.search("tanuki") rows1 = list(table.search("tanuki"))
assert len(rows1) == 1
assert rows1 == [
{
"rowid": 1,
"text": "tanuki are running tricksters",
"country": "Japan",
"not_searchable": "foo",
"rank": 0.0,
}
]
table.insert(search_records[1]) table.insert(search_records[1])
# Triggers will auto-populate FTS virtual table, not need to call populate_fts() # Triggers will auto-populate FTS virtual table, not need to call populate_fts()
assert [("racoons are biting trash pandas", "USA", "bar")] == table.search("usa") rows2 = list(table.search("usa"))
assert [] == table.search("bar") assert rows2 == [
{
"rowid": 2,
"text": "racoons are biting trash pandas",
"country": "USA",
"not_searchable": "bar",
"rank": 0.0,
}
]
assert [] == list(table.search("bar"))
@pytest.mark.parametrize("create_triggers", [True, False]) @pytest.mark.parametrize("create_triggers", [True, False])
@ -183,15 +256,29 @@ def test_rebuild_fts(fresh_db, table_to_fix):
table.insert(search_records[0]) table.insert(search_records[0])
table.enable_fts(["text", "country"]) table.enable_fts(["text", "country"])
# Run a search # Run a search
assert [("tanuki are running tricksters", "Japan", "foo")] == table.search("tanuki") rows = list(table.search("tanuki"))
assert len(rows) == 1
assert {
"rowid": 1,
"text": "tanuki are running tricksters",
"country": "Japan",
"not_searchable": "foo",
}.items() <= rows[0].items()
# Delete from searchable_fts_data # Delete from searchable_fts_data
fresh_db["searchable_fts_data"].delete_where() fresh_db["searchable_fts_data"].delete_where()
# This should have broken the index # This should have broken the index
with pytest.raises(sqlite3.DatabaseError): with pytest.raises(sqlite3.DatabaseError):
table.search("tanuki") list(table.search("tanuki"))
# Running rebuild_fts() should fix it # Running rebuild_fts() should fix it
fresh_db[table_to_fix].rebuild_fts() fresh_db[table_to_fix].rebuild_fts()
assert [("tanuki are running tricksters", "Japan", "foo")] == table.search("tanuki") rows2 = list(table.search("tanuki"))
assert len(rows2) == 1
assert {
"rowid": 1,
"text": "tanuki are running tricksters",
"country": "Japan",
"not_searchable": "foo",
}.items() <= rows2[0].items()
@pytest.mark.parametrize("invalid_table", ["does_not_exist", "not_searchable"]) @pytest.mark.parametrize("invalid_table", ["does_not_exist", "not_searchable"])

View file

@ -27,10 +27,6 @@ def test_tracer():
None, None,
), ),
("select name from sqlite_master where type = 'view'", None), ("select name from sqlite_master where type = 'view'", None),
(
'select * from "dogs" where rowid in (\n select rowid from [dogs_fts]\n where [dogs_fts] match :search\n)\norder by rowid',
("Cleopaws",),
),
] ]
@ -46,17 +42,25 @@ def test_with_tracer():
assert len(collected) == 0 assert len(collected) == 0
with db.tracer(tracer): with db.tracer(tracer):
db["dogs"].search("Cleopaws") list(db["dogs"].search("Cleopaws"))
assert len(collected) == 2 assert len(collected) == 7
assert collected == [ assert collected == [
("select name from sqlite_master where type = 'view'", None), ("select name from sqlite_master where type = 'view'", None),
("select name from sqlite_master where type = 'table'", None),
("PRAGMA table_info([dogs])", None),
( (
'select * from "dogs" where rowid in (\n select rowid from [dogs_fts]\n where [dogs_fts] match :search\n)\norder by rowid', "SELECT name FROM sqlite_master\n WHERE rootpage = 0\n AND (\n sql LIKE '%VIRTUAL TABLE%USING FTS%content=%dogs%'\n OR (\n tbl_name = \"dogs\"\n AND sql LIKE '%VIRTUAL TABLE%USING FTS%'\n )\n )",
("Cleopaws",), None,
),
("select name from sqlite_master where type = 'view'", None),
("select sql from sqlite_master where name = ?", ("dogs_fts",)),
(
"with original as (\n select\n rowid,\n *\n from [dogs]\n)\nselect\n original.*,\n [dogs_fts].rank as rank\nfrom\n [original]\n join [dogs_fts] on [original].rowid = [dogs_fts].rowid\nwhere\n [dogs_fts] match :query\norder by\n rank desc",
{"query": "Cleopaws"},
), ),
] ]
# Outside the with block collected should not be appended to # Outside the with block collected should not be appended to
db["dogs"].insert({"name": "Cleopaws"}) db["dogs"].insert({"name": "Cleopaws"})
assert len(collected) == 2 assert len(collected) == 7